Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
s3_contract.py298 linesDownload Raw Back to pipeline
1"""s3_contract.py — THE single dataset layout + manifest (finding V8/D-7/D-8).2 3Supersedes BOTH prior contracts: design-F1's `runs/<id>/{sft_corpus,dpo_pairs,4rl_task_pool,divergence_pairs,wm_tuples,holdout,diloco_rendezvous}` and5design-F2's `{traces,tasks,replay,task_grades,corpus}/v1/run_id=<id>` — the two6were never reconciled and coexisted in the grounding doc. One layout, one7manifest, two explicit serializers with a unit-tested leak guard.8 9Deliberate exclusions from the run layout:10  * `diloco_rendezvous/` — training-comms state, not dataset; lives in its own11    prefix/bucket (finding D-19).12  * `wm_tuples/` — emitted only when the P4 world-model ablation is scheduled13    (finding D-14); not part of Stage 0.14 15Layout (root = any local path or fsspec URI):16    <root>/runs/<run_id>/17        tasks/manifest.jsonl       policy-safe task rows (golden_diff -> sha256)18        tasks_full/manifest.jsonl  construction-side full rows (RESTRICTED prefix)19        traj/*.jsonl               CanonicalTrajectory records (audit trail)20        corpus_sft/rows.jsonl      admitted SFT rows (to_policy_row output)21        corpus_dpo/rows.jsonl      DPO-candidate rows22        holdout/tasks.jsonl        held-out task ids+rows (never rolled out)23        quarantine/*.jsonl         rejected trajectories w/ reasons (audit)24        manifest.json              RunManifest25        DATASET_CARD.md            human-readable card26"""27from __future__ import annotations28 29import dataclasses30import hashlib31import json32from dataclasses import dataclass, field33from typing import IO, Iterable34 35from composer_replication.datagen.schema import FeatureDeletionTask36 37SCHEMA_VERSION = "1"38 39 40def _is_local(root: str) -> bool:41    return "://" not in root or root.startswith("file://")42 43 44def _open(path: str, mode: str = "w") -> IO[str]:45    """Open a path for text IO; plain `open` locally, fsspec for s3:// etc.46 47    fsspec is lazy so the module (and all local-corpus runs) need no extra dep.48    """49    if _is_local(path):50        import os51        local = path.removeprefix("file://")52        os.makedirs(os.path.dirname(local), exist_ok=True)53        return open(local, mode, encoding="utf-8")54    try:55        import fsspec  # noqa: PLC0415 — lazy heavy dep56    except ImportError as e:57        raise RuntimeError(58            "Non-local corpus roots require fsspec; install with "59            "`pip install -e .[serverless]`. Got: " + repr(e)60        ) from e61    return fsspec.open(path, mode, encoding="utf-8").open()62 63 64def _exists(path: str) -> bool:65    if _is_local(path):66        import os67        return os.path.exists(path.removeprefix("file://"))68    import fsspec  # noqa: PLC041569    fs, _, paths = fsspec.get_fs_token_paths(path)70    return bool(fs.exists(paths[0]))71 72 73@dataclass(frozen=True)74class RunLayout:75    """Pure-path logic for one run's prefixes — testable without any IO."""76 77    root: str78    run_id: str79 80    def __post_init__(self) -> None:81        # Defense-in-depth (Wave-21 review P2): run_id is operator-supplied,82        # but a separator or `..` would silently escape the corpus root.83        if not self.run_id or "/" in self.run_id or "\\" in self.run_id \84                or ".." in self.run_id:85            raise ValueError(86                f"run_id {self.run_id!r} must be a single non-empty path "87                "segment (no separators, no '..')."88            )89 90    def _p(self, *parts: str) -> str:91        base = self.root.rstrip("/")92        return f"{base}/runs/{self.run_id}/" + "/".join(parts)93 94    @property95    def tasks_path(self) -> str:96        return self._p("tasks", "manifest.jsonl")97 98    @property99    def tasks_full_path(self) -> str:100        # RESTRICTED prefix: carries golden_diff/deleted_symbols. On S3 this101        # prefix gets a deny-by-default policy; locally it is still separated102        # so a naive `corpus_*` glob can never sweep it up.103        return self._p("tasks_full", "manifest.jsonl")104 105    @property106    def traj_path(self) -> str:107        return self._p("traj", "trajectories.jsonl")108 109    @property110    def sft_path(self) -> str:111        return self._p("corpus_sft", "rows.jsonl")112 113    @property114    def dpo_path(self) -> str:115        return self._p("corpus_dpo", "rows.jsonl")116 117    @property118    def holdout_path(self) -> str:119        return self._p("holdout", "tasks.jsonl")120 121    @property122    def quarantine_path(self) -> str:123        return self._p("quarantine", "rejected.jsonl")124 125    @property126    def manifest_path(self) -> str:127        return self._p("manifest.json")128 129    @property130    def card_path(self) -> str:131        return self._p("DATASET_CARD.md")132 133 134@dataclass135class RunManifest:136    """Run-level metadata: counts, cost, lineage, budget, acceptance status.137 138    `created_at` is CALLER-passed (never datetime.now() in here) so manifests139    are reproducible in tests. `parent_run_id` threads flywheel lineage so140    cross-generation dedup (finding D-12) can find prior signatures.141    """142 143    run_id: str144    created_at: str145    source: str = ""146    counts: dict = field(default_factory=dict)147    cost_usd: float = 0.0148    parent_run_id: str | None = None149    schema_version: str = SCHEMA_VERSION150    status: str = "building"          # building | accepted | rejected | partial151    budget_usd: float | None = None152 153    def spend(self, usd: float) -> None:154        self.cost_usd += usd155 156    @property157    def over_budget(self) -> bool:158        return self.budget_usd is not None and self.cost_usd >= self.budget_usd159 160    def write(self, layout: RunLayout) -> None:161        with _open(layout.manifest_path) as f:162            json.dump(dataclasses.asdict(self), f, indent=2)163 164    @classmethod165    def read(cls, layout: RunLayout) -> RunManifest:166        with _open(layout.manifest_path, "r") as f:167            return cls(**json.load(f))168 169 170# ---------------------------------------------------------------------171# Writers — the leak guard lives here (finding D-8)172# ---------------------------------------------------------------------173 174 175def _task_row_policy_safe(task: FeatureDeletionTask) -> dict:176    """Task row with the construction-side secrets REPLACED, not just hidden.177 178    `asdict()` includes `golden_diff` despite `repr=False` — that is exactly179    the leak D-8 flagged. We keep provenance via a sha256 (verifiable, not180    recoverable) and drop `deleted_symbols` entirely (they name the answer).181    """182    row = dataclasses.asdict(task)183    gold = row.pop("golden_diff", "")184    row.pop("deleted_symbols", None)185    row["golden_diff_sha256"] = hashlib.sha256(gold.encode()).hexdigest() if gold else ""186    return row187 188 189def write_tasks(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:190    """Write the POLICY-SAFE task manifest (the default everything reads)."""191    n = 0192    with _open(layout.tasks_path) as f:193        for t in tasks:194            f.write(json.dumps(_task_row_policy_safe(t)) + "\n")195            n += 1196    return n197 198 199def write_tasks_full(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:200    """Write FULL task rows (incl. golden_diff) to the RESTRICTED prefix.201 202    Only the validator/monitor side reads this; never corpus consumers.203    """204    n = 0205    with _open(layout.tasks_full_path) as f:206        for t in tasks:207            f.write(json.dumps(dataclasses.asdict(t)) + "\n")208            n += 1209    return n210 211 212def _write_jsonl(path: str, rows: Iterable[dict]) -> int:213    n = 0214    with _open(path) as f:215        for r in rows:216            f.write(json.dumps(r) + "\n")217            n += 1218    return n219 220 221def write_sft_rows(layout: RunLayout, rows: Iterable[dict]) -> int:222    return _write_jsonl(layout.sft_path, rows)223 224 225def write_dpo_rows(layout: RunLayout, rows: Iterable[dict]) -> int:226    return _write_jsonl(layout.dpo_path, rows)227 228 229def write_quarantine(layout: RunLayout, rows: Iterable[dict]) -> int:230    return _write_jsonl(layout.quarantine_path, rows)231 232 233def write_holdout(layout: RunLayout, tasks: Iterable[FeatureDeletionTask]) -> int:234    return _write_jsonl(layout.holdout_path, (_task_row_policy_safe(t) for t in tasks))235 236 237def write_trajectories(layout: RunLayout, rows: Iterable[dict]) -> int:238    return _write_jsonl(layout.traj_path, rows)239 240 241def write_dataset_card(layout: RunLayout, manifest: RunManifest,242                       *, license_tiers: dict[str, int] | None = None,243                       dedup_stats: dict | None = None,244                       decontamination_note: str = "") -> None:245    """A small human-readable dataset card (finding D-18)."""246    lines = [247        f"# Dataset card — run `{manifest.run_id}`",248        "",249        f"- **created:** {manifest.created_at}",250        f"- **source:** {manifest.source}",251        f"- **status:** {manifest.status}",252        f"- **schema_version:** {manifest.schema_version}",253        f"- **cost (USD):** {manifest.cost_usd:.2f}"254        + (f" / budget {manifest.budget_usd:.2f}" if manifest.budget_usd else ""),255        f"- **lineage:** parent_run_id={manifest.parent_run_id or 'none'}",256        "",257        "## Counts",258        "",259    ]260    for k, v in sorted(manifest.counts.items()):261        lines.append(f"- {k}: {v}")262    if license_tiers:263        lines += ["", "## License tiers seen", ""]264        lines += [f"- {k}: {v}" for k, v in sorted(license_tiers.items())]265    lines += ["", "## Decontamination", "",266              decontamination_note or267              "All source repos checked against the SWE-bench-family eval list "268              "(datagen.repo_gate.DECONTAMINATION_LIST) at ingest."]269    if dedup_stats:270        lines += ["", "## Dedup", ""]271        lines += [f"- {k}: {v}" for k, v in sorted(dedup_stats.items())]272    lines += ["", "Policy-safe rows only: `golden_diff` is sha256-hashed and "273              "`deleted_symbols` dropped in `tasks/`, `corpus_*/`, `holdout/` "274              "(full rows live in the restricted `tasks_full/`).", ""]275    with _open(layout.card_path) as f:276        f.write("\n".join(lines))277 278 279def manifest_exists(layout: RunLayout) -> bool:280    """Write-once guard for the driver (finding D-21 idempotency)."""281    return _exists(layout.manifest_path)282 283 284__all__ = [285    "SCHEMA_VERSION",286    "RunLayout",287    "RunManifest",288    "manifest_exists",289    "write_dataset_card",290    "write_dpo_rows",291    "write_holdout",292    "write_quarantine",293    "write_sft_rows",294    "write_tasks",295    "write_tasks_full",296    "write_trajectories",297]298