kingjones777's picture
Add files using upload-large-folder tool
da1a4ff verified
Raw History Blame Contribute Delete
7.36 kB
"""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}")