Mage-Flow-Edit-XPO3-NVFP4 / runtime /standard_transformer.py
ajh-code's picture
Add runtime/standard_transformer.py
371a6d9 verified
Raw
History Blame Contribute Delete
10 kB
"""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",
]