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 +21 -0
- creforge/cli.py +86 -0
- creforge/config.py +287 -0
- creforge/dataset.py +170 -0
- creforge/engine.py +380 -0
- creforge/generator.py +393 -0
- creforge/ids.py +61 -0
- creforge/profiles/__init__.py +0 -0
- creforge/profiles/baseline.yaml +153 -0
- creforge/profiles/stressed.yaml +20 -0
- creforge/validate.py +314 -0
- creforge-0.1.0.dist-info/METADATA +183 -0
- creforge-0.1.0.dist-info/RECORD +16 -0
- creforge-0.1.0.dist-info/WHEEL +4 -0
- creforge-0.1.0.dist-info/entry_points.txt +2 -0
- creforge-0.1.0.dist-info/licenses/LICENSE +202 -0
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)
|