"""Load a complete sharded Mage-Flow NVFP4 transformer component.""" from __future__ import annotations from contextlib import ExitStack import json import os from pathlib import Path from typing import Any import torch import torch.nn as nn from safetensors import safe_open from packed_artifact import ( assign_tensor_by_name, build_target_specs, instantiate_mage_transformer_on_meta, materialize_mage_rope_tensor_attributes, set_child_module, unregistered_meta_tensor_attribute_names, ) from torch_ops_native import ( PackedNvfp4LinearNativeOp, initialize_native_sm120_op, ) class StandardCheckpointError(RuntimeError): pass def fail(message: str) -> None: raise StandardCheckpointError(message) def read_object(path: Path) -> dict[str, Any]: try: value = json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as exc: fail(f"cannot read JSON object {path}: {exc}") if not isinstance(value, dict): fail(f"expected a JSON object: {path}") return value def _component_path(component_dir: Path, relative: str) -> Path: path = (component_dir / relative).resolve() if not path.is_relative_to(component_dir): fail(f"checkpoint index path escapes transformer component: {relative}") if not path.is_file(): fail(f"checkpoint shard is missing: {relative}") return path def _quantized_keys(module_key: str) -> dict[str, str]: return { "packed_weight": f"{module_key}.packed_weight", "weight_scales": f"{module_key}.weight_scales", "weight_scale": f"{module_key}.weight_scale", "bias": f"{module_key}.bias", } def _target_specs_for_config(depth: int, quant_config: dict[str, Any]) -> list[Any]: all_specs = {spec.module_key: spec for spec in build_target_specs(depth)} targets = quant_config.get("targets") if not isinstance(targets, list) or not all(isinstance(name, str) for name in targets): fail("transformer config has no valid target list") if len(set(targets)) != len(targets): fail("transformer target list contains duplicates") missing = [name for name in targets if name not in all_specs] if missing: fail(f"transformer target list contains an unknown module: {missing[0]}") declared_count = quant_config.get("target_count") if declared_count != len(targets): fail( "transformer quantization target count mismatch: " f"{declared_count!r} != {len(targets)}" ) return [all_specs[name] for name in targets] def _apply_runtime_defaults(quant_config: dict[str, Any]) -> dict[str, str]: defaults = quant_config.get("native_runtime_defaults") if not isinstance(defaults, dict): return {} applied: dict[str, str] = {} mapping = { "up_activation_scale_multiplier": "MAGE_NVFP4_UP_ACTIVATION_SCALE_MULTIPLIER", "down_activation_scale_multiplier": "MAGE_NVFP4_DOWN_ACTIVATION_SCALE_MULTIPLIER", "activation_scale_search": "MAGE_NVFP4_ACTIVATION_SCALE_SEARCH", } for key, env_name in mapping.items(): value = defaults.get(key) if value is None: continue os.environ.setdefault(env_name, str(value)) applied[env_name] = os.environ[env_name] return applied def load_standard_native_transformer( repo_root: str | Path, device: torch.device, ) -> tuple[nn.Module, dict[str, Any]]: """Load a complete standard-layout component without a BF16 base download.""" repo_root = Path(repo_root).resolve() component_dir = (repo_root / "transformer").resolve() if not component_dir.is_relative_to(repo_root) or not component_dir.is_dir(): fail("repository has no transformer component") if device.type != "cuda": fail("the native resident transformer requires a CUDA destination") config = read_object(component_dir / "config.json") quant_config = config.get("quantization_config") if not isinstance(quant_config, dict): fail("transformer config has no quantization_config") if quant_config.get("quant_method") != "mage_flow_nvfp4": fail( "unexpected transformer quantization method: " f"{quant_config.get('quant_method')!r}" ) if quant_config.get("quant_algo") != "NVFP4": fail("transformer config does not declare NVFP4") runtime_defaults = _apply_runtime_defaults(quant_config) if not initialize_native_sm120_op(allow_python_schema_fallback=False): fail("compiled native SM120 torch op did not load") depth = int(config.get("depth", 0)) target_specs = _target_specs_for_config(depth, quant_config) expected_targets = [spec.module_key for spec in target_specs] metadata = read_object(component_dir / "nvfp4_metadata.json") if metadata.get("artifact_kind") != ( "mage_flow_transformer_mlp_nvfp4_resident_v1" ): fail("unexpected transformer NVFP4 metadata kind") non_target_keys = metadata.get("non_target_keys") if not isinstance(non_target_keys, list) or not all( isinstance(key, str) for key in non_target_keys ): fail("transformer NVFP4 metadata has no non-target key list") recorded_targets = metadata.get("targets") if not isinstance(recorded_targets, list) or not all( isinstance(entry, dict) for entry in recorded_targets ): fail("transformer NVFP4 metadata has no valid targets list") recorded_modules = [entry.get("module_key") for entry in recorded_targets] if recorded_modules != expected_targets: fail("transformer config targets do not match NVFP4 metadata targets") index = read_object( component_dir / "diffusion_pytorch_model.safetensors.index.json" ) weight_map = index.get("weight_map") if not isinstance(weight_map, dict) or not all( isinstance(key, str) and isinstance(value, str) for key, value in weight_map.items() ): fail("transformer checkpoint has no valid weight map") quantized_keys = { key for spec in target_specs for key in _quantized_keys(spec.module_key).values() } expected_keys = set(non_target_keys) | quantized_keys actual_keys = set(weight_map) if actual_keys != expected_keys: missing = sorted(expected_keys - actual_keys) unexpected = sorted(actual_keys - expected_keys) fail( "transformer checkpoint key coverage mismatch; " f"missing={missing[:1]}, unexpected={unexpected[:1]}" ) original_target_weights = {spec.weight_key for spec in target_specs} leaked = sorted(actual_keys & original_target_weights) if leaked: fail(f"BF16 target weight leaked into quantized checkpoint: {leaked[0]}") shard_names = sorted(set(weight_map.values())) with ExitStack() as stack: handles = { name: stack.enter_context( safe_open( _component_path(component_dir, name), framework="pt", device="cpu", ) ) for name in shard_names } def tensor(key: str) -> torch.Tensor: try: handle = handles[weight_map[key]] except KeyError: fail(f"tensor is absent from checkpoint index: {key}") if key not in handle.keys(): fail(f"tensor is absent from its declared shard: {key}") return handle.get_tensor(key) model = instantiate_mage_transformer_on_meta(repo_root) for spec in target_specs: original = model.get_submodule(spec.module_key) if not isinstance(original, nn.Linear): fail( f"expected target {spec.module_key} to be nn.Linear, " f"found {type(original).__name__}" ) keys = _quantized_keys(spec.module_key) replacement = PackedNvfp4LinearNativeOp( in_features=int(original.in_features), out_features=int(original.out_features), packed_weight=tensor(keys["packed_weight"]).to(device), weight_scales=tensor(keys["weight_scales"]).to(device), weight_scale=tensor(keys["weight_scale"]).to(device), bias=tensor(keys["bias"]).to(device), ) set_child_module(model, spec.module_key, replacement) loaded_non_targets: list[str] = [] for key in non_target_keys: assign_tensor_by_name(model, key, tensor(key).to(device)) loaded_non_targets.append(key) materialized = materialize_mage_rope_tensor_attributes(model) meta_parameters = [ name for name, value in model.named_parameters() if value.is_meta ] meta_buffers = [ name for name, value in model.named_buffers() if value.is_meta ] unregistered_meta = unregistered_meta_tensor_attribute_names(model) if meta_parameters or meta_buffers or unregistered_meta: fail( "standard transformer loader left unresolved meta tensors: " f"{(meta_parameters + meta_buffers + unregistered_meta)[:4]}" ) report = { "layout": "huggingface_sharded_component", "checkpoint_shard_count": len(shard_names), "checkpoint_tensor_count": len(actual_keys), "loaded_non_target_tensor_count": len(loaded_non_targets), "loaded_quantized_projection_count": len(target_specs), "bf16_target_weight_reads": 0, "meta_parameter_names": meta_parameters, "meta_buffer_names": meta_buffers, "materialized_unregistered_tensor_attribute_names": materialized, "runtime_defaults_applied": runtime_defaults, } return model.eval().requires_grad_(False), report __all__ = [ "StandardCheckpointError", "load_standard_native_transformer", ]