ajh-code commited on
Commit
bc9418f
·
verified ·
1 Parent(s): efecef1

Add runtime/packed_artifact.py

Browse files
Files changed (1) hide show
  1. runtime/packed_artifact.py +1068 -0
runtime/packed_artifact.py ADDED
@@ -0,0 +1,1068 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import hashlib
7
+ import json
8
+ import math
9
+ import os
10
+ import sys
11
+ from collections import OrderedDict
12
+ from dataclasses import asdict, dataclass
13
+ from datetime import datetime, timezone
14
+ from pathlib import Path
15
+ from typing import Callable, Iterable
16
+
17
+ import safetensors
18
+ import torch
19
+ import torch.nn as nn
20
+ from safetensors import safe_open
21
+ from safetensors.torch import save_file
22
+
23
+
24
+ FORMAT_VERSION = 1
25
+ ARTIFACT_KIND = "mage_flow_transformer_mlp_nvfp4_resident_v1"
26
+ CANONICAL_SAFETENSORS_METADATA = {
27
+ "mage_nvfp4_contract": (
28
+ f"{ARTIFACT_KIND};format_version={FORMAT_VERSION}"
29
+ )
30
+ }
31
+ LEGACY_SAFETENSORS_METADATA = {
32
+ "artifact_kind": ARTIFACT_KIND,
33
+ "format_version": str(FORMAT_VERSION),
34
+ }
35
+ FP4_BLOCK_ELEMENTS = 16
36
+ SCALE_TILE_OUTER = 128
37
+ SCALE_TILE_INNER = 4
38
+ FP4_E2M1_MAX = 6.0
39
+ FP4_TENSOR_SCALE_MAX = 448.0
40
+ NVFP4_TENSOR_SCALE_DENOMINATOR = FP4_E2M1_MAX * FP4_TENSOR_SCALE_MAX
41
+ TARGET_DEPTH = 12
42
+
43
+ RELEASE_ROOT = Path(__file__).resolve().parents[1]
44
+ PROJECT_ROOT = RELEASE_ROOT
45
+ MAGE_ROOT = RELEASE_ROOT / "vendor"
46
+ RESIDENT_PYTHON_ROOT = RELEASE_ROOT / "runtime"
47
+ RESIDENT_SOURCE = RESIDENT_PYTHON_ROOT / "nvfp4_linear.cu"
48
+ RESIDENT_LIBRARY = RESIDENT_PYTHON_ROOT / "libmage_nvfp4_linear.so"
49
+ ARTIFACT_SCRIPT_PATH = Path(__file__).resolve()
50
+ _NATIVE_PACKER = None
51
+
52
+
53
+ class PackedArtifactError(RuntimeError):
54
+ pass
55
+
56
+
57
+ @dataclass(frozen=True)
58
+ class ScaleLayout:
59
+ inner_dim: int
60
+ outer_tiles: int
61
+ bytes: int
62
+
63
+
64
+ @dataclass(frozen=True)
65
+ class TargetSpec:
66
+ module_key: str
67
+ weight_key: str
68
+ bias_key: str
69
+ artifact_weight_key: str
70
+ artifact_scale_key: str
71
+ artifact_tensor_scale_key: str
72
+ artifact_bias_key: str
73
+
74
+
75
+ def fail(message: str) -> None:
76
+ raise PackedArtifactError(message)
77
+
78
+
79
+ def round_up(value: int, multiple: int) -> int:
80
+ return ((value + multiple - 1) // multiple) * multiple
81
+
82
+
83
+ def sha256_file(path: Path) -> str:
84
+ digest = hashlib.sha256()
85
+ with path.open("rb") as handle:
86
+ while True:
87
+ chunk = handle.read(1 << 20)
88
+ if not chunk:
89
+ break
90
+ digest.update(chunk)
91
+ return digest.hexdigest()
92
+
93
+
94
+ def sha256_bytes(data: bytes) -> str:
95
+ return hashlib.sha256(data).hexdigest()
96
+
97
+
98
+ def tensor_bytes(tensor: torch.Tensor) -> bytes:
99
+ if not tensor.is_contiguous():
100
+ tensor = tensor.contiguous()
101
+ return tensor.view(torch.uint8).cpu().numpy().tobytes()
102
+
103
+
104
+ def sha256_tensor(tensor: torch.Tensor) -> str:
105
+ return sha256_bytes(tensor_bytes(tensor))
106
+
107
+
108
+ def fsync_directory(path: Path) -> None:
109
+ descriptor = os.open(path, os.O_RDONLY | getattr(os, "O_DIRECTORY", 0))
110
+ try:
111
+ os.fsync(descriptor)
112
+ finally:
113
+ os.close(descriptor)
114
+
115
+
116
+ def fsync_file(path: Path) -> None:
117
+ descriptor = os.open(path, os.O_RDONLY)
118
+ try:
119
+ os.fsync(descriptor)
120
+ finally:
121
+ os.close(descriptor)
122
+
123
+
124
+ def write_bytes_once(path: Path, payload: bytes) -> None:
125
+ with path.open("xb") as handle:
126
+ handle.write(payload)
127
+ handle.flush()
128
+ os.fsync(handle.fileno())
129
+ fsync_directory(path.parent)
130
+
131
+
132
+ def host_scale_offset(outer: int, inner_scale: int, scale_inner_dim: int) -> int:
133
+ outer_tile = outer // SCALE_TILE_OUTER
134
+ local_outer = outer % SCALE_TILE_OUTER
135
+ local_inner = inner_scale % SCALE_TILE_INNER
136
+ inner_tile_start = inner_scale - local_inner
137
+ tile_base = (inner_tile_start + outer_tile * scale_inner_dim) * SCALE_TILE_OUTER
138
+ return tile_base + (local_outer % 32) * 16 + (local_outer // 32) * 4 + local_inner
139
+
140
+
141
+ def make_scale_layout(rows_k: int, outer_columns: int) -> ScaleLayout:
142
+ if rows_k <= 0 or outer_columns <= 0:
143
+ fail("scale layout requires positive rows_k and outer_columns")
144
+ if rows_k % FP4_BLOCK_ELEMENTS:
145
+ fail(
146
+ f"scale layout requires K divisible by {FP4_BLOCK_ELEMENTS}; "
147
+ f"got K={rows_k}"
148
+ )
149
+ inner_dim = round_up(rows_k // FP4_BLOCK_ELEMENTS, SCALE_TILE_INNER)
150
+ outer_tiles = (outer_columns + SCALE_TILE_OUTER - 1) // SCALE_TILE_OUTER
151
+ return ScaleLayout(
152
+ inner_dim=inner_dim,
153
+ outer_tiles=outer_tiles,
154
+ bytes=outer_tiles * inner_dim * SCALE_TILE_OUTER,
155
+ )
156
+
157
+
158
+ def build_target_specs(depth: int) -> list[TargetSpec]:
159
+ if depth != TARGET_DEPTH:
160
+ fail(
161
+ f"this artifact format is pinned to exactly {TARGET_DEPTH} transformer blocks; "
162
+ f"config reported depth={depth}"
163
+ )
164
+ specs: list[TargetSpec] = []
165
+ suffixes = (
166
+ "img_mlp.net.0.proj",
167
+ "img_mlp.net.2",
168
+ "txt_mlp.net.0.proj",
169
+ "txt_mlp.net.2",
170
+ )
171
+ for index in range(depth):
172
+ for suffix in suffixes:
173
+ module_key = f"transformer_blocks.{index}.{suffix}"
174
+ specs.append(
175
+ TargetSpec(
176
+ module_key=module_key,
177
+ weight_key=f"{module_key}.weight",
178
+ bias_key=f"{module_key}.bias",
179
+ artifact_weight_key=f"targets.{module_key}.packed_weight_e2m1",
180
+ artifact_scale_key=f"targets.{module_key}.packed_scales_ue4m3",
181
+ artifact_tensor_scale_key=f"targets.{module_key}.weight_tensor_scale",
182
+ artifact_bias_key=f"targets.{module_key}.bias_bf16",
183
+ )
184
+ )
185
+ return specs
186
+
187
+
188
+ def decode_fp4_e2m1(raw: int) -> float:
189
+ sign = -1.0 if (raw & 0x8) else 1.0
190
+ magnitude = raw & 0x7
191
+ table = (
192
+ 0.0,
193
+ 0.5,
194
+ 1.0,
195
+ 1.5,
196
+ 2.0,
197
+ 3.0,
198
+ 4.0,
199
+ 6.0,
200
+ )
201
+ return sign * table[magnitude]
202
+
203
+
204
+ def encode_fp4_e2m1(value: float) -> int:
205
+ candidates = [decode_fp4_e2m1(code) for code in range(16)]
206
+ best_code = 0
207
+ best_error = math.inf
208
+ for code, candidate in enumerate(candidates):
209
+ error = abs(candidate - value)
210
+ if error < best_error or (error == best_error and (code & 1) == 0 and (best_code & 1) == 1):
211
+ best_error = error
212
+ best_code = code
213
+ return best_code
214
+
215
+
216
+ def decode_fp8_e4m3(raw: int) -> float:
217
+ sign = -1.0 if (raw & 0x80) else 1.0
218
+ exponent = (raw >> 3) & 0x0F
219
+ mantissa = raw & 0x07
220
+ if exponent == 0:
221
+ if mantissa == 0:
222
+ return 0.0 * sign
223
+ return sign * (mantissa / 8.0) * (2.0 ** -6)
224
+ if exponent == 0x0F and mantissa == 0x07:
225
+ return math.nan
226
+ return sign * (1.0 + mantissa / 8.0) * (2.0 ** (exponent - 7))
227
+
228
+
229
+ def _build_positive_e4m3_table() -> list[tuple[int, float]]:
230
+ table: list[tuple[int, float]] = []
231
+ for raw in range(0x80):
232
+ value = decode_fp8_e4m3(raw)
233
+ if math.isnan(value) or value < 0.0:
234
+ continue
235
+ table.append((raw, value))
236
+ table.sort(key=lambda item: (item[1], item[0]))
237
+ return table
238
+
239
+
240
+ POSITIVE_E4M3_TABLE = _build_positive_e4m3_table()
241
+
242
+
243
+ def encode_fp8_e4m3_satfinite(value: float) -> int:
244
+ if value <= 0.0:
245
+ return 0
246
+ finite_values = [item for item in POSITIVE_E4M3_TABLE if item[1] <= FP4_TENSOR_SCALE_MAX]
247
+ best_raw = finite_values[-1][0]
248
+ best_error = math.inf
249
+ for raw, candidate in finite_values:
250
+ error = abs(candidate - value)
251
+ if error < best_error or (error == best_error and (raw & 1) == 0 and (best_raw & 1) == 1):
252
+ best_error = error
253
+ best_raw = raw
254
+ return best_raw
255
+
256
+
257
+ def host_tensor_scale_from_amax(amax: float) -> float:
258
+ return 1.0 if amax == 0.0 else amax / NVFP4_TENSOR_SCALE_DENOMINATOR
259
+
260
+
261
+ def native_packer():
262
+ global _NATIVE_PACKER
263
+ if _NATIVE_PACKER is None:
264
+ if str(RESIDENT_PYTHON_ROOT) not in sys.path:
265
+ sys.path.insert(0, str(RESIDENT_PYTHON_ROOT))
266
+ from packed_nvfp4_linear import NativeNvfp4Library
267
+
268
+ _NATIVE_PACKER = NativeNvfp4Library(RESIDENT_LIBRARY)
269
+ return _NATIVE_PACKER
270
+
271
+
272
+ def pack_weight_tensor(weight_nk_bf16: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, float]:
273
+ if weight_nk_bf16.dtype != torch.bfloat16 or weight_nk_bf16.device.type != "cpu":
274
+ fail("weight packer expects a CPU bfloat16 tensor")
275
+ if weight_nk_bf16.ndim != 2:
276
+ fail("weight packer expects a 2D [N,K] weight tensor")
277
+ if not weight_nk_bf16.is_contiguous():
278
+ weight_nk_bf16 = weight_nk_bf16.contiguous()
279
+
280
+ columns_n, rows_k = weight_nk_bf16.shape
281
+ if rows_k % 32 != 0 or columns_n % 8 != 0:
282
+ fail(
283
+ f"native NVFP4 packing requires K%32==0 and N%8==0; got N={columns_n} K={rows_k}"
284
+ )
285
+
286
+ amax = float(weight_nk_bf16.float().abs().max().item())
287
+ packed_fp4, packed_scales, tensor_scale = native_packer().pack_weight(
288
+ weight_nk_bf16
289
+ )
290
+ return packed_fp4, packed_scales, tensor_scale.reshape(1), amax
291
+
292
+
293
+ def load_transformer_config(source_repo: Path) -> dict:
294
+ config_path = source_repo / "transformer" / "config.json"
295
+ if not config_path.exists():
296
+ fail(f"missing transformer config: {config_path}")
297
+ return json.loads(config_path.read_text())
298
+
299
+
300
+ def source_transformer_checkpoint(source_repo: Path) -> Path:
301
+ checkpoint = source_repo / "transformer" / "diffusion_pytorch_model.safetensors"
302
+ if not checkpoint.exists():
303
+ fail(f"missing transformer checkpoint: {checkpoint}")
304
+ return checkpoint
305
+
306
+
307
+ def import_mage_transformer_symbols() -> tuple[type[nn.Module], object]:
308
+ if str(MAGE_ROOT) not in sys.path:
309
+ sys.path.insert(0, str(MAGE_ROOT))
310
+ try:
311
+ from mage_flow.models.mage_flow import MageFlow, MageFlowParams
312
+ except Exception as exc:
313
+ fail(f"failed to import local Mage transformer sources from {MAGE_ROOT}: {exc}")
314
+ return MageFlow, MageFlowParams
315
+
316
+
317
+ class PackedNvfp4LinearArtifactModule(nn.Module):
318
+ def __init__(
319
+ self,
320
+ in_features: int,
321
+ out_features: int,
322
+ packed_weight_e2m1: torch.Tensor,
323
+ packed_scales_ue4m3: torch.Tensor,
324
+ weight_tensor_scale: torch.Tensor,
325
+ bias_bf16: torch.Tensor | None,
326
+ ) -> None:
327
+ super().__init__()
328
+ self.in_features = int(in_features)
329
+ self.out_features = int(out_features)
330
+ self.register_buffer("packed_weight_e2m1", packed_weight_e2m1.contiguous())
331
+ self.register_buffer("packed_scales_ue4m3", packed_scales_ue4m3.contiguous())
332
+ self.register_buffer("weight_tensor_scale", weight_tensor_scale.contiguous())
333
+ if bias_bf16 is None:
334
+ self.bias_bf16 = None
335
+ else:
336
+ self.register_buffer("bias_bf16", bias_bf16.contiguous())
337
+
338
+ def forward(self, inputs: torch.Tensor) -> torch.Tensor:
339
+ raise RuntimeError(
340
+ "PackedNvfp4LinearArtifactModule is an artifact-only placeholder. "
341
+ "Attach the resident CUDA runtime before calling forward()."
342
+ )
343
+
344
+ def extra_repr(self) -> str:
345
+ return (
346
+ f"in_features={self.in_features}, out_features={self.out_features}, "
347
+ f"packed_weight_bytes={self.packed_weight_e2m1.numel()}, "
348
+ f"packed_scale_bytes={self.packed_scales_ue4m3.numel()}, "
349
+ f"has_bias={self.bias_bf16 is not None}"
350
+ )
351
+
352
+
353
+ def instantiate_mage_transformer_on_meta(source_repo: Path) -> nn.Module:
354
+ config = load_transformer_config(source_repo)
355
+ MageFlow, MageFlowParams = import_mage_transformer_symbols()
356
+ structure = {
357
+ key: value
358
+ for key, value in config.items()
359
+ if key
360
+ not in {
361
+ "_class_name",
362
+ "txt_max_length",
363
+ "max_sequence_length",
364
+ "param_dtype",
365
+ "packing",
366
+ "schedule_mode",
367
+ "static_shift",
368
+ "use_time_shift",
369
+ "rope_type",
370
+ "apply_text_rotary_emb",
371
+ "mlp_ratio",
372
+ "depth_single_blocks",
373
+ "theta",
374
+ "qkv_bias",
375
+ "guidance_embed",
376
+ "vec_in_dim",
377
+ "vec_type",
378
+ "time_type",
379
+ "double_block_type",
380
+ "quantization_config",
381
+ }
382
+ }
383
+ with torch.device("meta"):
384
+ model = MageFlow(MageFlowParams(**structure))
385
+ return model
386
+
387
+
388
+ def unregistered_meta_tensor_attribute_names(model: nn.Module) -> list[str]:
389
+ """Find direct tensor attributes that PyTorch's parameter/buffer walk misses."""
390
+ names: list[str] = []
391
+ for module_name, module in model.named_modules():
392
+ registered_names = set(module._parameters) | set(module._buffers)
393
+ for attribute_name, value in vars(module).items():
394
+ if attribute_name in registered_names:
395
+ continue
396
+ if isinstance(value, torch.Tensor) and value.is_meta:
397
+ prefix = f"{module_name}." if module_name else ""
398
+ names.append(f"{prefix}{attribute_name}")
399
+ return sorted(names)
400
+
401
+
402
+ def materialize_mage_rope_tensor_attributes(model: nn.Module) -> list[str]:
403
+ """Rebuild Mage's intentionally unregistered complex RoPE tensors on CPU."""
404
+ before = unregistered_meta_tensor_attribute_names(model)
405
+ expected = ["pos_embed.neg_freqs", "pos_embed.pos_freqs"]
406
+ if before != expected:
407
+ fail(
408
+ "unexpected unregistered meta tensor attributes before RoPE "
409
+ f"materialization: {before}"
410
+ )
411
+
412
+ rope = model.get_submodule("pos_embed")
413
+ rope_type = type(rope)
414
+ with torch.device("cpu"):
415
+ materialized = rope_type(
416
+ theta=rope.theta,
417
+ axes_dim=list(rope.axes_dim),
418
+ scale_rope=rope.scale_rope,
419
+ )
420
+ rope.pos_freqs = materialized.pos_freqs
421
+ rope.neg_freqs = materialized.neg_freqs
422
+ rope.video_freq_cache = {}
423
+
424
+ remaining = unregistered_meta_tensor_attribute_names(model)
425
+ if remaining:
426
+ fail(
427
+ "unresolved unregistered meta tensor attributes after RoPE "
428
+ f"materialization: {remaining}"
429
+ )
430
+ return before
431
+
432
+
433
+ def set_child_module(root: nn.Module, dotted_path: str, module: nn.Module) -> None:
434
+ parent_path, _, child_name = dotted_path.rpartition(".")
435
+ parent = root.get_submodule(parent_path) if parent_path else root
436
+ if child_name.isdigit() and isinstance(parent, (nn.Sequential, nn.ModuleList)):
437
+ parent[int(child_name)] = module
438
+ else:
439
+ setattr(parent, child_name, module)
440
+
441
+
442
+ def replace_targets_with_artifact_modules(
443
+ model: nn.Module,
444
+ artifact_path: Path,
445
+ target_specs: Iterable[TargetSpec],
446
+ ) -> None:
447
+ with safe_open(artifact_path, framework="pt", device="cpu") as handle:
448
+ for spec in target_specs:
449
+ packed_weight = handle.get_tensor(spec.artifact_weight_key)
450
+ packed_scales = handle.get_tensor(spec.artifact_scale_key)
451
+ weight_tensor_scale = handle.get_tensor(spec.artifact_tensor_scale_key)
452
+ bias = handle.get_tensor(spec.artifact_bias_key)
453
+ original = model.get_submodule(spec.module_key)
454
+ if not isinstance(original, nn.Linear):
455
+ fail(f"expected target module {spec.module_key} to be nn.Linear")
456
+ replacement = PackedNvfp4LinearArtifactModule(
457
+ in_features=int(original.in_features),
458
+ out_features=int(original.out_features),
459
+ packed_weight_e2m1=packed_weight,
460
+ packed_scales_ue4m3=packed_scales,
461
+ weight_tensor_scale=weight_tensor_scale,
462
+ bias_bf16=bias,
463
+ )
464
+ set_child_module(model, spec.module_key, replacement)
465
+
466
+
467
+ def replace_targets_with_resident_modules(
468
+ model: nn.Module,
469
+ artifact_path: Path,
470
+ target_specs: Iterable[TargetSpec],
471
+ device: torch.device,
472
+ ) -> None:
473
+ if device.type != "cuda":
474
+ fail("resident runtime modules require a CUDA destination")
475
+ if str(RESIDENT_PYTHON_ROOT) not in sys.path:
476
+ sys.path.insert(0, str(RESIDENT_PYTHON_ROOT))
477
+ from torch_ops import PackedNvfp4LinearOp
478
+
479
+ _replace_targets_with_registered_modules(
480
+ model,
481
+ artifact_path,
482
+ target_specs,
483
+ device,
484
+ PackedNvfp4LinearOp,
485
+ )
486
+
487
+
488
+ def replace_targets_with_native_resident_modules(
489
+ model: nn.Module,
490
+ artifact_path: Path,
491
+ target_specs: Iterable[TargetSpec],
492
+ device: torch.device,
493
+ ) -> None:
494
+ if device.type != "cuda":
495
+ fail("native resident runtime modules require a CUDA destination")
496
+ if str(RESIDENT_PYTHON_ROOT) not in sys.path:
497
+ sys.path.insert(0, str(RESIDENT_PYTHON_ROOT))
498
+ from torch_ops_native import (
499
+ PackedNvfp4LinearNativeOp,
500
+ initialize_native_sm120_op,
501
+ )
502
+
503
+ if not initialize_native_sm120_op(allow_python_schema_fallback=False):
504
+ fail("compiled native resident torch op did not load")
505
+ _replace_targets_with_registered_modules(
506
+ model,
507
+ artifact_path,
508
+ target_specs,
509
+ device,
510
+ PackedNvfp4LinearNativeOp,
511
+ )
512
+
513
+
514
+ def _replace_targets_with_registered_modules(
515
+ model: nn.Module,
516
+ artifact_path: Path,
517
+ target_specs: Iterable[TargetSpec],
518
+ device: torch.device,
519
+ module_cls: type[nn.Module],
520
+ ) -> None:
521
+ with safe_open(artifact_path, framework="pt", device="cpu") as handle:
522
+ for spec in target_specs:
523
+ original = model.get_submodule(spec.module_key)
524
+ if not isinstance(original, nn.Linear):
525
+ fail(f"expected target module {spec.module_key} to be nn.Linear")
526
+ replacement = module_cls(
527
+ in_features=int(original.in_features),
528
+ out_features=int(original.out_features),
529
+ packed_weight=handle.get_tensor(spec.artifact_weight_key).to(device),
530
+ weight_scales=handle.get_tensor(spec.artifact_scale_key).to(device),
531
+ weight_scale=handle.get_tensor(
532
+ spec.artifact_tensor_scale_key
533
+ ).to(device),
534
+ bias=handle.get_tensor(spec.artifact_bias_key).to(device),
535
+ )
536
+ set_child_module(model, spec.module_key, replacement)
537
+
538
+
539
+ def assign_tensor_by_name(model: nn.Module, key: str, tensor: torch.Tensor) -> None:
540
+ if "." not in key:
541
+ parent = model
542
+ leaf = key
543
+ else:
544
+ parent_path, _, leaf = key.rpartition(".")
545
+ parent = model.get_submodule(parent_path)
546
+ if leaf in parent._parameters:
547
+ requires_grad = parent._parameters[leaf].requires_grad
548
+ parent._parameters[leaf] = nn.Parameter(tensor, requires_grad=requires_grad)
549
+ return
550
+ if leaf in parent._buffers:
551
+ parent._buffers[leaf] = tensor
552
+ return
553
+ fail(f"destination key {key} was neither a parameter nor a buffer")
554
+
555
+
556
+ def pack_artifact(source_repo: Path, output_dir: Path) -> Path:
557
+ source_repo = source_repo.resolve()
558
+ output_dir = output_dir.resolve()
559
+ if output_dir.exists():
560
+ fail(f"output directory already exists: {output_dir}")
561
+ output_dir.mkdir(parents=True, exist_ok=False)
562
+
563
+ config = load_transformer_config(source_repo)
564
+ checkpoint_path = source_transformer_checkpoint(source_repo)
565
+ target_specs = build_target_specs(int(config["depth"]))
566
+ target_weight_keys = {spec.weight_key for spec in target_specs}
567
+ target_bias_keys = {spec.bias_key for spec in target_specs}
568
+ target_keys = target_weight_keys | target_bias_keys
569
+
570
+ artifact_tensors: OrderedDict[str, torch.Tensor] = OrderedDict()
571
+ target_metadata: list[dict] = []
572
+
573
+ with safe_open(checkpoint_path, framework="pt", device="cpu") as handle:
574
+ source_keys = list(handle.keys())
575
+ source_key_set = set(source_keys)
576
+ missing = sorted(target_keys - source_key_set)
577
+ if missing:
578
+ fail(f"source checkpoint is missing {len(missing)} target tensors, first={missing[0]}")
579
+
580
+ for spec in target_specs:
581
+ weight = handle.get_tensor(spec.weight_key)
582
+ bias = handle.get_tensor(spec.bias_key)
583
+ if weight.dtype != torch.bfloat16:
584
+ fail(f"{spec.weight_key} expected bfloat16, found {weight.dtype}")
585
+ if bias.dtype != torch.bfloat16:
586
+ fail(f"{spec.bias_key} expected bfloat16, found {bias.dtype}")
587
+ packed_weight, packed_scales, weight_tensor_scale, global_amax = pack_weight_tensor(weight)
588
+ scale_layout = make_scale_layout(weight.shape[1], weight.shape[0])
589
+ artifact_tensors[spec.artifact_bias_key] = bias.contiguous()
590
+ artifact_tensors[spec.artifact_scale_key] = packed_scales
591
+ artifact_tensors[spec.artifact_tensor_scale_key] = weight_tensor_scale
592
+ artifact_tensors[spec.artifact_weight_key] = packed_weight
593
+ target_metadata.append(
594
+ {
595
+ "module_key": spec.module_key,
596
+ "weight_key": spec.weight_key,
597
+ "bias_key": spec.bias_key,
598
+ "weight_shape": list(weight.shape),
599
+ "bias_shape": list(bias.shape),
600
+ "weight_tensor_scale_key": spec.artifact_tensor_scale_key,
601
+ "artifact_weight_key": spec.artifact_weight_key,
602
+ "artifact_scale_key": spec.artifact_scale_key,
603
+ "artifact_bias_key": spec.artifact_bias_key,
604
+ "weight_tensor_scale": float(weight_tensor_scale.item()),
605
+ "weight_global_amax": global_amax,
606
+ "packed_weight_bytes": int(packed_weight.numel()),
607
+ "packed_scale_bytes": int(packed_scales.numel()),
608
+ "scale_layout": asdict(scale_layout),
609
+ "source_weight_sha256": sha256_tensor(weight),
610
+ "source_bias_sha256": sha256_tensor(bias),
611
+ }
612
+ )
613
+
614
+ non_target_keys = sorted(set(source_keys) - target_keys)
615
+ artifact_path = output_dir / "packed_transformer.safetensors"
616
+ save_file(
617
+ OrderedDict(sorted(artifact_tensors.items())),
618
+ artifact_path,
619
+ metadata=CANONICAL_SAFETENSORS_METADATA,
620
+ )
621
+ fsync_file(artifact_path)
622
+ fsync_directory(output_dir)
623
+
624
+ library_hashes = {
625
+ "artifact_script_sha256": sha256_file(ARTIFACT_SCRIPT_PATH),
626
+ "resident_source_sha256": sha256_file(RESIDENT_SOURCE),
627
+ "resident_library_sha256": sha256_file(RESIDENT_LIBRARY),
628
+ "mage_flow_py_sha256": sha256_file(MAGE_ROOT / "mage_flow" / "models" / "mage_flow.py"),
629
+ "mage_layers_py_sha256": sha256_file(MAGE_ROOT / "mage_flow" / "models" / "modules" / "mage_layers.py"),
630
+ "pipeline_py_sha256": sha256_file(MAGE_ROOT / "mage_flow" / "pipeline.py"),
631
+ }
632
+
633
+ metadata = OrderedDict(
634
+ (
635
+ ("format_version", FORMAT_VERSION),
636
+ ("artifact_kind", ARTIFACT_KIND),
637
+ ("created_utc", datetime.now(timezone.utc).isoformat()),
638
+ (
639
+ "container",
640
+ {
641
+ "format": "safetensors",
642
+ "header_metadata": CANONICAL_SAFETENSORS_METADATA,
643
+ "header_encoding": (
644
+ "single deterministic contract key; legacy two-key "
645
+ "draft headers remain readable"
646
+ ),
647
+ },
648
+ ),
649
+ (
650
+ "source",
651
+ OrderedDict(
652
+ (
653
+ ("transformer_config_path", str(source_repo / "transformer" / "config.json")),
654
+ ("transformer_checkpoint_path", str(checkpoint_path)),
655
+ ("transformer_config_sha256", sha256_file(source_repo / "transformer" / "config.json")),
656
+ ("transformer_checkpoint_sha256", sha256_file(checkpoint_path)),
657
+ )
658
+ ),
659
+ ),
660
+ ("library_hashes", library_hashes),
661
+ (
662
+ "environment",
663
+ {
664
+ "python_version": sys.version,
665
+ "torch_version": torch.__version__,
666
+ "safetensors_version": safetensors.__version__,
667
+ },
668
+ ),
669
+ (
670
+ "model",
671
+ {
672
+ "depth": int(config["depth"]),
673
+ "hidden_size": int(config["hidden_size"]),
674
+ "num_heads": int(config["num_heads"]),
675
+ "context_in_dim": int(config["context_in_dim"]),
676
+ "in_channels": int(config["in_channels"]),
677
+ "out_channels": int(config["out_channels"]),
678
+ "patch_size": int(config["patch_size"]),
679
+ },
680
+ ),
681
+ (
682
+ "quantization",
683
+ {
684
+ "format": "nvfp4_two_level",
685
+ "block_elements": FP4_BLOCK_ELEMENTS,
686
+ "scale_tile_outer": SCALE_TILE_OUTER,
687
+ "scale_tile_inner": SCALE_TILE_INNER,
688
+ "bias_policy": "artifact_bfloat16",
689
+ },
690
+ ),
691
+ ("targets", target_metadata),
692
+ ("non_target_keys", non_target_keys),
693
+ )
694
+ )
695
+ write_bytes_once(
696
+ output_dir / "metadata.json",
697
+ (json.dumps(metadata, indent=2, sort_keys=False) + "\n").encode("utf-8"),
698
+ )
699
+ return output_dir
700
+
701
+
702
+ def load_validated_artifact_metadata(
703
+ artifact_dir: Path,
704
+ source_repo: Path,
705
+ *,
706
+ require_resident_runtime: bool = False,
707
+ ) -> dict:
708
+ artifact_dir = artifact_dir.resolve()
709
+ source_repo = source_repo.resolve()
710
+ metadata_path = artifact_dir / "metadata.json"
711
+ artifact_path = artifact_dir / "packed_transformer.safetensors"
712
+ if not metadata_path.is_file() or not artifact_path.is_file():
713
+ fail(f"artifact dir missing metadata or safetensors: {artifact_dir}")
714
+ try:
715
+ metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
716
+ except (OSError, json.JSONDecodeError) as exc:
717
+ fail(f"invalid artifact metadata {metadata_path}: {exc}")
718
+ if metadata.get("format_version") != FORMAT_VERSION:
719
+ fail(f"unsupported format_version: {metadata.get('format_version')}")
720
+ if metadata.get("artifact_kind") != ARTIFACT_KIND:
721
+ fail(f"unexpected artifact_kind: {metadata.get('artifact_kind')}")
722
+
723
+ config_path = source_repo / "transformer" / "config.json"
724
+ checkpoint_path = source_transformer_checkpoint(source_repo)
725
+ expected_config_hash = sha256_file(config_path)
726
+ expected_checkpoint_hash = sha256_file(checkpoint_path)
727
+ try:
728
+ source_metadata = metadata["source"]
729
+ recorded_config_hash = source_metadata["transformer_config_sha256"]
730
+ recorded_checkpoint_hash = source_metadata[
731
+ "transformer_checkpoint_sha256"
732
+ ]
733
+ model_metadata = metadata["model"]
734
+ recorded_depth = int(model_metadata["depth"])
735
+ recorded_targets = metadata["targets"]
736
+ recorded_non_target_keys = metadata["non_target_keys"]
737
+ except (KeyError, TypeError, ValueError) as exc:
738
+ fail(f"artifact metadata schema is incomplete or invalid: {exc}")
739
+ if not isinstance(recorded_targets, list):
740
+ fail("artifact metadata targets must be a list")
741
+ if not isinstance(recorded_non_target_keys, list) or not all(
742
+ isinstance(key, str) for key in recorded_non_target_keys
743
+ ):
744
+ fail("artifact metadata non_target_keys must be a list of strings")
745
+ if recorded_config_hash != expected_config_hash:
746
+ fail("transformer config hash mismatch")
747
+ if recorded_checkpoint_hash != expected_checkpoint_hash:
748
+ fail("transformer checkpoint hash mismatch")
749
+
750
+ config = load_transformer_config(source_repo)
751
+ if recorded_depth != int(config["depth"]):
752
+ fail("artifact model depth does not match source config")
753
+ specs = build_target_specs(int(config["depth"]))
754
+ expected_modules = [spec.module_key for spec in specs]
755
+ if not all(isinstance(entry, dict) for entry in recorded_targets):
756
+ fail("artifact metadata target entries must be objects")
757
+ recorded_modules = [entry.get("module_key") for entry in recorded_targets]
758
+ if recorded_modules != expected_modules:
759
+ fail("artifact target allowlist/order mismatch")
760
+ target_source_keys = {
761
+ key for spec in specs for key in (spec.weight_key, spec.bias_key)
762
+ }
763
+ expected_artifact_keys = {
764
+ key
765
+ for spec in specs
766
+ for key in (
767
+ spec.artifact_weight_key,
768
+ spec.artifact_scale_key,
769
+ spec.artifact_tensor_scale_key,
770
+ spec.artifact_bias_key,
771
+ )
772
+ }
773
+ with safe_open(artifact_path, framework="pt", device="cpu") as artifact_handle:
774
+ actual_artifact_keys = set(artifact_handle.keys())
775
+ header_metadata = artifact_handle.metadata()
776
+ if actual_artifact_keys != expected_artifact_keys:
777
+ fail("artifact tensor key coverage mismatch")
778
+ if header_metadata not in (
779
+ CANONICAL_SAFETENSORS_METADATA,
780
+ LEGACY_SAFETENSORS_METADATA,
781
+ ):
782
+ fail("artifact safetensors header metadata mismatch")
783
+
784
+ with safe_open(checkpoint_path, framework="pt", device="cpu") as source_handle:
785
+ source_keys = set(source_handle.keys())
786
+ missing_target_keys = sorted(target_source_keys - source_keys)
787
+ if missing_target_keys:
788
+ fail(
789
+ "source checkpoint is missing target tensors, first="
790
+ f"{missing_target_keys[0]}"
791
+ )
792
+ expected_non_target_keys = sorted(source_keys - target_source_keys)
793
+ if recorded_non_target_keys != expected_non_target_keys:
794
+ fail("artifact non-target source manifest mismatch")
795
+
796
+ if require_resident_runtime:
797
+ hashes = metadata.get("library_hashes", {})
798
+ if not isinstance(hashes, dict):
799
+ fail("artifact metadata library_hashes must be an object")
800
+ if hashes.get("resident_source_sha256") != sha256_file(RESIDENT_SOURCE):
801
+ fail("resident source hash mismatch")
802
+ if hashes.get("resident_library_sha256") != sha256_file(RESIDENT_LIBRARY):
803
+ fail("resident library hash mismatch")
804
+ return metadata
805
+
806
+
807
+ def validate_artifact(artifact_dir: Path, source_repo: Path) -> None:
808
+ artifact_dir = artifact_dir.resolve()
809
+ source_repo = source_repo.resolve()
810
+ metadata = load_validated_artifact_metadata(artifact_dir, source_repo)
811
+ artifact_path = artifact_dir / "packed_transformer.safetensors"
812
+ checkpoint_path = source_transformer_checkpoint(source_repo)
813
+ config = load_transformer_config(source_repo)
814
+ specs = {
815
+ spec.module_key: spec for spec in build_target_specs(int(config["depth"]))
816
+ }
817
+
818
+ with safe_open(artifact_path, framework="pt", device="cpu") as artifact_handle, safe_open(
819
+ checkpoint_path, framework="pt", device="cpu"
820
+ ) as source_handle:
821
+ for entry in metadata["targets"]:
822
+ spec = specs[entry["module_key"]]
823
+ weight = source_handle.get_tensor(spec.weight_key)
824
+ bias = source_handle.get_tensor(spec.bias_key)
825
+ packed_weight, packed_scales, weight_tensor_scale, global_amax = pack_weight_tensor(weight)
826
+
827
+ candidate_weight = artifact_handle.get_tensor(spec.artifact_weight_key)
828
+ candidate_scales = artifact_handle.get_tensor(spec.artifact_scale_key)
829
+ candidate_tensor_scale = artifact_handle.get_tensor(spec.artifact_tensor_scale_key)
830
+ candidate_bias = artifact_handle.get_tensor(spec.artifact_bias_key)
831
+
832
+ if not torch.equal(candidate_weight, packed_weight):
833
+ fail(f"packed weight mismatch for {spec.module_key}")
834
+ if not torch.equal(candidate_scales, packed_scales):
835
+ fail(f"packed scales mismatch for {spec.module_key}")
836
+ if not torch.equal(candidate_tensor_scale, weight_tensor_scale):
837
+ fail(f"tensor scale mismatch for {spec.module_key}")
838
+ if not torch.equal(candidate_bias, bias):
839
+ fail(f"bias mismatch for {spec.module_key}")
840
+ if abs(float(entry["weight_global_amax"]) - global_amax) > 0.0:
841
+ fail(f"global amax mismatch for {spec.module_key}")
842
+
843
+
844
+ def load_clean_transformer_from_artifact(
845
+ artifact_dir: Path,
846
+ source_repo: Path,
847
+ assign_non_target: bool = True,
848
+ ) -> nn.Module:
849
+ artifact_dir = artifact_dir.resolve()
850
+ source_repo = source_repo.resolve()
851
+ metadata = load_validated_artifact_metadata(artifact_dir, source_repo)
852
+ target_specs = build_target_specs(int(metadata["model"]["depth"]))
853
+ model = instantiate_mage_transformer_on_meta(source_repo)
854
+ replace_targets_with_artifact_modules(model, artifact_dir / "packed_transformer.safetensors", target_specs)
855
+
856
+ if assign_non_target:
857
+ skip_keys = {
858
+ key
859
+ for spec in target_specs
860
+ for key in (spec.weight_key, spec.bias_key)
861
+ }
862
+ source_tensor_keys_read: list[str] = []
863
+ with safe_open(source_transformer_checkpoint(source_repo), framework="pt", device="cpu") as source_handle:
864
+ for key in source_handle.keys():
865
+ if key in skip_keys:
866
+ continue
867
+ tensor = source_handle.get_tensor(key)
868
+ source_tensor_keys_read.append(key)
869
+ assign_tensor_by_name(model, key, tensor)
870
+ target_reads = sorted(set(source_tensor_keys_read) & skip_keys)
871
+ if target_reads:
872
+ fail(f"clean CPU loader read target source tensors: {target_reads[0]}")
873
+ meta_parameters = [
874
+ name for name, parameter in model.named_parameters() if parameter.is_meta
875
+ ]
876
+ meta_buffers = [
877
+ name for name, buffer in model.named_buffers() if buffer.is_meta
878
+ ]
879
+ if meta_parameters or meta_buffers:
880
+ first = (meta_parameters + meta_buffers)[0]
881
+ fail(f"clean CPU loader left unresolved meta tensors, first={first}")
882
+ materialize_mage_rope_tensor_attributes(model)
883
+ return model
884
+
885
+
886
+ def load_clean_resident_transformer(
887
+ artifact_dir: Path,
888
+ source_repo: Path,
889
+ device: torch.device,
890
+ ) -> tuple[nn.Module, dict]:
891
+ return _load_clean_cuda_transformer(
892
+ artifact_dir,
893
+ source_repo,
894
+ device,
895
+ replace_targets_with_resident_modules,
896
+ )
897
+
898
+
899
+ def load_clean_native_resident_transformer(
900
+ artifact_dir: Path,
901
+ source_repo: Path,
902
+ device: torch.device,
903
+ ) -> tuple[nn.Module, dict]:
904
+ return _load_clean_cuda_transformer(
905
+ artifact_dir,
906
+ source_repo,
907
+ device,
908
+ replace_targets_with_native_resident_modules,
909
+ )
910
+
911
+
912
+ def _load_clean_cuda_transformer(
913
+ artifact_dir: Path,
914
+ source_repo: Path,
915
+ device: torch.device,
916
+ target_replacement_fn: Callable[[nn.Module, Path, Iterable[TargetSpec], torch.device], None],
917
+ ) -> tuple[nn.Module, dict]:
918
+ artifact_dir = artifact_dir.resolve()
919
+ source_repo = source_repo.resolve()
920
+ metadata = load_validated_artifact_metadata(
921
+ artifact_dir,
922
+ source_repo,
923
+ require_resident_runtime=True,
924
+ )
925
+ target_specs = build_target_specs(int(metadata["model"]["depth"]))
926
+ target_source_keys = {
927
+ key
928
+ for spec in target_specs
929
+ for key in (spec.weight_key, spec.bias_key)
930
+ }
931
+ model = instantiate_mage_transformer_on_meta(source_repo)
932
+ target_replacement_fn(
933
+ model,
934
+ artifact_dir / "packed_transformer.safetensors",
935
+ target_specs,
936
+ device,
937
+ )
938
+
939
+ loaded_source_keys: list[str] = []
940
+ skipped_target_source_keys: list[str] = []
941
+ source_tensor_keys_read: list[str] = []
942
+ with safe_open(
943
+ source_transformer_checkpoint(source_repo),
944
+ framework="pt",
945
+ device="cpu",
946
+ ) as source_handle:
947
+ source_keys = list(source_handle.keys())
948
+ missing_target_keys = sorted(target_source_keys - set(source_keys))
949
+ if missing_target_keys:
950
+ fail(
951
+ "source checkpoint is missing target tensors, first="
952
+ f"{missing_target_keys[0]}"
953
+ )
954
+ for key in source_keys:
955
+ if key in target_source_keys:
956
+ skipped_target_source_keys.append(key)
957
+ continue
958
+ tensor = source_handle.get_tensor(key)
959
+ source_tensor_keys_read.append(key)
960
+ assign_tensor_by_name(model, key, tensor.to(device))
961
+ loaded_source_keys.append(key)
962
+
963
+ materialized_tensor_attributes = materialize_mage_rope_tensor_attributes(model)
964
+ target_source_reads = sorted(
965
+ set(source_tensor_keys_read) & target_source_keys
966
+ )
967
+ meta_parameters = [
968
+ name for name, parameter in model.named_parameters() if parameter.is_meta
969
+ ]
970
+ meta_buffers = [
971
+ name for name, buffer in model.named_buffers() if buffer.is_meta
972
+ ]
973
+ report = {
974
+ "source_tensor_count_loaded": len(loaded_source_keys),
975
+ "source_tensor_keys_loaded": loaded_source_keys,
976
+ "source_tensor_count_read": len(source_tensor_keys_read),
977
+ "source_tensor_keys_read": source_tensor_keys_read,
978
+ "target_source_tensor_count_skipped": len(skipped_target_source_keys),
979
+ "target_source_tensor_keys_skipped": skipped_target_source_keys,
980
+ "target_source_tensor_reads": len(target_source_reads),
981
+ "target_source_tensor_keys_read": target_source_reads,
982
+ "meta_parameter_names": meta_parameters,
983
+ "meta_buffer_names": meta_buffers,
984
+ "materialized_unregistered_tensor_attribute_names": (
985
+ materialized_tensor_attributes
986
+ ),
987
+ "unregistered_meta_tensor_attribute_names": (
988
+ unregistered_meta_tensor_attribute_names(model)
989
+ ),
990
+ }
991
+ return model.eval().requires_grad_(False), report
992
+
993
+
994
+ def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
995
+ parser = argparse.ArgumentParser(description=__doc__)
996
+ subparsers = parser.add_subparsers(dest="command", required=True)
997
+
998
+ pack_parser = subparsers.add_parser("pack", help="build a packed artifact directory")
999
+ pack_parser.add_argument("--source-repo", type=Path, required=True)
1000
+ pack_parser.add_argument("--output-dir", type=Path, required=True)
1001
+
1002
+ validate_parser = subparsers.add_parser("validate", help="recompute and validate a packed artifact")
1003
+ validate_parser.add_argument("--artifact-dir", type=Path, required=True)
1004
+ validate_parser.add_argument("--source-repo", type=Path, required=True)
1005
+
1006
+ plan_load_parser = subparsers.add_parser("plan-load", help="instantiate on meta and replace target modules")
1007
+ plan_load_parser.add_argument("--artifact-dir", type=Path, required=True)
1008
+ plan_load_parser.add_argument("--source-repo", type=Path, required=True)
1009
+ plan_load_parser.add_argument("--skip-non-target", action="store_true")
1010
+
1011
+ runtime_parser = subparsers.add_parser(
1012
+ "validate-runtime",
1013
+ help="placeholder for future single-GPU resident validation",
1014
+ )
1015
+ runtime_parser.add_argument("--artifact-dir", type=Path, required=True)
1016
+ runtime_parser.add_argument("--source-repo", type=Path, required=True)
1017
+ return parser.parse_args(argv)
1018
+
1019
+
1020
+ def main(argv: list[str] | None = None) -> int:
1021
+ args = parse_args(argv)
1022
+ try:
1023
+ if args.command == "pack":
1024
+ artifact_dir = pack_artifact(args.source_repo, args.output_dir)
1025
+ print(artifact_dir)
1026
+ return 0
1027
+ if args.command == "validate":
1028
+ validate_artifact(args.artifact_dir, args.source_repo)
1029
+ print("ok")
1030
+ return 0
1031
+ if args.command == "plan-load":
1032
+ model = load_clean_transformer_from_artifact(
1033
+ args.artifact_dir, args.source_repo, assign_non_target=not args.skip_non_target
1034
+ )
1035
+ packed_count = sum(
1036
+ 1 for _name, module in model.named_modules() if isinstance(module, PackedNvfp4LinearArtifactModule)
1037
+ )
1038
+ meta_parameters = [
1039
+ name for name, parameter in model.named_parameters() if parameter.is_meta
1040
+ ]
1041
+ meta_buffers = [
1042
+ name for name, buffer in model.named_buffers() if buffer.is_meta
1043
+ ]
1044
+ print(
1045
+ json.dumps(
1046
+ {
1047
+ "packed_module_count": packed_count,
1048
+ "meta_parameter_count": len(meta_parameters),
1049
+ "meta_buffer_count": len(meta_buffers),
1050
+ "non_target_assignment_skipped": bool(args.skip_non_target),
1051
+ },
1052
+ indent=2,
1053
+ )
1054
+ )
1055
+ return 0
1056
+ if args.command == "validate-runtime":
1057
+ fail(
1058
+ "validate-runtime is intentionally not implemented in this CPU-only slice. "
1059
+ "Use the future CUDA resident path on CUDA_VISIBLE_DEVICES=3."
1060
+ )
1061
+ fail(f"unsupported command: {args.command}")
1062
+ except PackedArtifactError as exc:
1063
+ print(f"error: {exc}", file=sys.stderr)
1064
+ return 1
1065
+
1066
+
1067
+ if __name__ == "__main__":
1068
+ raise SystemExit(main())