Instructions to use kingjones777/Ming-Image-0.1-Design-ROCm-INT8 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use kingjones777/Ming-Image-0.1-Design-ROCm-INT8 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("kingjones777/Ming-Image-0.1-Design-ROCm-INT8", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
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}")
|