File size: 18,517 Bytes
12f320c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 | #!/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()
|