Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150 / build_native_checkpoint.py
Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
18.5 kB
#!/usr/bin/env python3
"""Stream original Qwen3.5-9B text and MTP safetensors into a standalone TT-native PTQ.
Run only with exclusive TT device ownership. TTNN host conversion
initializes device metadata. Example inside the pinned container:
python /work/build_native_checkpoint.py --weights /weights/Qwen3.5-9B \
--output /work/native-candidate --profile current-bfp4 \
--overrides /work/candidate.json --importance /work/importance/importance.npz \
--container-image IMAGE_ID --device-ownership-confirmed
Each target tensor is remapped individually; linear matrices are transposed
BEFORE native TILE quantization. Other tensors, including BF16 embeddings and
original FP32 nonlinear parameters, are preserved as individual safetensors.
All original mtp.* tensors bypass remapping and quantization, preserving their
runtime keys, values and source dtype for the existing one-layer MTP runtime.
Importance only measures/ranks precision error: it does not optimize rounding,
implement an imatrix-aware quantizer, or measure output KL. Candidate quality
must be evaluated on held-out text with the separate TT runner/scorer.
"""
import argparse
import datetime
import importlib.util
import json
import platform
import shutil
from pathlib import Path
from native_checkpoint import DTYPES, FORMAT, MANIFEST, MTP_KEYS, digest, local_file, matrix_family, save_json, tensor_hash, tensor_precision, validate_mtp_keys
ROOT = Path(__file__).resolve().parent
ASSETS = ("config.json", "tokenizer.json", "tokenizer_config.json", "vocab.json", "merges.txt",
"chat_template.jinja", "special_tokens_map.json", "added_tokens.json", "generation_config.json",
"tokenizer.model", "preprocessor_config.json", "processor_config.json",
"video_preprocessor_config.json")
def load_remapper(path):
spec = importlib.util.spec_from_file_location("native_export_weight_mapping", path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.remap_qwen36_state_dict
def weight_error(original, restored, importance, module, torch):
"""Bounded row chunks; diagonal input second moments, not output KL."""
moments, count = None, None
if importance is not None and module + ".sumsq" in importance.files:
import numpy as np
sumsq = importance[module + ".sumsq"]
count = int(importance[module + ".count"].item())
if count < 1 or sumsq.shape != (original.shape[1],) or not np.isfinite(sumsq).all() or (sumsq < 0).any():
raise ValueError(f"Invalid activation importance for {module}")
moments = torch.from_numpy(sumsq.copy()).to(torch.float64) / count
error_sum, original_sum, weighted_error, weighted_signal, max_error = 0.0, 0.0, 0.0, 0.0, 0.0
for start in range(0, original.shape[0], 128):
source = original[start:start + 128].to(torch.float32)
error = restored[start:start + 128].to(torch.float32) - source
squared = error.square()
error_sum += squared.sum(dtype=torch.float64).item()
original_sum += source.square().sum(dtype=torch.float64).item()
max_error = max(max_error, error.abs().max().item())
if moments is not None:
weighted_error += (squared.sum(dim=0, dtype=torch.float64) * moments).sum().item()
weighted_signal += (source.square().sum(dim=0, dtype=torch.float64) * moments).sum().item()
result = {"sum_squared_error": error_sum, "mean_squared_error": error_sum / original.numel(),
"max_abs_error": max_error,
"relative_squared_error": error_sum / original_sum if original_sum else None,
"importance_available": moments is not None}
if moments is not None:
result.update({"activation_rows": count, "importance_weighted_squared_error": weighted_error,
"importance_weighted_mean_per_output": weighted_error / original.shape[0],
"importance_weighted_relative_squared_error": weighted_error / weighted_signal if weighted_signal else None})
return result
def export(args):
import numpy as np
import torch
import ttnn
from safetensors import safe_open
from safetensors.torch import save_file
from tt_eval import precision_map
torch.set_num_threads(args.cpu_threads)
weights, output = args.weights.resolve(), args.output.resolve()
if not output.is_relative_to(ROOT) or output == ROOT or output.is_relative_to(weights) or weights.is_relative_to(output):
raise ValueError(f"Output must be isolated beneath {ROOT}, outside original weights")
if output.exists() and any(output.iterdir()):
raise ValueError("Refusing to overwrite a nonempty output directory")
config = json.loads((weights / "config.json").read_text())
text = config["text_config"]
if (text["num_hidden_layers"], text["hidden_size"], text["vocab_size"]) != (32, 4096, 248320):
raise ValueError("Only original Qwen3.5-9B is supported")
overrides = json.loads(args.overrides.read_text()) if args.overrides else None
precision = precision_map(args.profile, overrides, text["layer_types"])
remap_path = ROOT / "runtime/qwen36/tt/weight_mapping.py"
remap = load_remapper(remap_path)
index_path = weights / "model.safetensors.index.json"
index = json.loads(index_path.read_text())["weight_map"]
validate_mtp_keys(index)
shards = {}
for name, filename in index.items():
if name.startswith("mtp.") or name == "lm_head.weight" or (name.startswith("model.language_model.") and ".mtp." not in name and ".visual." not in name):
shards.setdefault(filename, []).append(name)
if not shards:
raise ValueError("Expected original model.language_model text checkpoint keys")
importance = np.load(args.importance, allow_pickle=False) if args.importance else None
output.mkdir(parents=True, exist_ok=True)
(output / "tensors").mkdir()
(output / "provenance").mkdir()
files, tensors, source_shards, errors = {}, {}, {}, {}
def register(path):
path.chmod(0o644)
name = str(path.relative_to(output))
files[name] = {"bytes": path.stat().st_size, "sha256": digest(path)}
return name
for name in ASSETS:
source = weights / name
if source.is_file():
shutil.copyfile(source, output / name)
register(output / name)
licenses = [p for p in weights.iterdir() if p.is_file() and p.name.upper().startswith(("LICENSE", "NOTICE", "COPYING"))]
if not licenses or not all((output / name).is_file() for name in ("config.json", "tokenizer.json", "tokenizer_config.json")):
raise ValueError("Original config/tokenizer/license assets are required for standalone export")
for source in licenses:
shutil.copyfile(source, output / source.name)
register(output / source.name)
for source, destination in ((Path(__file__), output / "provenance/build_native_checkpoint.py"),
(ROOT / "native_checkpoint.py", output / "native_checkpoint.py"),
(remap_path, output / "provenance/weight_mapping.py"),
(ROOT / "tt_eval.py", output / "provenance/tt_eval.py"),
(ROOT / "precision-plan.json", output / "provenance/precision-plan.json"),
(index_path, output / "provenance/original.safetensors.index.json")):
shutil.copyfile(source, destination)
register(destination)
if args.overrides:
shutil.copyfile(args.overrides, output / "provenance/overrides.json")
register(output / "provenance/overrides.json")
if args.importance:
shutil.copyfile(args.importance, output / "provenance/importance.npz")
register(output / "provenance/importance.npz")
if (args.importance.parent / "metadata.json").is_file():
shutil.copyfile(args.importance.parent / "metadata.json", output / "provenance/importance-metadata.json")
register(output / "provenance/importance-metadata.json")
counter, quantized, fp32, mtp_fp32 = 0, 0, 0, 0
for filename, names in sorted(shards.items()):
shard = local_file(weights, filename)
source_shards[filename] = {"bytes": shard.stat().st_size, "sha256": digest(shard)}
with safe_open(str(shard), framework="pt", device="cpu") as source:
for original_name in sorted(names):
original = source.get_tensor(original_name)
original_hash = tensor_hash(original)
is_mtp = original_name.startswith("mtp.")
remapped = {original_name: original} if is_mtp else remap({original_name: original})
for name, value in remapped.items():
if name in tensors or "visual" in name or ("mtp" in name and not is_mtp):
raise ValueError(f"Unexpected duplicate/non-text remapped tensor: {name}")
dtype_name = None if is_mtp else tensor_precision(name, precision)
entry = {"source_name": original_name, "source_shard": filename,
"source_shape": list(original.shape), "source_torch_dtype": str(original.dtype),
"source_tensor_sha256": original_hash, "shape": list(value.shape),
"family": matrix_family(name), "restored_torch_dtype": str(value.dtype)}
if dtype_name is None:
value_hash = original_hash if is_mtp else tensor_hash(value)
path = output / "tensors" / f"{counter:05d}.safetensors"
save_file({name: value.contiguous()}, str(path))
with safe_open(str(path), framework="pt", device="cpu") as saved:
restored = saved.get_tensor(name)
if (restored.dtype != value.dtype or not torch.equal(restored, value)
or (is_mtp and tensor_hash(restored) != original_hash)):
raise ValueError(f"Lossless roundtrip failed: {name}")
del restored
if is_mtp:
mtp_fp32 += int(value.dtype == torch.float32)
else:
fp32 += int(value.dtype == torch.float32)
entry.update({"storage": "safetensors-lossless", "tensor_sha256": value_hash,
"roundtrip": {"exact_values": True, "exact_dtype": True}})
else:
if value.ndim != 2 or value.dtype != torch.bfloat16:
raise ValueError(f"Expected original BF16 linear matrix: {name} {value.shape} {value.dtype}")
path = output / "tensors" / f"{counter:05d}.tensorbin"
oriented = value.T.contiguous()
native = ttnn.from_torch(oriented, dtype=getattr(ttnn, DTYPES[dtype_name]), layout=ttnn.TILE_LAYOUT)
del oriented
ttnn.dump_tensor(str(path), native)
reloaded = ttnn.load_tensor(str(path))
rounded = ttnn.to_torch(reloaded).to(torch.bfloat16)
if not torch.equal(ttnn.to_torch(native), rounded):
raise ValueError(f"Native serialization or BF16 host restoration changes values: {name}")
del native, reloaded
restored = rounded.T.contiguous()
# Deliberately exercise the same transpose/copy that runtime converters do.
second = ttnn.from_torch(restored.T.contiguous(), dtype=getattr(ttnn, DTYPES[dtype_name]), layout=ttnn.TILE_LAYOUT)
second_path = output / "tensors" / f"{counter:05d}.roundtrip.tensorbin"
ttnn.dump_tensor(str(second_path), second)
exact_values = torch.equal(rounded, ttnn.to_torch(second))
exact_bytes = digest(path) == digest(second_path)
if not exact_values or not exact_bytes:
raise ValueError(f"Native requantization is not idempotent: {name}; values={exact_values}, bytes={exact_bytes}")
second_path.unlink()
del second, rounded
errors[name] = {"source_module": original_name.removesuffix(".weight"), "precision": dtype_name,
**weight_error(value, restored, importance, original_name.removesuffix(".weight"), torch)}
del restored
entry.update({"storage": "ttnn-tile", "precision": dtype_name,
"native_dtype": DTYPES[dtype_name], "native_shape": [value.shape[1], value.shape[0]],
"orientation": "input,output", "layout": "TILE_LAYOUT",
"roundtrip": {"exact_values": exact_values, "exact_serialized_bytes": exact_bytes}})
quantized += 1
entry["file"] = register(path)
entry["sha256"] = files[entry["file"]]["sha256"]
tensors[name] = entry
counter += 1
print(json.dumps({"tensor": name, "precision": dtype_name or "lossless", "bytes": files[entry["file"]]["bytes"]}), flush=True)
del original, remapped, value
if importance is not None:
importance.close()
if not {"tok_embeddings.weight", "output.weight", "norm.weight"} <= tensors.keys() or fp32 != 48:
raise ValueError(f"Missing top-level tensors or original FP32 nonlinear tensors: fp32={fp32}, expected 48")
validate_mtp_keys(tensors)
save_json(output / "quantization-error.json", {"method": "Unmodified TTNN rounding; importance-weighted precision ranking only, not optimized values or output KL",
"formula": "sum_out,in ((W-Wq)^2 * input_sumsq[in]/input_count)",
"tensors": errors})
register(output / "quantization-error.json")
manifest = {"format": FORMAT, "schema_version": 1, "scope": "text-only-no-vision-with-mtp",
"mtp": {"enabled": True, "num_speculative_tokens": 1,
"source_tensor_count": len(MTP_KEYS), "storage": "lossless"},
"status": "tensor_roundtrip_verified", "precision": precision,
"end_to_end_verification": "Requires separate manifest-bound equivalence.json (MTP=0) or equivalence-mtp.json (MTP=1); serialization alone is not validation",
"profile": args.profile, "files": files, "tensors": tensors,
"roundtrip": {"all_passed": True, "native_matrices": quantized, "lossless_tensors": counter - quantized,
"preserved_original_fp32_tensors": fp32,
"preserved_original_mtp_fp32_tensors": mtp_fp32},
"source": {"original_index_sha256": digest(index_path), "shards": source_shards,
"declared_revision": "c202236235762e1c871ad0ccb60c8ee5ba337b9a"},
"toolchain": {"python": platform.python_version(), "torch": torch.__version__,
"ttnn": getattr(ttnn, "__version__", None), "container_image": args.container_image,
"tt_metal_commit": "de59f8a658b1ceafd230c8266026b1a72bb198d7"},
"load_contract": {"weights": "native dump -> CPU TT -> BF16 torch [out,in] -> runtime transpose -> requantize",
"nonlinear": "Lossless original tensors, preserving FP32 until normal runtime conversion",
"mtp": "Lossless original mtp.* tensors; existing Qwen36MTP uses target layer-0 precision policy with one speculative token",
"gdn_derived": "Runtime rebuilds AB and QKVABZ from quant-rounded components, then requantizes; no derived tensor stored",
"supported_runtime": "single-device Qwen36 text with optional MTP-1 at manifest precision, generic or explicitly dtype-gated packed families, only under verified runtime sources/environment",
"risks": ["Block exponent grouping is orientation/layout dependent", "Component idempotence does not prove GDN derived or full-model parity", "Every generic/packed configuration requires its own equivalence evidence; TP is excluded"],
"required_evidence": "native_checkpoint.py CHECKPOINT --baseline ORIGINAL_AT_SAME_PRECISION --restored NATIVE_RELOAD_RUN writes equivalence.json for recorded MTP=0 or equivalence-mtp.json for recorded MTP=1 only after exact full-logit comparison; actual speculation requires separate live-cycle evidence"},
"created_at": datetime.datetime.now(datetime.timezone.utc).isoformat()}
save_json(output / MANIFEST, manifest)
print(json.dumps({"manifest": str(output / MANIFEST), "tensors": counter, "native_matrices": quantized,
"artifact_bytes": sum(info["bytes"] for info in files.values()), "end_to_end_verified": False}), flush=True)
def main():
from tt_eval import PROFILES
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--weights", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--profile", choices=tuple(PROFILES), default="current-bfp4")
parser.add_argument("--overrides", type=Path)
parser.add_argument("--importance", type=Path)
parser.add_argument("--cpu-threads", type=int, default=8)
parser.add_argument("--container-image", required=True, help="Exact image identity reported by docker image inspect")
parser.add_argument("--device-ownership-confirmed", action="store_true")
args = parser.parse_args()
if not args.device_ownership_confirmed:
parser.error("TTNN host conversion requires exclusive device ownership; stop other P150 workloads first")
if args.cpu_threads < 1:
parser.error("--cpu-threads must be positive")
export(args)
if __name__ == "__main__":
main()