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
Download code/quant/load_int8.py from kingjones777/Ming-Image-0.1-Design-ROCm-INT8: direct link, hf CLI and curl.
- Browser
- Download file 7.36 kB
-
https://huggingface.co/kingjones777/Ming-Image-0.1-Design-ROCm-INT8/resolve/main/code/quant/load_int8.py
- Command line
-
hf download hf://kingjones777/Ming-Image-0.1-Design-ROCm-INT8/code/quant/load_int8.py
-
curl -L -o load_int8.py https://huggingface.co/kingjones777/Ming-Image-0.1-Design-ROCm-INT8/resolve/main/code/quant/load_int8.py
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}") | |