jeeves-mlx / jeeves_support.py
cowWhySo's picture
Publish Jeeves MLX 4bit (CPU diagnostics; not benchmarked)
6f4220a verified
Raw History Blame Contribute Delete
7.92 kB
"""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("<Q", raw)[0]
if length < 2 or length > 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