#!/usr/bin/env python3 """Freeze the byte-level corpus and window schedule for K3 AttnRes Round 05.""" from __future__ import annotations import argparse import hashlib import json import os import urllib.request from pathlib import Path from typing import Any import pyarrow.parquet as pq PROTOCOL_ID = "llm-atlas-k3-attnres-gradient-scale-v1" DATASET_REPO = "Salesforce/wikitext" DATASET_REVISION = "b08601e04326c79dfdd32d625aee71d232d685c3" DATASET_VARIANT = "wikitext-2-raw-v1" SPLITS = ("train", "validation", "test") SEEDS = (2026073001, 2026073002, 2026073003) CONTEXT = 256 FORMAL_STEPS = 8000 FORMAL_BATCH = 32 VALIDATION_WINDOWS = 64 DIAGNOSTIC_WINDOWS = 16 GATE_STEPS = (0, 1, 7999) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--cache-dir", type=Path, required=True) parser.add_argument("--manifest", type=Path, required=True) return parser.parse_args() def file_sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for block in iter(lambda: handle.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() def atomic_json(path: Path, value: dict[str, Any]) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_suffix(path.suffix + ".tmp") temporary.write_text( json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n" ) os.replace(temporary, path) def download(url: str, path: Path) -> None: if path.exists(): return path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_suffix(path.suffix + ".part") request = urllib.request.Request( url, headers={"User-Agent": "llm-atlas-k3-attnres-gradient-scale/1.0"}, ) with urllib.request.urlopen(request, timeout=120) as response: with temporary.open("wb") as output: while block := response.read(1024 * 1024): output.write(block) os.replace(temporary, path) def hashed_start(fields: list[str], corpus_length: int) -> int: value = int.from_bytes( hashlib.sha256("\0".join(fields).encode()).digest()[:8], "big" ) return value % (corpus_length - (CONTEXT + 1)) def fixed_window_start(label: str, index: int, corpus_length: int) -> int: return hashed_start([PROTOCOL_ID, label, str(index)], corpus_length) def train_window_start(seed: int, step: int, row: int, corpus_length: int) -> int: return hashed_start( [PROTOCOL_ID, "train-window", str(seed), str(step), str(row)], corpus_length, ) def concatenate_split(parquet_path: Path) -> tuple[bytes, int]: table = pq.read_table(parquet_path, columns=["text"]) rows = table.column("text").to_pylist() payload = b"".join(((row or "") + "\n").encode("utf-8") for row in rows) return payload, len(rows) def tensor_hash(payload: bytes, starts: list[int]) -> str: digest = hashlib.sha256() for start in starts: digest.update(payload[start : start + CONTEXT + 1]) return digest.hexdigest() def main() -> None: args = parse_args() args.cache_dir.mkdir(parents=True, exist_ok=True) split_manifest: dict[str, Any] = {} split_bytes: dict[str, bytes] = {} for split in SPLITS: relative = f"{DATASET_VARIANT}/{split}-00000-of-00001.parquet" url = ( f"https://huggingface.co/datasets/{DATASET_REPO}/resolve/" f"{DATASET_REVISION}/{relative}" ) parquet_path = args.cache_dir / f"{split}.parquet" download(url, parquet_path) payload, rows = concatenate_split(parquet_path) binary_path = args.cache_dir / f"{split}.bin" if not binary_path.exists() or file_sha256(binary_path) != hashlib.sha256( payload ).hexdigest(): temporary = binary_path.with_suffix(".bin.tmp") temporary.write_bytes(payload) os.replace(temporary, binary_path) split_bytes[split] = payload split_manifest[split] = { "source_path": relative, "source_url": url, "parquet_bytes": parquet_path.stat().st_size, "parquet_sha256": file_sha256(parquet_path), "rows": rows, "concatenated_bytes": len(payload), "concatenated_sha256": hashlib.sha256(payload).hexdigest(), "binary_path": str(binary_path), "binary_sha256": file_sha256(binary_path), } train = split_bytes["train"] validation = split_bytes["validation"] schedule_digest = hashlib.sha256() schedule_cells = 0 for seed in SEEDS: for step in range(1, FORMAL_STEPS + 1): for row in range(FORMAL_BATCH): start = train_window_start(seed, step, row, len(train)) schedule_digest.update(start.to_bytes(8, "big")) schedule_cells += 1 validation_starts = [ fixed_window_start("validation-window", index, len(validation)) for index in range(VALIDATION_WINDOWS) ] diagnostic_starts = [ fixed_window_start("diagnostic-window", index, len(validation)) for index in range(DIAGNOSTIC_WINDOWS) ] gate_tensor_hashes: dict[str, dict[str, str]] = {} for seed in SEEDS: gate_tensor_hashes[str(seed)] = {} for step in GATE_STEPS: starts = [ train_window_start(seed, step, row, len(train)) for row in range(FORMAL_BATCH) ] gate_tensor_hashes[str(seed)][str(step)] = tensor_hash(train, starts) manifest = { "schema_version": 1, "protocol_id": PROTOCOL_ID, "status": "frozen-before-model-output", "dataset": { "repository": DATASET_REPO, "revision": DATASET_REVISION, "variant": DATASET_VARIANT, "preprocessing": ( "parquet row order; (text or empty string) + LF; UTF-8; " "no normalization; vocabulary is raw bytes 0..255" ), "splits": split_manifest, }, "windows": { "context": CONTEXT, "target_bytes_per_window": CONTEXT, "seeds": list(SEEDS), "formal_steps": FORMAL_STEPS, "formal_batch": FORMAL_BATCH, "formal_schedule_cells": schedule_cells, "formal_schedule_sha256": schedule_digest.hexdigest(), "validation_starts": validation_starts, "validation_tensor_sha256": tensor_hash( validation, validation_starts ), "diagnostic_starts": diagnostic_starts, "diagnostic_tensor_sha256": tensor_hash( validation, diagnostic_starts ), "gate_steps": list(GATE_STEPS), "gate_training_tensor_sha256": gate_tensor_hashes, }, } atomic_json(args.manifest, manifest) print(json.dumps(manifest, ensure_ascii=False, indent=2)) if __name__ == "__main__": main()