"""Standard-library helpers for the Jeeves CPU conversion notebook.""" from __future__ import annotations import ast import hashlib import json import math import os import struct from pathlib import Path FORMAT_VERSION = "jeeves-mlx-cpu-v1" def write_json(path, value): path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) temp = path.with_name(path.name + ".tmp") temp.write_text(json.dumps(value, indent=2, ensure_ascii=False, allow_nan=False) + "\n", encoding="utf-8") os.replace(temp, path) def read_json(path): return json.loads(Path(path).read_text(encoding="utf-8")) def sha256(path, block_size=8 * 1024 * 1024): h = hashlib.sha256() with Path(path).open("rb") as stream: for chunk in iter(lambda: stream.read(block_size), b""): h.update(chunk) return h.hexdigest() def fingerprint(value): return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest() def adapt_config(original): """Adapt metadata only. Leave all tensor transforms to MLX-LM.""" cfg = json.loads(json.dumps(original)) if "text_config" in cfg: raise ValueError("This notebook targets the flat Jeeves export, not a new nested export.") required = ("layer_types", "num_hidden_layers", "rope_theta", "partial_rotary_factor", "hidden_size") for key in required: if key not in cfg: raise ValueError(f"Missing source configuration field: {key}") expected = ["full_attention" if (i + 1) % 4 == 0 else "linear_attention" for i in range(cfg["num_hidden_layers"])] if cfg["layer_types"] != expected: raise ValueError("Layer schedule changed. Do not guess full_attention_interval.") if cfg.get("hidden_act", "silu") != "silu" or cfg.get("num_experts", 0): raise ValueError("Unsupported architecture change in this source revision.") if cfg.get("quantization") or cfg.get("quantization_config"): raise ValueError("Expected unquantized, fused Jeeves weights.") cfg["model_type"] = "qwen3_5" cfg["full_attention_interval"] = 4 cfg["rope_parameters"] = { "type": "default", "rope_theta": cfg["rope_theta"], "partial_rotary_factor": cfg["partial_rotary_factor"], } cfg["torch_dtype"] = cfg.get("dtype", "bfloat16") return cfg def safetensors_header(path): """Validate container length/offsets without loading multi-GB tensors.""" path = Path(path) size = path.stat().st_size with path.open("rb") as stream: raw = stream.read(8) if len(raw) != 8: raise ValueError(f"Truncated Safetensors file: {path}") length = struct.unpack(" min(64 * 1024 * 1024, size - 8): raise ValueError(f"Invalid Safetensors header length: {path}") header = json.loads(stream.read(length)) data_bytes = size - 8 - length tensors = {k: v for k, v in header.items() if k != "__metadata__"} widths = {"BOOL": 1, "U8": 1, "I8": 1, "F8_E4M3": 1, "F8_E5M2": 1, "I16": 2, "U16": 2, "F16": 2, "BF16": 2, "I32": 4, "U32": 4, "F32": 4, "F64": 8, "I64": 8, "U64": 8} spans = [] for name, t in tensors.items(): start, end = t["data_offsets"] shape = t["shape"] if any(not isinstance(x, int) or x < 0 for x in shape): raise ValueError(f"Invalid tensor shape: {name}") if not (0 <= start <= end <= data_bytes): raise ValueError(f"Invalid tensor offsets: {path.name}:{name}") if t["dtype"] in widths and end - start != math.prod(shape) * widths[t["dtype"]]: raise ValueError(f"Tensor byte count mismatch: {name}") spans.append((start, end)) cursor = 0 for start, end in sorted(spans): if start != cursor: raise ValueError(f"Tensor gaps/overlaps in {path}") cursor = end if cursor != data_bytes: raise ValueError(f"Unexpected trailing data in {path}") return tensors def inspect_model(directory): directory = Path(directory) index = read_json(directory / "model.safetensors.index.json") mapping = index["weight_map"] found, nbytes = {}, 0 for filename in sorted(set(mapping.values())): if Path(filename).name != filename: raise ValueError("Unexpected nested shard path") path = directory / filename header = safetensors_header(path) for key, info in header.items(): if key in found or mapping.get(key) != filename: raise ValueError(f"Duplicate or incorrectly indexed tensor: {key}") found[key] = info nbytes += info["data_offsets"][1] - info["data_offsets"][0] if set(mapping) != set(found): raise ValueError("Shard/index tensor inventory mismatch") advertised = index.get("metadata", {}).get("total_size") if advertised is not None and advertised != nbytes: raise ValueError(f"Tensor byte total mismatch: {advertised} != {nbytes}") return {"tensor_count": len(found), "tensor_bytes": nbytes, "tensors": found} def extract_encoder(source_text): """Retain upstream encoder AST nodes; remove only training/PyTorch imports.""" parsed = ast.parse(source_text) wanted = {"sanitize", "Example", "Markers", "Encoder", "readout_positions"} assignments = {"STATE", "Q", "OPT", "OPT_END", "DECIDE", "THINK", "THINK_END", "CONTROL_RE"} nodes, seen = [], set() for node in parsed.body: if isinstance(node, (ast.ClassDef, ast.FunctionDef)) and node.name in wanted: nodes.append(node) seen.add(node.name) elif isinstance(node, ast.Assign): names = {n.id for target in node.targets for n in ast.walk(target) if isinstance(n, ast.Name)} if names and names <= assignments: nodes.append(node) seen.update(names) if seen != wanted | assignments: raise ValueError(f"Upstream encoder structure changed: missing {(wanted | assignments) - seen}") imports = ast.parse("from __future__ import annotations\nimport re\nfrom dataclasses import dataclass\n" "from transformers import AutoTokenizer\nfrom dataformat import DataFormat, Question\n").body module = ast.Module(body=imports + nodes, type_ignores=[]) ast.fix_missing_locations(module) text = ("# Extracted from PostHog/jeeves loader/dataloader.py (MIT).\n" "# Only training imports/unused nodes were removed; see LICENSE_JEEVES_CODE.\n" + ast.unparse(module) + "\n") compile(text, "jeeves_format.py", "exec") return text def build_manifest(directory, identity): directory = Path(directory) files = {} for path in sorted(directory.rglob("*")): if not path.is_file() or path.name == "PACKAGE_MANIFEST.json" or "__pycache__" in path.parts: continue name = path.relative_to(directory).as_posix() files[name] = {"bytes": path.stat().st_size, "sha256": sha256(path)} value = {"schema": FORMAT_VERSION, "identity": identity, "files": files} write_json(directory / "PACKAGE_MANIFEST.json", value) return value def verify_manifest(directory, identity=None): directory = Path(directory) manifest = read_json(directory / "PACKAGE_MANIFEST.json") if identity is not None and manifest["identity"] != identity: raise ValueError("Package belongs to a different conversion identity.") for name, item in manifest["files"].items(): path = directory / name if not path.is_relative_to(directory) or ".." in Path(name).parts: raise ValueError("Unsafe manifest path") if not path.is_file() or path.stat().st_size != item["bytes"] or sha256(path) != item["sha256"]: raise ValueError(f"Package integrity check failed: {name}") return manifest