creforge 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
creforge/__init__.py ADDED
@@ -0,0 +1,21 @@
1
+ """creforge: PII-safe synthetic credit bureau data from explicit behavioural rules."""
2
+
3
+ __version__ = "0.1.0"
4
+
5
+ from .config import Config, Profile, list_profiles, load_profile # noqa: E402
6
+ from .dataset import Dataset, DiskDataset, generate, write_dataset # noqa: E402
7
+ from .validate import Report, validate # noqa: E402
8
+
9
+ __all__ = [
10
+ "Config",
11
+ "Dataset",
12
+ "DiskDataset",
13
+ "Profile",
14
+ "Report",
15
+ "__version__",
16
+ "generate",
17
+ "list_profiles",
18
+ "load_profile",
19
+ "validate",
20
+ "write_dataset",
21
+ ]
creforge/cli.py ADDED
@@ -0,0 +1,86 @@
1
+ """Command-line interface: ``creforge generate | validate | profiles``."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import sys
7
+ import time
8
+ from pathlib import Path
9
+
10
+ import click
11
+ import yaml
12
+
13
+ from . import __version__
14
+ from .config import Config, list_profiles, load_profile
15
+ from .dataset import write_dataset
16
+ from .validate import validate
17
+
18
+
19
+ @click.group()
20
+ @click.version_option(__version__, prog_name="creforge")
21
+ def main() -> None:
22
+ """PII-safe synthetic credit bureau data from explicit behavioural rules."""
23
+
24
+
25
+ @main.command()
26
+ @click.option("--profile", "-p", default="baseline", show_default=True,
27
+ help="Built-in profile name or path to a YAML profile.")
28
+ @click.option("--subjects", "-n", type=click.IntRange(min=1), default=10_000, show_default=True)
29
+ @click.option("--months", "-m", type=click.IntRange(12, 120), default=36, show_default=True)
30
+ @click.option("--seed", "-s", type=click.IntRange(min=0), default=0, show_default=True)
31
+ @click.option("--start-month", default="2023-01", show_default=True, help="YYYY-MM")
32
+ @click.option("--out", "-o", type=click.Path(file_okay=False, path_type=Path), required=True)
33
+ @click.option("--format", "fmt", type=click.Choice(["parquet", "csv"]), default="parquet",
34
+ show_default=True)
35
+ @click.option("--workers", "-w", type=click.IntRange(min=1), default=1, show_default=True)
36
+ @click.option("--chunk-size", type=click.IntRange(min=1), default=50_000, show_default=True)
37
+ def generate(profile, subjects, months, seed, start_month, out, fmt, workers, chunk_size) -> None:
38
+ """Generate a dataset into OUT (one part file per chunk, plus manifest.json)."""
39
+ cfg = Config.from_profile(profile, subjects=subjects, months=months, seed=seed,
40
+ start_month=start_month, chunk_size=chunk_size)
41
+ t0 = time.perf_counter()
42
+ manifest = write_dataset(cfg, out, format=fmt, workers=workers)
43
+ secs = time.perf_counter() - t0
44
+ click.echo(f"Wrote {out} in {secs:.1f}s (profile={cfg.profile.name}, seed={seed})")
45
+ for table, n in manifest["row_counts"].items():
46
+ click.echo(f" {table:<14} {n:>14,}")
47
+
48
+
49
+ @main.command(name="validate")
50
+ @click.argument("path", type=click.Path(exists=True, file_okay=False, path_type=Path))
51
+ @click.option("--json", "as_json", is_flag=True, help="Emit the report as JSON.")
52
+ @click.option("--report", type=click.Path(dir_okay=False, path_type=Path),
53
+ help="Also write the Markdown report to this file.")
54
+ @click.option("--strict", is_flag=True, help="Exit non-zero on calibration failures too.")
55
+ def validate_cmd(path, as_json, report, strict) -> None:
56
+ """Check integrity and calibration of a dataset written by `generate`."""
57
+ rep = validate(path)
58
+ md = rep.to_markdown()
59
+ if report:
60
+ report.write_text(md, encoding="utf-8")
61
+ click.echo(json.dumps(rep.to_dict(), indent=2, default=str) if as_json else md)
62
+ if not rep.integrity_ok or (strict and not rep.calibration_ok):
63
+ sys.exit(1)
64
+
65
+
66
+ @main.group()
67
+ def profiles() -> None:
68
+ """Inspect built-in profiles."""
69
+
70
+
71
+ @profiles.command(name="list")
72
+ def profiles_list() -> None:
73
+ for name in list_profiles():
74
+ desc = " ".join(load_profile(name).description.split())
75
+ click.echo(f"{name:<12} {desc}")
76
+
77
+
78
+ @profiles.command(name="show")
79
+ @click.argument("name")
80
+ def profiles_show(name) -> None:
81
+ """Print the fully resolved profile (inheritance applied) as YAML."""
82
+ click.echo(yaml.safe_dump(load_profile(name).model_dump(mode="json"), sort_keys=False))
83
+
84
+
85
+ if __name__ == "__main__":
86
+ main()
creforge/config.py ADDED
@@ -0,0 +1,287 @@
1
+ """Profile and run configuration (validated with pydantic)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import json
7
+ import re
8
+ from importlib import resources
9
+ from pathlib import Path
10
+ from typing import Annotated, Any, Literal
11
+
12
+ import numpy as np
13
+ import yaml
14
+ from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
15
+
16
+ Prob = Annotated[float, Field(ge=0.0, le=1.0)]
17
+ NonNeg = Annotated[float, Field(ge=0.0)]
18
+ Positive = Annotated[float, Field(gt=0.0)]
19
+
20
+ _SUM_TOL = 1e-6
21
+ _NAME = re.compile(r"^[a-z][a-z0-9_]*$")
22
+
23
+
24
+ class _Model(BaseModel):
25
+ model_config = ConfigDict(extra="forbid", frozen=True)
26
+
27
+
28
+ def _check_shares(items: dict[str, Any], what: str) -> None:
29
+ total = sum(v.share for v in items.values())
30
+ if abs(total - 1.0) > _SUM_TOL:
31
+ raise ValueError(f"{what} shares must sum to 1.0, got {total:.6f}")
32
+
33
+
34
+ class LogNormal(_Model):
35
+ median: Positive
36
+ sigma: NonNeg
37
+
38
+
39
+ class Behaviour(_Model):
40
+ """Monthly transition parameters for one product (before grade/seasoning/macro).
41
+
42
+ ``roll[i]``: bucket i -> i+1 for C, D1..D4. ``cure[i]``: D1..D5 -> current.
43
+ ``back[i]``: D2..D5 -> one bucket lower. ``restructure``: D2+ -> RS.
44
+ """
45
+
46
+ roll: list[Prob] = Field(min_length=5, max_length=5)
47
+ cure: list[Prob] = Field(min_length=5, max_length=5)
48
+ back: list[Prob] = Field(min_length=4, max_length=4)
49
+ restructure: Prob = 0.0
50
+ rs_cure: Prob = 0.0
51
+ rs_redefault: Prob = 0.0
52
+ close: Prob = 0.0
53
+
54
+ @model_validator(mode="after")
55
+ def _rows_are_substochastic(self) -> Behaviour:
56
+ rows = {
57
+ "C": self.roll[0] + self.close,
58
+ "D1": self.roll[1] + self.cure[0],
59
+ "D2": self.roll[2] + self.cure[1] + self.back[0] + self.restructure,
60
+ "D3": self.roll[3] + self.cure[2] + self.back[1] + self.restructure,
61
+ "D4": self.roll[4] + self.cure[3] + self.back[2] + self.restructure,
62
+ "D5": self.cure[4] + self.back[3] + self.restructure,
63
+ "RS": self.rs_cure + self.rs_redefault,
64
+ }
65
+ bad = {k: round(v, 6) for k, v in rows.items() if v > 1.0 + _SUM_TOL}
66
+ if bad:
67
+ raise ValueError(f"outflow probabilities exceed 1.0 for states {bad}")
68
+ return self
69
+
70
+
71
+ class Product(_Model):
72
+ kind: Literal["installment", "revolving"]
73
+ holding_rate: NonNeg = Field(description="Expected open accounts per subject at window start")
74
+ inquiry_share: NonNeg = Field(description="Share of new inquiries for this product")
75
+ amount: LogNormal = Field(description="Principal (installment) or limit (revolving)")
76
+ round_to: Positive = 100.0
77
+ tenor_months: list[Annotated[int, Field(ge=1)]] | None = None
78
+ max_age_months: Annotated[int, Field(ge=0)] = 240
79
+ interest_rate: Annotated[float, Field(ge=0.0, le=1.0)]
80
+ rate_grade_sensitivity: NonNeg = 1.0
81
+ secured: bool = False
82
+ min_payment_pct: Prob | None = None
83
+ behaviour: Behaviour
84
+
85
+ @model_validator(mode="after")
86
+ def _kind_consistency(self) -> Product:
87
+ if self.kind == "installment":
88
+ if not self.tenor_months:
89
+ raise ValueError("installment products need tenor_months")
90
+ else:
91
+ if self.min_payment_pct is None or self.min_payment_pct <= 0:
92
+ raise ValueError("revolving products need min_payment_pct > 0")
93
+ if self.behaviour.restructure > 0:
94
+ raise ValueError("restructure is only supported for installment products")
95
+ return self
96
+
97
+
98
+ class Grade(_Model):
99
+ share: Prob
100
+ roll_mult: Positive
101
+ cure_mult: Positive
102
+ inquiry_rate: NonNeg = Field(description="Expected inquiries per subject per month")
103
+ approval: Annotated[float, Field(gt=0.0, lt=1.0)]
104
+ utilization: Prob
105
+ transactor_share: Prob
106
+ rate_spread: NonNeg
107
+
108
+
109
+ class IncomeBand(_Model):
110
+ share: Prob
111
+ amount_mult: Positive
112
+ holding_mult: NonNeg
113
+
114
+
115
+ class Population(_Model):
116
+ age_years: tuple[int, int] = (21, 70)
117
+ regions: Annotated[int, Field(ge=1, le=99)] = 13
118
+ income_bands: dict[str, IncomeBand]
119
+
120
+ @model_validator(mode="after")
121
+ def _valid(self) -> Population:
122
+ lo, hi = self.age_years
123
+ if not 18 <= lo <= hi <= 100:
124
+ raise ValueError("age_years must satisfy 18 <= min <= max <= 100")
125
+ _check_shares(self.income_bands, "income band")
126
+ return self
127
+
128
+
129
+ class Seasoning(_Model):
130
+ """Roll-rate multiplier by months on book: floor at 0, peak at peak_age, tail after."""
131
+
132
+ floor: Positive = 0.5
133
+ peak: Positive = 1.35
134
+ peak_age: Annotated[int, Field(ge=1)] = 15
135
+ tail: Positive = 0.85
136
+
137
+ def curve(self, age: np.ndarray) -> np.ndarray:
138
+ a = np.maximum(np.asarray(age, dtype=np.float64), 0.0) / self.peak_age
139
+ hump = a * np.exp(1.0 - a)
140
+ base = np.where(a <= 1.0, self.floor, self.tail)
141
+ return base + (self.peak - base) * hump
142
+
143
+
144
+ class Macro(_Model):
145
+ """Piecewise-linear roll multiplier by window month; flat beyond the given points."""
146
+
147
+ points: list[tuple[Annotated[int, Field(ge=0)], Positive]] = [(0, 1.0)]
148
+
149
+ @field_validator("points")
150
+ @classmethod
151
+ def _sorted(cls, v: list[tuple[int, float]]) -> list[tuple[int, float]]:
152
+ if not v:
153
+ raise ValueError("macro.points must not be empty")
154
+ months = [m for m, _ in v]
155
+ if months != sorted(set(months)):
156
+ raise ValueError("macro.points months must be strictly increasing")
157
+ return v
158
+
159
+ def path(self, months: int) -> np.ndarray:
160
+ xs = np.array([m for m, _ in self.points], dtype=np.float64)
161
+ ys = np.array([y for _, y in self.points], dtype=np.float64)
162
+ return np.interp(np.arange(months, dtype=np.float64), xs, ys)
163
+
164
+ @property
165
+ def is_flat(self) -> bool:
166
+ return len({y for _, y in self.points}) == 1
167
+
168
+
169
+ class Inquiries(_Model):
170
+ recent_window_months: Annotated[int, Field(ge=1)] = 6
171
+ recent_penalty: NonNeg = Field(0.35, description="Logit penalty per recent inquiry")
172
+ withdrawn_share: Prob = 0.05
173
+
174
+
175
+ class Target(_Model):
176
+ dpd30_share: tuple[Prob, Prob]
177
+ annual_writeoff_rate: tuple[Prob, Prob]
178
+
179
+
180
+ class Profile(_Model):
181
+ name: str
182
+ description: str = ""
183
+ sources: list[str] = []
184
+ population: Population
185
+ grades: dict[str, Grade]
186
+ products: dict[str, Product]
187
+ seasoning: Seasoning = Seasoning()
188
+ macro: Macro = Macro()
189
+ inquiries: Inquiries = Inquiries()
190
+ lenders: Annotated[int, Field(ge=1, le=9999)] = 40
191
+ writeoff_after_months: Annotated[int, Field(ge=1, le=36)] = Field(
192
+ 6, description="Months spent in the 120+ bucket before write-off"
193
+ )
194
+ targets: dict[str, Target] = {}
195
+
196
+ @model_validator(mode="after")
197
+ def _valid(self) -> Profile:
198
+ _check_shares(self.grades, "grade")
199
+ if len(self.grades) < 2:
200
+ raise ValueError("at least two grades are required")
201
+ if not self.products:
202
+ raise ValueError("at least one product is required")
203
+ for name in self.products:
204
+ if not _NAME.match(name):
205
+ raise ValueError(f"product name {name!r} must be snake_case")
206
+ inq = sum(p.inquiry_share for p in self.products.values())
207
+ if abs(inq - 1.0) > _SUM_TOL:
208
+ raise ValueError(f"product inquiry_share must sum to 1.0, got {inq:.6f}")
209
+ unknown = set(self.targets) - set(self.products)
210
+ if unknown:
211
+ raise ValueError(f"targets for unknown products: {sorted(unknown)}")
212
+ return self
213
+
214
+
215
+ class Config(_Model):
216
+ profile: Profile
217
+ subjects: Annotated[int, Field(ge=1)]
218
+ months: Annotated[int, Field(ge=12, le=120)] = 36
219
+ seed: Annotated[int, Field(ge=0)] = 0
220
+ start_month: str = "2023-01"
221
+ chunk_size: Annotated[int, Field(ge=1)] = 50_000
222
+
223
+ @field_validator("start_month")
224
+ @classmethod
225
+ def _month(cls, v: str) -> str:
226
+ if not re.fullmatch(r"\d{4}-(0[1-9]|1[0-2])", v):
227
+ raise ValueError("start_month must look like YYYY-MM")
228
+ return v
229
+
230
+ @classmethod
231
+ def from_profile(cls, profile: str | Path = "baseline", **kwargs: Any) -> Config:
232
+ return cls(profile=load_profile(profile), **kwargs)
233
+
234
+ @property
235
+ def n_chunks(self) -> int:
236
+ return -(-self.subjects // self.chunk_size)
237
+
238
+ def chunk_bounds(self, chunk: int) -> tuple[int, int]:
239
+ start = chunk * self.chunk_size
240
+ return start, min(start + self.chunk_size, self.subjects)
241
+
242
+ def sha256(self) -> str:
243
+ payload = self.model_dump(mode="json", exclude={"chunk_size"})
244
+ blob = json.dumps(payload, sort_keys=True, separators=(",", ":"))
245
+ return hashlib.sha256(blob.encode()).hexdigest()
246
+
247
+
248
+ # --- profile loading -----------------------------------------------------------------
249
+
250
+
251
+ def _deep_merge(base: dict[str, Any], over: dict[str, Any]) -> dict[str, Any]:
252
+ out = dict(base)
253
+ for k, v in over.items():
254
+ out[k] = _deep_merge(out[k], v) if isinstance(v, dict) and isinstance(out.get(k), dict) else v
255
+ return out
256
+
257
+
258
+ def _read_profile_dict(ref: str | Path, _seen: tuple[str, ...] = ()) -> dict[str, Any]:
259
+ path = Path(ref)
260
+ if path.suffix in {".yaml", ".yml"} and path.exists():
261
+ text, key, here = path.read_text(encoding="utf-8"), str(path.resolve()), path.parent
262
+ else:
263
+ name = str(ref)
264
+ if name not in list_profiles():
265
+ raise FileNotFoundError(f"unknown profile {name!r}; built-ins: {list_profiles()}")
266
+ text = resources.files("creforge.profiles").joinpath(f"{name}.yaml").read_text("utf-8")
267
+ key, here = f"builtin:{name}", None
268
+ if key in _seen:
269
+ raise ValueError(f"circular profile inheritance: {' -> '.join(_seen + (key,))}")
270
+ data = yaml.safe_load(text) or {}
271
+ parent = data.pop("extends", None)
272
+ if parent is None:
273
+ return data
274
+ parent_ref: str | Path = parent
275
+ if here is not None and (here / parent).exists():
276
+ parent_ref = here / parent
277
+ return _deep_merge(_read_profile_dict(parent_ref, _seen + (key,)), data)
278
+
279
+
280
+ def load_profile(ref: str | Path) -> Profile:
281
+ """Load a built-in profile by name, or a YAML file by path (``extends:`` supported)."""
282
+ return Profile.model_validate(_read_profile_dict(ref))
283
+
284
+
285
+ def list_profiles() -> list[str]:
286
+ root = resources.files("creforge.profiles")
287
+ return sorted(p.name[:-5] for p in root.iterdir() if p.name.endswith(".yaml"))
creforge/dataset.py ADDED
@@ -0,0 +1,170 @@
1
+ """Public generation API: in-memory ``Dataset`` and streaming ``write_dataset``."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import multiprocessing
7
+ from collections.abc import Iterator
8
+ from concurrent.futures import ProcessPoolExecutor
9
+ from dataclasses import dataclass
10
+ from pathlib import Path
11
+ from typing import Literal
12
+
13
+ import polars as pl
14
+
15
+ from . import __version__
16
+ from .config import Config
17
+ from .generator import TABLES, generate_chunk, make_context
18
+
19
+ Format = Literal["parquet", "csv"]
20
+ MANIFEST = "manifest.json"
21
+ PRIVACY_STATEMENT = (
22
+ "Generated entirely from explicit behavioural rules and random draws. No real, "
23
+ "row-level or personal data was read or used; identifiers are keyed permutations "
24
+ "of row counters. Any resemblance to real persons or accounts is coincidental."
25
+ )
26
+
27
+
28
+ @dataclass
29
+ class Dataset:
30
+ """Generated tables, kept per chunk so validation can stream over them."""
31
+
32
+ config: Config
33
+ chunks: list[dict[str, pl.DataFrame]]
34
+
35
+ def table(self, name: str) -> pl.DataFrame:
36
+ return pl.concat([c[name] for c in self.chunks])
37
+
38
+ @property
39
+ def subject(self) -> pl.DataFrame:
40
+ return self.table("subject")
41
+
42
+ @property
43
+ def inquiry(self) -> pl.DataFrame:
44
+ return self.table("inquiry")
45
+
46
+ @property
47
+ def account(self) -> pl.DataFrame:
48
+ return self.table("account")
49
+
50
+ @property
51
+ def account_month(self) -> pl.DataFrame:
52
+ return self.table("account_month")
53
+
54
+ def iter_chunks(self) -> Iterator[dict[str, pl.DataFrame]]:
55
+ yield from self.chunks
56
+
57
+ def write(self, out: str | Path, format: Format = "parquet") -> dict:
58
+ out = Path(out)
59
+ counts = [_write_chunk(out, i, c, format) for i, c in enumerate(self.chunks)]
60
+ return _write_manifest(out, self.config, format, counts)
61
+
62
+
63
+ def generate(config: Config) -> Dataset:
64
+ """Generate all chunks in memory. Use :func:`write_dataset` for large runs."""
65
+ ctx = make_context(config)
66
+ return Dataset(config, [generate_chunk(config, i, ctx) for i in range(config.n_chunks)])
67
+
68
+
69
+ def write_dataset(
70
+ config: Config, out: str | Path, format: Format = "parquet", workers: int = 1
71
+ ) -> dict:
72
+ """Generate chunk by chunk straight to disk; memory stays bounded by one chunk per worker.
73
+
74
+ Output is byte-identical for a given config regardless of ``workers``.
75
+ """
76
+ out = Path(out)
77
+ out.mkdir(parents=True, exist_ok=True)
78
+ jobs = range(config.n_chunks)
79
+ if workers <= 1 or config.n_chunks == 1:
80
+ ctx = make_context(config)
81
+ counts = [_write_chunk(out, i, generate_chunk(config, i, ctx), format) for i in jobs]
82
+ else:
83
+ # "spawn": forking after Polars' thread pool has started can deadlock.
84
+ spawn = multiprocessing.get_context("spawn")
85
+ with ProcessPoolExecutor(max_workers=workers, mp_context=spawn) as pool:
86
+ counts = list(pool.map(_chunk_job, [(config, i, str(out), format) for i in jobs]))
87
+ return _write_manifest(out, config, format, counts)
88
+
89
+
90
+ def _chunk_job(args: tuple[Config, int, str, Format]) -> dict[str, int]:
91
+ config, chunk, out, format = args
92
+ return _write_chunk(Path(out), chunk, generate_chunk(config, chunk), format)
93
+
94
+
95
+ def part_path(out: Path, table: str, chunk: int, format: Format) -> Path:
96
+ return out / table / f"part-{chunk:05d}.{format}"
97
+
98
+
99
+ def _write_chunk(out: Path, chunk: int, frames: dict[str, pl.DataFrame], format: Format) -> dict:
100
+ for name in TABLES:
101
+ path = part_path(out, name, chunk, format)
102
+ path.parent.mkdir(parents=True, exist_ok=True)
103
+ df = frames[name]
104
+ if format == "parquet":
105
+ df.write_parquet(path, compression="zstd", statistics=True)
106
+ else:
107
+ df.write_csv(path)
108
+ return {name: frames[name].height for name in TABLES}
109
+
110
+
111
+ def _write_manifest(out: Path, config: Config, format: Format, counts: list[dict]) -> dict:
112
+ out.mkdir(parents=True, exist_ok=True)
113
+ manifest = {
114
+ "creforge_version": __version__,
115
+ "format": format,
116
+ "chunks": config.n_chunks,
117
+ "config_sha256": config.sha256(),
118
+ "row_counts": {t: sum(c[t] for c in counts) for t in TABLES},
119
+ "privacy": PRIVACY_STATEMENT,
120
+ "config": config.model_dump(mode="json"),
121
+ }
122
+ (out / MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
123
+ return manifest
124
+
125
+
126
+ @dataclass
127
+ class DiskDataset:
128
+ """A dataset previously written by :func:`write_dataset`, read one chunk at a time."""
129
+
130
+ path: Path
131
+ manifest: dict
132
+
133
+ @classmethod
134
+ def open(cls, path: str | Path) -> DiskDataset:
135
+ path = Path(path)
136
+ return cls(path, json.loads((path / MANIFEST).read_text(encoding="utf-8")))
137
+
138
+ @property
139
+ def config(self) -> Config:
140
+ return Config.model_validate(self.manifest["config"])
141
+
142
+ def iter_chunks(self) -> Iterator[dict[str, pl.DataFrame]]:
143
+ fmt = self.manifest["format"]
144
+ for i in range(self.manifest["chunks"]):
145
+ yield {t: _read(part_path(self.path, t, i, fmt), fmt) for t in TABLES}
146
+
147
+
148
+ # CSV loses types; restore the non-string columns explicitly instead of guessing.
149
+ _CSV_TYPES: dict[str, pl.DataType] = {
150
+ **dict.fromkeys(("created_month", "inquiry_date", "open_date", "close_date", "as_of_month"),
151
+ pl.Date()),
152
+ **dict.fromkeys(("requested_amount", "credit_limit", "principal", "interest_rate", "balance",
153
+ "amount_due", "amount_paid"), pl.Float64()),
154
+ **dict.fromkeys(("birth_year", "tenor_months", "months_in_arrears"), pl.Int16()),
155
+ "secured": pl.Boolean(),
156
+ }
157
+
158
+
159
+ def _read(path: Path, fmt: str) -> pl.DataFrame:
160
+ if fmt == "parquet":
161
+ return pl.read_parquet(path)
162
+ df = pl.read_csv(path, infer_schema=False)
163
+ def restore(c: str, t: pl.DataType) -> pl.Expr:
164
+ if t == pl.Date():
165
+ return pl.col(c).str.to_date()
166
+ if t == pl.Boolean():
167
+ return pl.col(c) == "true"
168
+ return pl.col(c).cast(t)
169
+
170
+ return df.with_columns(restore(c, t) for c, t in _CSV_TYPES.items() if c in df.columns)