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: 8,722 Bytes
18c1466 | 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 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 | """Weight-only symmetric per-output-channel INT8 linear.
Scales stay float32 across dtype casts. ``module.to(dtype=torch.bfloat16)``
(and ``.bfloat16()`` / ``.half()`` / ``.to(device, dtype)``) must not touch them;
device moves still do. The int8 weight codes are likewise dtype-stable.
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
# Leaf names of Linear modules whose 2-D weights are quantized.
# Exact match: the routers are `gate` / `image_gate` / `audio_gate`, NOT `gate_proj`.
QUANT_LEAVES = frozenset(
{"query_key_value", "dense", "gate_proj", "up_proj", "down_proj"}
)
QUANT_RULE = (
"Quantize ONLY 2-D .weight tensors under model.model.layers. whose owning "
"module's leaf name is exactly one of query_key_value, dense, gate_proj, "
"up_proj, down_proj. Everything else stays byte-identical BF16: embeddings, "
"lm_head, all norms, the vision tower, linear_proj, and the three routers "
"(modules named gate, image_gate, audio_gate — leaf match, not a substring). "
"Per-output-channel symmetric: scale = absmax/127, "
"q = clamp(round(w/scale), -127, 127). All-zero rows: scale 1.0, q 0."
)
# Real checkpoint keys look like `model.model.layers.N...`. A bare
# `layers.N...` name is the same stack with the root prefix omitted (tests).
_DECODER_LAYER_PREFIXES = ((), ("model", "model"))
def _weight_leaf(tensor_name: str) -> str | None:
"""Owning module's leaf name if `tensor_name` ends in `.weight`, else None."""
if not isinstance(tensor_name, str) or not tensor_name.endswith(".weight"):
return None
module = tensor_name[: -len(".weight")]
if not module:
return None
return module.rsplit(".", 1)[-1]
def _under_decoder_layers(tensor_name: str) -> bool:
"""True when the tensor lives under the MLLM decoder `model.model.layers` stack.
`layers` must be its own path component, followed by a layer index. The
components before it must be empty or end in `model.model` — so a vision
tower that happens to contain the substring "layers" is not selected, and
`gate` is never selected just because `gate_proj` contains those letters.
"""
parts = tensor_name.split(".")
for i, part in enumerate(parts):
if part != "layers":
continue
if i + 1 >= len(parts) or not parts[i + 1].isdigit():
continue
prefix = tuple(parts[:i])
if prefix in _DECODER_LAYER_PREFIXES:
return True
if len(prefix) >= 2 and prefix[-2:] == ("model", "model"):
return True
return False
def quant_rule_leaf(tensor_name: str) -> str | None:
"""Leaf name if the name matches the quantize rule, ignoring rank.
Returns None when the tensor is not a candidate. A candidate whose rank is
not 2 is a hard error for the stream (see quantize_stream); ``is_quantizable``
itself returns False for that case.
"""
leaf = _weight_leaf(tensor_name)
if leaf not in QUANT_LEAVES:
return None
if not _under_decoder_layers(tensor_name):
return None
return leaf
def is_quantizable(tensor_name: str, shape) -> bool:
"""True only for 2-D quantize-rule weights. See ``QUANT_RULE``."""
if quant_rule_leaf(tensor_name) is None:
return False
try:
rank = len(shape)
except TypeError:
return False
return rank == 2
def quantize_weight(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Per-output-channel symmetric int8.
``scale = absmax(row) / 127``, ``q = clamp(round(w / scale), -127, 127)``.
An all-zero row gets scale 1.0 and q 0 (no div-by-zero, no NaN/Inf).
"""
if weight.ndim != 2:
raise ValueError(
f"quantize_weight expects a 2-D weight, got shape {tuple(weight.shape)}"
)
wf = weight.detach().to(dtype=torch.float32)
absmax = wf.abs().amax(dim=1)
scale = absmax / 127.0
zero = scale == 0
# All-zero rows would divide by 0. Force scale 1 and q 0 instead of NaN.
scale = torch.where(zero, torch.ones_like(scale), scale)
q = torch.round(wf / scale[:, None]).clamp(-127, 127).to(dtype=torch.int8)
q = torch.where(zero[:, None], torch.zeros_like(q), q)
return q.contiguous(), scale.to(dtype=torch.float32).contiguous()
def _scale_name(weight_name: str) -> str:
if not weight_name.endswith(".weight"):
raise ValueError(f"not a weight tensor name: {weight_name}")
return weight_name[: -len("weight")] + "scale"
class Int8Linear(nn.Module):
"""``F.linear`` on a weight dequantized from int8 + per-row float32 scale.
``weight`` is int8 ``[out, in]``, ``scale`` is float32 ``[out]``, ``bias``
(optional) keeps the source dtype. All three are buffers.
"""
def __init__(self, weight: torch.Tensor, scale: torch.Tensor, bias: torch.Tensor | None):
super().__init__()
if weight.dtype != torch.int8 or weight.ndim != 2:
raise ValueError(
f"weight must be int8 [out, in], got dtype={weight.dtype} shape={tuple(weight.shape)}"
)
if scale.dtype != torch.float32 or tuple(scale.shape) != (weight.shape[0],):
raise ValueError(
f"scale must be float32 [{weight.shape[0]}], got dtype={scale.dtype} shape={tuple(scale.shape)}"
)
if bias is not None:
if bias.ndim != 1 or bias.shape[0] != weight.shape[0]:
raise ValueError(
f"bias must be [{weight.shape[0]}], got shape={tuple(bias.shape)}"
)
self.in_features = int(weight.shape[1])
self.out_features = int(weight.shape[0])
self.register_buffer("weight", weight)
self.register_buffer("scale", scale)
self.register_buffer("bias", bias)
def _apply(self, fn, *args, **kwargs):
# Pull dtype-stable buffers out before Module._apply. Putting them back
# with only a device move (never fn's dtype cast) keeps scale float32
# and weight int8. Bias is left in the dict so it follows the cast.
saved: dict[str, torch.Tensor] = {}
for name in ("weight", "scale"):
buf = self._buffers.get(name, None)
if buf is not None:
saved[name] = buf
self._buffers[name] = None
try:
out = super()._apply(fn, *args, **kwargs)
finally:
for name, buf in saved.items():
self._buffers[name] = _move_device_keep_dtype(buf, fn)
return out
def forward(self, x: torch.Tensor) -> torch.Tensor:
# One dequant in fp32, one cast to the activation dtype, then linear.
w = (self.weight.float() * self.scale[:, None]).to(dtype=x.dtype)
return F.linear(x, w, self.bias)
@classmethod
def from_linear(cls, linear: nn.Linear) -> "Int8Linear":
if not isinstance(linear, nn.Linear):
raise TypeError(f"from_linear expects nn.Linear, got {type(linear).__name__}")
q, scale = quantize_weight(linear.weight.data)
if linear.bias is None:
bias = None
else:
bias = linear.bias.detach().clone()
return cls(q, scale, bias)
@classmethod
def shell(
cls,
in_features: int,
out_features: int,
bias: bool,
bias_dtype: torch.dtype,
device,
) -> "Int8Linear":
"""Empty buffers (for ``meta``). Does not read or write any weight values."""
dev = torch.device(device) if not isinstance(device, torch.device) else device
weight = torch.empty((out_features, in_features), dtype=torch.int8, device=dev)
scale = torch.empty((out_features,), dtype=torch.float32, device=dev)
if bias:
bias_t: torch.Tensor | None = torch.empty(
(out_features,), dtype=bias_dtype, device=dev
)
else:
bias_t = None
return cls(weight, scale, bias_t)
def extra_repr(self) -> str:
return (
f"in_features={self.in_features}, out_features={self.out_features}, "
f"bias={self.bias is not None}"
)
def _move_device_keep_dtype(buf: torch.Tensor, fn) -> torch.Tensor:
"""Apply only the device change implied by ``fn``, preserving ``buf``'s dtype and values.
Probed with a 0-element tensor so a dtype cast cannot round the real scale.
"""
try:
probe = torch.empty((), dtype=buf.dtype, device=buf.device)
moved = fn(probe)
except Exception:
return buf
if not torch.is_tensor(moved) or moved.device == buf.device:
return buf
return buf.to(device=moved.device)
|