Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
16 kB
#!/usr/bin/env python3
"""TT-native text-only checkpoint contract and lossless host reconstruction.
Native matrices are TTNN TILE dumps in [input, output] orientation. Loading
returns remapped torch [output, input] BF16 tensors; the runtime requantizes
them. This is NOT direct device-cache loading, GGUF, QAT, or an official
Unsloth quant. Exact per-matrix idempotence is necessary but insufficient:
GDN AB/mega tensors concatenate quant-rounded components, and packed/fused
layouts can regroup block exponents. End-to-end equivalence is required for
each supported runtime configuration. Single-device generic or dtype-gated
mixed packed execution with optional one-token MTP can be verified; vision
and tensor parallelism are excluded. MTP source tensors are stored losslessly.
TTNN host conversion may initialize driver metadata: even host-only commands
need exclusive device ownership in this environment. No original HF checkpoint
is used by this loader. The installed compatible TT runtime is still required.
"""
import argparse
import hashlib
import json
import os
from pathlib import Path
MANIFEST = "native_manifest.json"
FORMAT = "qwen35-9b-ttnn-native-text-v1"
DTYPES = {"bf16": "bfloat16", "bfp8": "bfloat8_b", "bfp4": "bfloat4_b"}
MTP_KEYS = frozenset({
"mtp.fc.weight",
"mtp.norm.weight",
"mtp.pre_fc_norm_embedding.weight",
"mtp.pre_fc_norm_hidden.weight",
"mtp.layers.0.input_layernorm.weight",
"mtp.layers.0.post_attention_layernorm.weight",
"mtp.layers.0.self_attn.q_norm.weight",
"mtp.layers.0.self_attn.k_norm.weight",
"mtp.layers.0.self_attn.q_proj.weight",
"mtp.layers.0.self_attn.k_proj.weight",
"mtp.layers.0.self_attn.v_proj.weight",
"mtp.layers.0.self_attn.o_proj.weight",
"mtp.layers.0.mlp.gate_proj.weight",
"mtp.layers.0.mlp.up_proj.weight",
"mtp.layers.0.mlp.down_proj.weight",
})
def tensor_hash(tensor):
import torch
raw = tensor.contiguous().view(torch.uint8).reshape(-1).numpy()
return hashlib.sha256(memoryview(raw)).hexdigest()
def validate_mtp_keys(keys):
actual = {name for name in keys if name.startswith("mtp.")}
if actual != MTP_KEYS:
raise ValueError(f"Incomplete/unsupported one-layer MTP subtree: missing={sorted(MTP_KEYS - actual)}, unexpected={sorted(actual - MTP_KEYS)}")
def equivalence_filename(environment):
mode = environment.get("QWEN36_MTP")
if mode not in ("0", "1"):
raise ValueError("Equivalence requires explicit QWEN36_MTP=0 or QWEN36_MTP=1")
return "equivalence-mtp.json" if mode == "1" else "equivalence.json"
def digest(path):
value = hashlib.sha256()
with Path(path).open("rb") as stream:
for block in iter(lambda: stream.read(8 * 1024 * 1024), b""):
value.update(block)
return value.hexdigest()
def save_json(path, value):
path = Path(path)
temporary = path.with_suffix(path.suffix + ".partial")
temporary.write_text(json.dumps(value, sort_keys=True, indent=2, allow_nan=False) + "\n")
temporary.replace(path)
def local_file(root, name):
root = Path(root).resolve()
path = (root / name).resolve()
if not path.is_relative_to(root) or path == root:
raise ValueError(f"Artifact path escapes checkpoint: {name}")
return path
def read_manifest(root):
manifest = json.loads((Path(root) / MANIFEST).read_text())
if manifest.get("format") != FORMAT or manifest.get("scope") != "text-only-no-vision-with-mtp":
raise ValueError("Unsupported native checkpoint format/scope")
if not manifest.get("tensors") or not manifest.get("roundtrip", {}).get("all_passed"):
raise ValueError("Incomplete native checkpoint or failed tensor roundtrip")
mtp = manifest.get("mtp", {})
if (mtp.get("enabled") is not True or type(mtp.get("num_speculative_tokens")) is not int
or mtp["num_speculative_tokens"] != 1 or type(mtp.get("source_tensor_count")) is not int
or mtp["source_tensor_count"] != len(MTP_KEYS) or mtp.get("storage") != "lossless"):
raise ValueError("Checkpoint requires complete lossless MTP-1 metadata")
validate_mtp_keys(manifest["tensors"])
for name in MTP_KEYS:
entry = manifest["tensors"][name]
source_hash = entry.get("source_tensor_sha256")
file_info = manifest.get("files", {}).get(entry.get("file"), {})
if (entry.get("storage") != "safetensors-lossless" or entry.get("source_name") != name
or not entry.get("source_shape") or entry["source_shape"] != entry.get("shape")
or not entry.get("source_torch_dtype")
or entry["source_torch_dtype"] != entry.get("restored_torch_dtype")
or not isinstance(source_hash, str) or len(source_hash) != 64
or any(c not in "0123456789abcdef" for c in source_hash)
or source_hash != entry.get("tensor_sha256")
or not file_info.get("sha256") or entry.get("sha256") != file_info["sha256"]
or entry.get("roundtrip", {}).get("exact_values") is not True
or entry.get("roundtrip", {}).get("exact_dtype") is not True):
raise ValueError(f"MTP tensor is not bound to lossless original values/dtype: {name}")
return manifest
def verify_files(root, manifest):
for name, info in manifest["files"].items():
path = local_file(root, name)
if path.stat().st_size != info["bytes"] or digest(path) != info["sha256"]:
raise ValueError(f"Checkpoint file hash/size mismatch: {name}")
def require_equivalence(root, manifest):
proof = json.loads((Path(root) / equivalence_filename(os.environ)).read_text())
if (proof.get("manifest_sha256") != digest(Path(root) / MANIFEST)
or proof.get("exact_logits_equal") is not True or proof.get("tokens_compared", 0) < 1):
raise ValueError("Native checkpoint lacks matching exact end-to-end equivalence evidence")
if proof.get("precision") != manifest["precision"]:
raise ValueError("Equivalence evidence precision mismatch")
environment = proof.get("runtime_environment", {})
if equivalence_filename(environment) != equivalence_filename(os.environ):
raise ValueError("Equivalence evidence MTP mode mismatch")
return proof
def validate_runtime(root, runtime_root, num_devices):
"""Reject using equivalence evidence under a different execution contract."""
if num_devices != 1:
raise ValueError("Native checkpoint equivalence covers one device only")
proof = require_equivalence(root, read_manifest(root))
for name, value in proof["runtime_environment"].items():
if os.environ.get(name) != value:
raise ValueError(f"Runtime environment differs from native equivalence: {name}")
for name, value in proof["runtime_sources"].items():
if digest(local_file(runtime_root, name)) != value:
raise ValueError(f"Runtime source differs from native equivalence: {name}")
def matrix_family(name):
if name == "output.weight":
return "lm_head"
if not name.startswith("layers.") or not name.endswith(".weight"):
return None
suffix = ".".join(name.split(".")[2:])
if suffix in {f"self_attn.{p}_proj.weight" for p in ("q", "k", "v", "o")}:
return "attention"
if suffix == "linear_attn.out_proj.weight":
return "gdn_output"
if suffix in {f"linear_attn.{p}.weight" for p in ("qkv_proj", "in_proj_a", "in_proj_b", "in_proj_z")}:
return "gdn"
if suffix in {"mlp.gate_proj.weight", "mlp.up_proj.weight"}:
return "mlp_gate_up"
if suffix == "mlp.down_proj.weight":
return "mlp_down"
return None
def tensor_precision(name, plan):
family = matrix_family(name)
if family is None:
return None
if family == "lm_head":
return plan[family]
values = plan["layers"][name.split(".")[1]]
return values.get("gdn_output", values["gdn"]) if family == "gdn_output" else values[family]
def load_native_checkpoint(root, verify_hashes=True, allow_unverified=False):
"""Reconstruct text and MTP tensors; allow_unverified is for equivalence runs only.
Small/nonlinear tensors (including original FP32 A_log/dt_bias) retain
original values/dtypes. Embeddings and all MTP tensors are stored losslessly. This loader
materializes one remapped model, never an HF model or an original copy.
"""
import torch
import ttnn
from safetensors import safe_open
root = Path(root).resolve()
manifest = read_manifest(root)
if verify_hashes:
verify_files(root, manifest)
if not allow_unverified:
require_equivalence(root, manifest)
result = {}
for name, entry in manifest["tensors"].items():
path = local_file(root, entry["file"])
if entry["storage"] == "ttnn-tile":
if entry["precision"] != tensor_precision(name, manifest["precision"]):
raise ValueError(f"Tensor precision disagrees with checkpoint plan: {name}")
tensor = ttnn.load_tensor(str(path))
if list(tensor.shape) != entry["native_shape"] or tensor.dtype != getattr(ttnn, DTYPES[entry["precision"]]):
raise ValueError(f"Native tensor shape/dtype mismatch: {name}")
if tensor.layout != ttnn.TILE_LAYOUT:
raise ValueError(f"Native tensor layout mismatch: {name}")
value = ttnn.to_torch(tensor).to(torch.bfloat16).T.contiguous()
del tensor
elif entry["storage"] == "safetensors-lossless":
with safe_open(str(path), framework="pt", device="cpu") as source:
value = source.get_tensor(name)
else:
raise ValueError(f"Unsupported tensor storage: {name}")
if list(value.shape) != entry["shape"] or str(value.dtype) != entry["restored_torch_dtype"]:
raise ValueError(f"Restored tensor shape/dtype mismatch: {name}")
if name in MTP_KEYS and tensor_hash(value) != entry["source_tensor_sha256"]:
raise ValueError(f"Restored MTP tensor differs from original source hash: {name}")
result[name] = value
if not {"tok_embeddings.weight", "output.weight", "norm.weight"} <= result.keys():
raise ValueError("Native checkpoint is missing required top-level tensors")
return result
def record_equivalence(root, baseline, restored):
"""Compare actual full-vocabulary runner artifacts, not weight-error proxies.
Baseline must evaluate ORIGINAL weights at the SAME chosen precision and
runtime, not baseline-bf16 against a mixed candidate. Exact logits equality
is deliberately strict. Evidence covers the supplied records only.
The destination is selected by the recorded MTP runtime environment, never
by the environment of this proof-writing process. MTP proof does not replace
a separate live speculative-cycle verification.
"""
import numpy as np
root, baseline, restored = map(Path, (root, baseline, restored))
manifest = read_manifest(root)
verify_files(root, manifest)
left = json.loads((baseline / "metadata.json").read_text())
right = json.loads((restored / "metadata.json").read_text())
for metadata in (left, right):
if metadata.get("status") != "complete" or metadata.get("backend") != "ttnn-native":
raise ValueError("Both evaluations must be completed TTNN runs")
if metadata.get("precision") != manifest["precision"] or not metadata.get("full_vocabulary"):
raise ValueError("Both evaluations must use the exported precision and full logits")
for key in ("precision", "selection", "input_sha256", "runtime_environment", "alignment", "max_seq_len"):
if left.get(key) != right.get(key):
raise ValueError(f"End-to-end comparison configuration mismatch: {key}")
proof_filename = equivalence_filename(left.get("runtime_environment", {}))
if left["source_identity"]["sources"] != right["source_identity"]["sources"]:
raise ValueError("Runtime source hashes differ between equivalence runs")
if right["source_identity"].get("native_manifest_sha256") != digest(root / MANIFEST):
raise ValueError("Restored evaluation is not bound to this native manifest")
if left["source_identity"].get("native_manifest_sha256") is not None:
raise ValueError("Baseline must use the original HF checkpoint")
if left["source_identity"].get("index_sha256") != manifest["source"]["original_index_sha256"]:
raise ValueError("Baseline original checkpoint index differs from exported source")
for metadata in (left, right):
if metadata["source_identity"].get("config_sha256") != manifest["files"]["config.json"]["sha256"]:
raise ValueError("Evaluation config differs from native checkpoint")
rows_a = [json.loads(line) for line in (baseline / "records.jsonl").read_text().splitlines() if line.strip()]
rows_b = [json.loads(line) for line in (restored / "records.jsonl").read_text().splitlines() if line.strip()]
if not rows_a or len(rows_a) != len(rows_b):
raise ValueError("Evaluation record coverage differs or is empty")
for metadata, rows in ((left, rows_a), (right, rows_b)):
if metadata.get("completed_records") != len(rows) or [row["id"] for row in rows] != metadata["selection"]["record_ids"]:
raise ValueError("Evaluation records do not cover the declared completed selection")
evidence, tokens = [], 0
for a, b in zip(rows_a, rows_b):
for key in ("id", "split", "token_ids"):
if a[key] != b[key]:
raise ValueError(f"Evaluation record mismatch: {key}")
file_a = local_file(baseline, a["logits_file"])
file_b = local_file(restored, b["logits_file"])
x, y = np.load(file_a, mmap_mode="r"), np.load(file_b, mmap_mode="r")
expected = (len(a["token_ids"]) - 1, left["vocab_size"])
if x.shape != expected or y.shape != expected:
raise ValueError(f"Unexpected full-logits shape for record {a['id']}")
for start in range(0, x.shape[0], 16):
if not np.isfinite(x[start:start + 16]).all() or not np.array_equal(x[start:start + 16], y[start:start + 16]):
raise ValueError(f"Native reload logits differ for record {a['id']} at chunk {start}")
tokens += x.shape[0]
evidence.append({"id": a["id"], "baseline_sha256": digest(file_a), "restored_sha256": digest(file_b)})
proof = {"manifest_sha256": digest(root / MANIFEST), "exact_logits_equal": True,
"tokens_compared": tokens, "records": evidence, "precision": manifest["precision"],
"runtime_sources": left["source_identity"]["sources"],
"runtime_environment": left["runtime_environment"],
"baseline_metadata_sha256": digest(baseline / "metadata.json"),
"restored_metadata_sha256": digest(restored / "metadata.json"),
"scope": "Exact parity only on recorded full-vocabulary teacher-forced sequences; not a quality certification"}
save_json(root / proof_filename, proof)
return proof
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("checkpoint", type=Path)
parser.add_argument("--baseline", type=Path)
parser.add_argument("--restored", type=Path)
args = parser.parse_args()
if bool(args.baseline) != bool(args.restored):
parser.error("--baseline and --restored must be supplied together")
if args.baseline:
print(json.dumps(record_equivalence(args.checkpoint, args.baseline, args.restored), indent=2))
else:
manifest = read_manifest(args.checkpoint)
verify_files(args.checkpoint, manifest)
print(json.dumps({"files_verified": True, "tensors": len(manifest["tensors"]),
"equivalence": require_equivalence(args.checkpoint, manifest)}, indent=2))
if __name__ == "__main__":
main()