Team Ai
Modelpublic

physicsrob/torchwright-doom-e1m1

sourceHugging Faceupdated 2mo agoView on Hugging Face
3likes34downloads
pretty_text.py343 linesDownload Raw Back to tools
1"""Pure-stdlib, bundle-driven Doom text prettifier.2 3This file is also copied into published bundles as ``tools/pretty_text.py``.4It therefore must not import TorchWright, torch, Transformers, or5``torchwright_doom``.6"""7 8from __future__ import annotations9 10import argparse11import hashlib12import json13import sys14from pathlib import Path15 16 17def _sha256(path: Path) -> str:18    digest = hashlib.sha256()19    with path.open("rb") as handle:20        for chunk in iter(lambda: handle.read(1024 * 1024), b""):21            digest.update(chunk)22    return digest.hexdigest()23 24 25def _scan(text: str) -> list[tuple[str, list[str] | None]]:26    body = "\n".join(line.split("#", 1)[0] for line in text.splitlines())27    out: list[tuple[str, list[str] | None]] = []28    i, n = 0, len(body)29    while i < n:30        if body[i].isspace():31            i += 132            continue33        start = i34        while i < n and not body[i].isspace() and body[i] != "(":35            i += 136        name = body[start:i]37        args = None38        if i < n and body[i] == "(":39            close = body.find(")", i)40            if close < 0:41                raise ValueError(f"unclosed token arguments after {name!r}")42            inner = body[i + 1 : close].strip()43            args = [part.strip() for part in inner.split(",")] if inner else []44            i = close + 145        if name:46            out.append((name, args))47    return out48 49 50def _label(name: str, args: list[str] | None, *, compact: bool = False) -> str:51    if not args:52        return name53    separator = "," if compact else ", "54    return f"{name}({separator.join(args)})"55 56 57def _fmt_decimal(value: float, places: int) -> str:58    if places <= 0:59        return str(int(round(value)))60    text = f"{value:.{places}f}".rstrip("0").rstrip(".")61    return text or "0"62 63 64def _decode_float(lo: float, hi: float, encoded: float) -> float:65    return lo + (float(encoded) + 1.0) * 0.5 * (hi - lo)66 67 68def _encode_float(lo: float, hi: float, value: float) -> float:69    return (2.0 / (hi - lo)) * float(value) - (hi + lo) / (hi - lo)70 71 72def _level(value: float, steps: int) -> int:73    return round((float(value) + 1.0) * 0.5 * steps)74 75 76class DoomTextFormatter:77    def __init__(self, vocab: dict, tables: dict):78        self.vocab_blob = vocab79        self.tables = tables80        self.words = list(vocab["words"])81        self.labels = list(vocab["labels"])82        if len(self.words) != int(vocab["n_rows"]) or len(self.labels) != len(83            self.words84        ):85            raise ValueError("frozen Doom vocabulary arrays have inconsistent widths")86        self.word_to_id = {word: row for row, word in enumerate(self.words)}87        self.label_to_id = {label: row for row, label in enumerate(self.labels)}88        if len(self.word_to_id) != len(self.words) or len(self.label_to_id) != len(89            self.labels90        ):91            raise ValueError("frozen Doom vocabulary is not injective")92 93        carrier = tables["carrier"]94        self.value_start = int(carrier["value"]["start"])95        self.value_size = int(carrier["value"]["size"])96        self.angle_start = int(carrier["angle"]["start"])97        self.angle_size = int(carrier["angle"]["size"])98        self.angle_lo = int(carrier["angle"]["lo"])99        self.value_steps = int(tables["value_steps"])100        self.angle_bam = int(tables["angle_bam"])101        # Sentinel encoding for "no back sector": one-sided walls have no102        # back-sector heights, so the prompt carries this reserved value,103        # rendered as "none".104        self.sentinel_value = float(tables["back_height_sentinel"])105        self.marker_range = {106            key: (float(value[0]), float(value[1]))107            for key, value in tables["marker_range"].items()108        }109        self.angle_markers = set(tables["angle_markers"])110        self.sentinel_markers = set(tables["sentinel_markers"])111        self.x_markers = set(tables["x_coord_markers"])112        self.y_markers = set(tables["y_coord_markers"])113        origin = tables.get("origin", [0.0, 0.0])114        self.origin = (float(origin[0]), float(origin[1]))115        self.header_levels = {116            str(key): int(value)117            for key, value in tables.get("header_levels", {}).items()118        }119        layout = tables.get("layout", {})120        self.indent_unit = int(layout.get("indent_unit", 2))121        self.field_indent = int(layout.get("field_indent", 4))122 123    @classmethod124    def from_bundle(125        cls, bundle_dir: str | Path, *, allow_incomplete: bool = False126    ) -> "DoomTextFormatter":127        directory = Path(bundle_dir)128        manifest_path = directory / "doom_bundle_manifest.json"129        manifest = json.loads(manifest_path.read_text(encoding="utf-8"))130        if not allow_incomplete and not manifest.get("validation", {}).get("complete"):131            raise ValueError("Doom bundle manifest is not complete")132        files = manifest.get("files", {})133        for name in ("doom_vocab.json", "doom_tables.json"):134            path = directory / name135            if not path.is_file():136                raise FileNotFoundError(f"Doom bundle is missing {name}")137            expected = files.get(name, {}).get("sha256")138            if expected and _sha256(path) != expected:139                raise ValueError(f"Doom bundle hash mismatch for {name}")140        vocab = json.loads((directory / "doom_vocab.json").read_text(encoding="utf-8"))141        tables = json.loads(142            (directory / "doom_tables.json").read_text(encoding="utf-8")143        )144        if int(vocab["n_rows"]) != int(manifest["vocab_size"]):145            raise ValueError("Doom formatter vocabulary width disagrees with manifest")146        if vocab.get("fingerprint") != manifest.get("row_vocab_fingerprint"):147            raise ValueError("Doom formatter row-vocabulary fingerprint mismatch")148        screen = vocab.get("screen", {})149        manifest_screen = manifest.get("screen", {})150        if (screen.get("width"), screen.get("height")) != (151            manifest_screen.get("width"),152            manifest_screen.get("height"),153        ):154            raise ValueError("Doom formatter screen identity mismatch")155        words_digest = hashlib.sha256(156            json.dumps(157                vocab["words"], ensure_ascii=False, separators=(",", ":")158            ).encode("utf-8")159        ).hexdigest()160        if words_digest != manifest.get("tokenizer_vocab_sha256"):161            raise ValueError("Doom formatter tokenizer-word identity mismatch")162        return cls(vocab, tables)163 164    def rows_from_raw_text(self, raw_text: str) -> list[int]:165        rows = []166        for word in raw_text.split():167            try:168                rows.append(self.word_to_id[word])169            except KeyError:170                raise ValueError(f"unknown canonical Doom word: {word!r}") from None171        return rows172 173    def raw_text_from_rows(self, rows: list[int]) -> str:174        # Same explicit non-negative contract as tokenizer/codec.py (the175        # project-side codec): reject negative rows rather than inheriting176        # Python list wraparound. Parity tests pin the two implementations.177        out = []178        for row in rows:179            try:180                index = int(row)181            except (TypeError, ValueError):182                raise ValueError("Doom row outside frozen vocabulary") from None183            if not 0 <= index < len(self.words):184                raise ValueError("Doom row outside frozen vocabulary")185            out.append(self.words[index])186        return " ".join(out)187 188    def _carrier_kind(self, row: int) -> str | None:189        if self.value_start <= row < self.value_start + self.value_size:190            return "value"191        if self.angle_start <= row < self.angle_start + self.angle_size:192            return "angle"193        return None194 195    def _origin_shift(self, marker: str) -> float:196        if marker in self.x_markers:197            return self.origin[0]198        if marker in self.y_markers:199            return self.origin[1]200        return 0.0201 202    def _shortest_value(203        self, lo: float, hi: float, carrier: float, shift: float204    ) -> str:205        target = _level(carrier, self.value_steps)206        physical = _decode_float(lo, hi, carrier) + shift207        for places in range(10):208            candidate = round(physical, places)209            if (210                _level(_encode_float(lo, hi, candidate - shift), self.value_steps)211                == target212            ):213                return _fmt_decimal(candidate, places)214        return repr(physical)215 216    def _render_carrier(self, marker: str, row: int) -> str:217        if self._carrier_kind(row) == "value":218            if marker not in self.marker_range:219                raise ValueError(f"value follows non-marker {marker!r}")220            lo, hi = self.marker_range[marker]221            carrier = -1.0 + (row - self.value_start) / self.value_steps * 2.0222            shift = self._origin_shift(marker)223            physical = _decode_float(lo, hi, carrier) + shift224            if (225                marker in self.sentinel_markers226                and abs(physical - self.sentinel_value) < 0.5227            ):228                return "none"229            return self._shortest_value(lo, hi, carrier, shift)230        if marker not in self.angle_markers:231            raise ValueError(f"angle carrier follows non-angle marker {marker!r}")232        bam = row - self.angle_start + self.angle_lo233        physical = bam * 360.0 / self.angle_bam234        for places in range(10):235            candidate = round(physical, places)236            if round(candidate * self.angle_bam / 360.0) == bam:237                return _fmt_decimal(candidate, places)238        return repr(physical)239 240    def _encode_carrier(self, marker: str, value: str) -> int:241        if marker in self.marker_range:242            lo, hi = self.marker_range[marker]243            if marker in self.sentinel_markers and value == "none":244                carrier = _encode_float(lo, hi, self.sentinel_value)245            else:246                carrier = _encode_float(247                    lo, hi, float(value) - self._origin_shift(marker)248                )249            return self.value_start + _level(carrier, self.value_steps)250        if marker in self.angle_markers:251            bam = round(float(value) * self.angle_bam / 360.0)252            return self.angle_start + bam - self.angle_lo253        raise ValueError(f"token {marker!r} cannot carry value {value!r}")254 255    def _pretty_flat(self, rows: list[int]) -> str:256        units = []257        i = 0258        while i < len(rows):259            row = rows[i]260            if self._carrier_kind(row):261                raise ValueError(f"carrier at row-stream position {i} has no marker")262            label = self.labels[row]263            if i + 1 < len(rows) and self._carrier_kind(rows[i + 1]):264                name, args = _scan(label)[0]265                args = list(args or ())266                args.append(self._render_carrier(name, rows[i + 1]))267                label = _label(name, args)268                i += 1269            units.append(label)270            i += 1271        return " ".join(units)272 273    def _layout(self, flat: str) -> str:274        lines: list[str] = []275        group: list[str] = []276        level = 0277 278        def flush() -> None:279            if not group:280                return281            if group[0].split("(", 1)[0] in self.header_levels:282                lines.append(" " * (level * self.indent_unit) + group[0])283                if len(group) > 1:284                    lines.append(285                        " " * (level * self.indent_unit + self.field_indent)286                        + " ".join(group[1:])287                    )288            else:289                lines.append(" ".join(group))290            group.clear()291 292        for name, args in _scan(flat):293            if name in self.header_levels:294                flush()295                level = self.header_levels[name]296            group.append(_label(name, args))297        flush()298        return "\n".join(lines)299 300    def format_text(self, raw_tokenizer_text: str) -> str:301        return self._layout(302            self._pretty_flat(self.rows_from_raw_text(raw_tokenizer_text))303        )304 305    def parse_pretty_text(self, pretty_text: str) -> str:306        rows: list[int] = []307        for name, args in _scan(pretty_text):308            pretty = _label(name, args)309            row = self.label_to_id.get(pretty)310            if row is not None:311                rows.append(row)312                continue313            if not args:314                raise ValueError(f"unknown pretty Doom token: {pretty!r}")315            base = _label(name, args[:-1])316            try:317                rows.append(self.label_to_id[base])318            except KeyError:319                raise ValueError(f"unknown pretty Doom token: {pretty!r}") from None320            rows.append(self._encode_carrier(name, args[-1]))321        return self.raw_text_from_rows(rows)322 323 324def main(argv: list[str] | None = None) -> int:325    parser = argparse.ArgumentParser(description="Format canonical Doom tokenizer text")326    parser.add_argument("--bundle", type=Path)327    parser.add_argument("--input", type=Path)328    parser.add_argument("--output", type=Path)329    args = parser.parse_args(argv)330    bundle = args.bundle or Path(__file__).resolve().parent.parent331    formatter = DoomTextFormatter.from_bundle(bundle)332    raw = args.input.read_text(encoding="utf-8") if args.input else sys.stdin.read()333    rendered = formatter.format_text(raw) + "\n"334    if args.output:335        args.output.write_text(rendered, encoding="utf-8")336    else:337        sys.stdout.write(rendered)338    return 0339 340 341if __name__ == "__main__":342    raise SystemExit(main())343