physicsrob/torchwright-doom-e1m1
334
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 