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()