#!/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()