File size: 7,358 Bytes
da1a4ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
"""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}")