"""Load a streamed INT8 Ming MLLM checkpoint onto a meta-initialized model. ``model`` must already exist with parameters on ``meta`` (for example under ``accelerate.init_empty_weights()``). Quantized modules listed in ``int8_manifest.json`` are swapped from ``nn.Linear`` to ``Int8Linear.shell`` before the shards are assigned in. """ from __future__ import annotations import json from pathlib import Path import torch from safetensors.torch import load_file from torch import nn try: # imported as the `quant` package (modeling_bailingmm2.py) from .int8_linear import Int8Linear except ImportError: # run from inside quant/ (CLI, tests) from int8_linear import Int8Linear MANIFEST_NAME = "int8_manifest.json" INDEX_NAME = "model.safetensors.index.json" def load_int8_mllm_(model: nn.Module, int8_dir, device) -> dict: """Swap quantize-rule linears for INT8 shells and assign shard tensors. Returns ``{"modules_swapped", "tensors_loaded", "bytes_loaded"}``. Raises ``RuntimeError`` on a bad manifest, a module that is not an ``nn.Linear``, an unexpected checkpoint key, or any parameter / persistent buffer still on ``meta``. Non-persistent buffers (rotary ``inv_freq``) may stay on CPU; the caller moves the model afterwards. """ int8_dir = Path(int8_dir) dev = torch.device(device) if not isinstance(device, torch.device) else device manifest_path = int8_dir / MANIFEST_NAME if not manifest_path.is_file(): raise RuntimeError(f"missing int8 manifest: {manifest_path}") manifest = json.loads(manifest_path.read_text(encoding="utf-8")) if manifest.get("format") != "ming-int8-wo-v1": raise RuntimeError( f"unsupported int8 manifest format: {manifest.get('format')!r} ({manifest_path})" ) module_names = manifest.get("quantized_modules") if not isinstance(module_names, list) or not all(isinstance(n, str) for n in module_names): raise RuntimeError(f"{manifest_path} quantized_modules is not a list of strings") swapped = _swap_linears(model, module_names) index_path = int8_dir / INDEX_NAME if not index_path.is_file(): raise RuntimeError(f"missing index: {index_path}") index = json.loads(index_path.read_text(encoding="utf-8")) weight_map = index.get("weight_map") if not isinstance(weight_map, dict) or not weight_map: raise RuntimeError(f"{index_path} has no weight_map") shard_names: list[str] = [] seen: set[str] = set() for shard in weight_map.values(): if shard not in seen: seen.add(shard) shard_names.append(shard) tensors_loaded = 0 bytes_loaded = 0 unexpected: list[str] = [] for shard in shard_names: rel = Path(shard) if rel.is_absolute() or ".." in rel.parts: raise RuntimeError(f"unsafe shard path in index: {shard}") path = int8_dir / rel if not path.is_file(): raise RuntimeError(f"missing shard: {path}") sd = load_file(str(path), device=str(dev)) for tensor in sd.values(): tensors_loaded += 1 bytes_loaded += tensor.numel() * tensor.element_size() incompatible = model.load_state_dict(sd, strict=False, assign=True) unexpected.extend(incompatible.unexpected_keys) del sd if unexpected: listed = "\n".join(f" {key}" for key in unexpected) raise RuntimeError( f"unexpected keys in checkpoint (not present on the model):\n{listed}" ) _assert_loaded(model, module_names, dev) return { "modules_swapped": swapped, "tensors_loaded": tensors_loaded, "bytes_loaded": bytes_loaded, } def _swap_linears(model: nn.Module, module_names: list[str]) -> int: for name in module_names: try: linear = model.get_submodule(name) except AttributeError as exc: raise RuntimeError(f"manifest module not found on model: {name}") from exc if not isinstance(linear, nn.Linear): raise RuntimeError( f"{name} is {type(linear).__name__}, expected nn.Linear " "(refusing to swap a router or other non-linear)" ) parent_name, _, leaf = name.rpartition(".") if not leaf: raise RuntimeError(f"cannot place shell for {name}") parent = model.get_submodule(parent_name) if parent_name else model has_bias = linear.bias is not None bias_dtype = linear.bias.dtype if has_bias else torch.float32 shell = Int8Linear.shell( in_features=linear.in_features, out_features=linear.out_features, bias=has_bias, bias_dtype=bias_dtype, device="meta", ) setattr(parent, leaf, shell) return len(module_names) def _assert_loaded(model: nn.Module, module_names: list[str], dev: torch.device) -> None: offenders: list[str] = [] for name, param in model.named_parameters(remove_duplicate=False): if param is not None and param.device.type == "meta": offenders.append(f"parameter {name} dtype={param.dtype} device={param.device}") for mod_name, mod in model.named_modules(): nonpersist = getattr(mod, "_non_persistent_buffers_set", set()) for buf_name, buf in mod._buffers.items(): if buf is None: continue full = f"{mod_name}.{buf_name}" if mod_name else buf_name if buf.device.type != "meta": # Non-persistent buffers (rotary inv_freq) are not in the # checkpoint. accelerate leaves them on CPU; that is not an error. continue if buf_name in nonpersist: offenders.append( f"non-persistent buffer {full} dtype={buf.dtype} device={buf.device}" ) else: offenders.append(f"buffer {full} dtype={buf.dtype} device={buf.device}") if offenders: listed = "\n".join(f" {line}" for line in offenders) raise RuntimeError(f"tensors still on meta after load:\n{listed}") for name in module_names: mod = model.get_submodule(name) if not isinstance(mod, Int8Linear): raise RuntimeError(f"{name} was not swapped to Int8Linear") if mod.weight is None or mod.weight.dtype != torch.int8: raise RuntimeError(f"{name}.weight is not int8 after load") if mod.scale is None or mod.scale.dtype != torch.float32: raise RuntimeError(f"{name}.scale is not float32 after load") if mod.weight.device.type == "meta" or mod.scale.device.type == "meta": raise RuntimeError(f"{name} still has meta tensors after load") if mod.weight.device != dev or mod.scale.device != dev: raise RuntimeError( f"{name} loaded on weight={mod.weight.device} scale={mod.scale.device}, " f"expected {dev}" ) if tuple(mod.scale.shape) != (mod.out_features,): raise RuntimeError( f"{name}.scale shape {tuple(mod.scale.shape)} != ({mod.out_features},)" ) if mod.bias is not None and mod.bias.device != dev: raise RuntimeError(f"{name}.bias is on {mod.bias.device}, expected {dev}")