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
| diff --git a/configuration_bailingmm2.py b/configuration_bailingmm2.py | |
| index 5ab2542..b20eca3 100644 | |
| --- a/configuration_bailingmm2.py | |
| +++ b/configuration_bailingmm2.py | |
| from configuration_bailing_moe_v2 import BailingMoeV2Config | |
| class BailingMM2Config(PretrainedConfig): | |
| model_type = "bailingmm_moe_v2_lite" | |
| + # Declared so transformers' `_attn_implementation` setter recurses into both towers. | |
| + # Without it an explicit attn_implementation (e.g. "eager" on ROCm, which has no | |
| + # flash-attn) never reaches them, and their "flash_attention_2" defaults raise at | |
| + # model construction. | |
| + sub_configs = {"vision_config": Qwen2_5_VLVisionConfig, "llm_config": BailingMoeV2Config} | |
| def __init__( | |
| self, | |
| diff --git a/diffusion/transformer.py b/diffusion/transformer.py | |
| index d47ca9f..89a2845 100644 | |
| --- a/diffusion/transformer.py | |
| +++ b/diffusion/transformer.py | |
| ADALN_EMBED_DIM = 256 | |
| SEQ_MULTI_OF = 32 | |
| +def _native_sdpa_is_active(processor) -> bool: | |
| + """True when attention would go to diffusers' default native SDPA backend (no per-model backend, | |
| + no context parallelism, and the active global backend is NATIVE).""" | |
| + if processor._attention_backend is not None or processor._parallel_config is not None: | |
| + return False | |
| + try: | |
| + from diffusers.models.attention_dispatch import AttentionBackendName, _AttentionBackendRegistry | |
| + | |
| + name, _ = _AttentionBackendRegistry.get_active_backend() | |
| + return name == AttentionBackendName.NATIVE | |
| + except Exception: | |
| + return False | |
| + | |
| + | |
| class TimestepEmbedder(nn.Module): | |
| def __init__(self, out_size, mid_size=None, frequency_embedding_size=256): | |
| super().__init__() | |
| class SingleStreamAttentionProcessor: | |
| attention_mask = attention_mask[:, None, None, :] | |
| # Compute joint attention | |
| - hidden_states = dispatch_attention_fn( | |
| - query, | |
| - key, | |
| - value, | |
| - attn_mask=attention_mask, | |
| - dropout_p=0.0, | |
| - is_causal=False, | |
| - backend=self._attention_backend, | |
| - parallel_config=self._parallel_config, | |
| - ) | |
| + if _native_sdpa_is_active(self): | |
| + # What diffusers' default "native" backend computes, but SDPA receives contiguous | |
| + # [B, H, L, D] tensors instead of permuted views. PyTorch's math SDPA (the only SDPA | |
| + # kernel that runs on ROCm gfx1151) is ~2x faster on contiguous inputs, with | |
| + # bit-identical output. | |
| + q, k, v = (x.transpose(1, 2).contiguous() for x in (query, key, value)) | |
| + hidden_states = F.scaled_dot_product_attention( | |
| + q, k, v, attn_mask=attention_mask, dropout_p=0.0, is_causal=False | |
| + ).transpose(1, 2) | |
| + else: | |
| + hidden_states = dispatch_attention_fn( | |
| + query, | |
| + key, | |
| + value, | |
| + attn_mask=attention_mask, | |
| + dropout_p=0.0, | |
| + is_causal=False, | |
| + backend=self._attention_backend, | |
| + parallel_config=self._parallel_config, | |
| + ) | |
| # Reshape back | |
| hidden_states = hidden_states.flatten(2, 3) | |
| diff --git a/generate_paired.sh b/generate_paired.sh | |
| new file mode 100755 | |
| index 0000000..864f462 | |
| --- /dev/null | |
| +++ b/generate_paired.sh | |
| +#!/usr/bin/env bash | |
| +# Paired pipeline: Ling-3.0-flash-VL prompt enhancement -> Ming-Image text-to-image. | |
| +# | |
| +# Stage 1 pe_ling.py caption -> validated structured JSON prompt | |
| +# (system prompt: assets/t2i_rewriter_system_prompt.txt) | |
| +# Stage 2 infer.py --task text-to-image --prompt <json file> -> PNG | |
| +# (infer.py reads --prompt as a file when the path exists) | |
| +# | |
| +# Artifacts land in --output-dir: enhanced_prompt.json (overwritten per run) | |
| +# plus the PNG(s) infer.py writes (image_00.png for text-to-image). | |
| +# Fails loudly at every stage (set -Eeuo pipefail + ERR trap + stage checks). | |
| +set -Eeuo pipefail | |
| +trap 'printf "generate_paired: FAILED at line %d (exit %d)\n" "$LINENO" "$?" >&2' ERR | |
| + | |
| +SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" | |
| +PYTHON="${PYTHON:-python3}" | |
| + | |
| +# Local llama-server seat serving Ling-3.0-flash-VL on the target box. | |
| +DEFAULT_BASE_URL="http://127.0.0.1:8090/v1" | |
| +DEFAULT_PE_MODEL="ling-3.0-flash-vl-mtp-halo-STRIX_LEAN" | |
| + | |
| +usage() { | |
| + cat <<'EOF' | |
| +Usage: generate_paired.sh --model DIR_OR_REPO CAPTION [options] [-- EXTRA_INFER_ARGS...] | |
| + | |
| +Enhances CAPTION with Ling-3.0-flash-VL (pe_ling.py), validates the structured | |
| +JSON rewrite, then renders it with infer.py --task text-to-image. | |
| + | |
| +Required: | |
| + CAPTION free-form design caption (positional) | |
| + --model DIR_OR_REPO Ming checkpoint directory or HF repo id | |
| + (may also be set via the MING_MODEL environment variable) | |
| + | |
| +Passthrough to infer.py (all optional; infer.py defaults in parentheses): | |
| + --resolution N resolution bucket, 1024 or 2048 for text-to-image; | |
| + other positive values snap to the nearest bucket (2048) | |
| + --seed N generation seed (42) | |
| + --steps N diffusion steps (12) | |
| + -- everything after this is passed to infer.py verbatim | |
| + (e.g. -- --validate-only --dtype float16) | |
| + | |
| +Prompt-enhancement endpoint: | |
| + --base-url URL OpenAI-compatible base URL (http://127.0.0.1:8090/v1) | |
| + --pe-model ID chat model id served there | |
| + (ling-3.0-flash-vl-mtp-halo-STRIX_LEAN) | |
| + LITELLM_API_KEY env exported key is sent as a Bearer token (for a gated | |
| + OpenAI-compatible gateway such as LiteLLM) | |
| + | |
| +Other: | |
| + --output-dir DIR artifact directory (outputs/paired) | |
| + -h, --help this help | |
| + | |
| +Examples: | |
| + ./generate_paired.sh --model /models/Ming-Image-0.1-Design \ | |
| + "espresso machine product poster, warm morning light" --resolution 2048 | |
| + | |
| + LITELLM_API_KEY=sk-... ./generate_paired.sh \ | |
| + --base-url http://<gateway-host>:4000/v1 --pe-model <gateway-model-name> \ | |
| + --model /models/Ming-Image-0.1-Design "a caption" --seed 7 | |
| +EOF | |
| +} | |
| + | |
| +die() { | |
| + printf 'generate_paired: %s\n' "$*" >&2 | |
| + exit 1 | |
| +} | |
| + | |
| +model="${MING_MODEL:-}" | |
| +base_url="$DEFAULT_BASE_URL" | |
| +pe_model="$DEFAULT_PE_MODEL" | |
| +output_dir="outputs/paired" | |
| +resolution="" | |
| +seed="" | |
| +steps="" | |
| +caption="" | |
| +extra_infer_args=() | |
| + | |
| +while [[ $# -gt 0 ]]; do | |
| + case "$1" in | |
| + --model) [[ $# -ge 2 ]] || die "--model requires a value"; model="$2"; shift 2 ;; | |
| + --base-url) [[ $# -ge 2 ]] || die "--base-url requires a value"; base_url="$2"; shift 2 ;; | |
| + --pe-model) [[ $# -ge 2 ]] || die "--pe-model requires a value"; pe_model="$2"; shift 2 ;; | |
| + --output-dir) [[ $# -ge 2 ]] || die "--output-dir requires a value"; output_dir="$2"; shift 2 ;; | |
| + --resolution) [[ $# -ge 2 ]] || die "--resolution requires a value"; resolution="$2"; shift 2 ;; | |
| + --seed) [[ $# -ge 2 ]] || die "--seed requires a value"; seed="$2"; shift 2 ;; | |
| + --steps) [[ $# -ge 2 ]] || die "--steps requires a value"; steps="$2"; shift 2 ;; | |
| + -h|--help) usage; exit 0 ;; | |
| + --) shift; extra_infer_args+=("$@"); break ;; | |
| + -*) usage >&2; die "unknown option: $1" ;; | |
| + *) | |
| + if [[ -n "$caption" ]]; then | |
| + usage >&2 | |
| + die "unexpected extra argument: $1 (CAPTION was already given)" | |
| + fi | |
| + caption="$1" | |
| + shift | |
| + ;; | |
| + esac | |
| +done | |
| + | |
| +[[ -n "$caption" ]] || { usage >&2; die "CAPTION is required"; } | |
| +[[ -n "$model" ]] || { usage >&2; die "--model DIR_OR_REPO is required (or set MING_MODEL)"; } | |
| +if [[ -n "$resolution" && ! "$resolution" =~ ^[0-9]+$ ]]; then | |
| + die "--resolution must be a positive integer, got: $resolution" | |
| +fi | |
| +if [[ -n "$seed" && ! "$seed" =~ ^-?[0-9]+$ ]]; then | |
| + die "--seed must be an integer, got: $seed" | |
| +fi | |
| +if [[ -n "$steps" && ! "$steps" =~ ^[0-9]+$ ]]; then | |
| + die "--steps must be a positive integer, got: $steps" | |
| +fi | |
| +[[ -f "$SCRIPT_DIR/pe_ling.py" ]] || die "missing stage-1 script: $SCRIPT_DIR/pe_ling.py" | |
| +[[ -f "$SCRIPT_DIR/infer.py" ]] || die "missing stage-2 script: $SCRIPT_DIR/infer.py" | |
| +command -v "$PYTHON" >/dev/null 2>&1 || die "python interpreter not found: $PYTHON (override with PYTHON=...)" | |
| + | |
| +mkdir -p -- "$output_dir" || die "cannot create output directory: $output_dir" | |
| +prompt_json="$output_dir/enhanced_prompt.json" | |
| + | |
| +printf '== stage 1/2: prompt enhancement (pe_ling.py, model %s @ %s)\n' "$pe_model" "$base_url" >&2 | |
| +"$PYTHON" "$SCRIPT_DIR/pe_ling.py" "$caption" \ | |
| + --out "$prompt_json" \ | |
| + --base-url "$base_url" \ | |
| + --model "$pe_model" | |
| +[[ -s "$prompt_json" ]] || die "prompt enhancement produced no prompt file: $prompt_json" | |
| + | |
| +printf '== stage 2/2: Ming-Image text-to-image (infer.py, model %s)\n' "$model" >&2 | |
| +infer_args=( | |
| + --model "$model" | |
| + --task text-to-image | |
| + --prompt "$prompt_json" | |
| + --output-dir "$output_dir" | |
| +) | |
| +if [[ -n "$resolution" ]]; then infer_args+=(--resolution "$resolution"); fi | |
| +if [[ -n "$seed" ]]; then infer_args+=(--seed "$seed"); fi | |
| +if [[ -n "$steps" ]]; then infer_args+=(--steps "$steps"); fi | |
| +if [[ ${#extra_infer_args[@]} -gt 0 ]]; then infer_args+=("${extra_infer_args[@]}"); fi | |
| +validate_only=0 | |
| +for arg in ${extra_infer_args[@]+"${extra_infer_args[@]}"}; do | |
| + if [[ "$arg" == "--validate-only" ]]; then validate_only=1; fi | |
| +done | |
| +"$PYTHON" "$SCRIPT_DIR/infer.py" "${infer_args[@]}" | |
| + | |
| +if [[ "$validate_only" -eq 1 ]]; then | |
| + printf 'generate_paired: --validate-only dry run, no PNG expected; enhanced prompt: %s\n' \ | |
| + "$prompt_json" >&2 | |
| + exit 0 | |
| +fi | |
| + | |
| +# infer.py exits non-zero on failure (set -e above); additionally verify the | |
| +# promised PNG artifacts actually exist so a silent no-write still fails | |
| +# loudly. -newer pins the check to THIS run: stage 2 always writes its PNG | |
| +# after stage 1 wrote enhanced_prompt.json, so stale PNGs do not satisfy it. | |
| +pngs=() | |
| +while IFS= read -r png; do | |
| + pngs+=("$png") | |
| +done < <(find "$output_dir" -maxdepth 1 -name '*.png' -type f -newer "$prompt_json" | sort) | |
| +if [[ ${#pngs[@]} -eq 0 ]]; then | |
| + die "infer.py exited 0 but wrote no PNG under $output_dir in this run" | |
| +fi | |
| +printf 'generate_paired: enhanced prompt: %s\n' "$prompt_json" >&2 | |
| +printf 'generate_paired: %d PNG(s):\n' "${#pngs[@]}" >&2 | |
| +printf '%s\n' "${pngs[@]}" | |
| diff --git a/infer.py b/infer.py | |
| index 9dda84a..814e8f3 100644 | |
| --- a/infer.py | |
| +++ b/infer.py | |
| def parse_args() -> argparse.Namespace: | |
| action="store_true", | |
| help="Validate model profile and task arguments without loading weights", | |
| ) | |
| + parser.add_argument( | |
| + "--attention-bf16-reduction", | |
| + action="store_true", | |
| + help=( | |
| + "Let PyTorch's math attention kernel (the only SDPA kernel that runs on ROCm gfx1151) " | |
| + "stay in bf16 instead of upcasting to fp32: faster, less precise." | |
| + ), | |
| + ) | |
| + parser.add_argument( | |
| + "--release-mllm-after-conditioning", | |
| + action="store_true", | |
| + help=( | |
| + "Free the MLLM, vision tower and connector as soon as the conditioning is computed, " | |
| + "before the diffusion steps. Lowers peak memory; one image per process." | |
| + ), | |
| + ) | |
| return parser.parse_args() | |
| def load_model_and_processor(model_directory: Path, args): | |
| ) | |
| processor = load_bailingmm2_processor(processor_directory) | |
| + if getattr(args, "attention_bf16_reduction", False): | |
| + # The math SDPA kernel upcasts bf16 inputs to fp32 by default; this keeps it in bf16. | |
| + torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(True) | |
| + | |
| dtype = _dtype(args.dtype) | |
| load_kwargs = { | |
| "torch_dtype": dtype, | |
| def load_model_and_processor(model_directory: Path, args): | |
| model = model.to(device=args.device, dtype=dtype) | |
| elif device_plan is not None: | |
| _validate_balanced_placement(model, device_plan, torch) | |
| + if getattr(args, "release_mllm_after_conditioning", False): | |
| + _release_mllm_before_sampling(model) | |
| return model, processor | |
| +def _release_mllm_before_sampling(model) -> None: | |
| + """Free the MLLM-side modules once the conditioning exists (--release-mllm-after-conditioning). | |
| + | |
| + Wraps the diffusion sampler: by the time it is called the conditioning tensors are computed, | |
| + so the language model, vision tower and connector are moved to the meta device (releasing | |
| + their memory) before the diffusion steps start. The model cannot generate again afterwards. | |
| + """ | |
| + import gc | |
| + | |
| + import torch | |
| + | |
| + original_sample = model.diffusion_loss.sample | |
| + | |
| + def sample_after_release(*args, **kwargs): | |
| + for name in ("model", "vision", "linear_proj", "connector"): | |
| + module = getattr(model, name, None) | |
| + if module is not None: | |
| + module.to("meta") | |
| + gc.collect() | |
| + if torch.cuda.is_available(): | |
| + torch.cuda.empty_cache() | |
| + return original_sample(*args, **kwargs) | |
| + | |
| + model.diffusion_loss.sample = sample_after_release | |
| + | |
| + | |
| def run_generation( | |
| model, | |
| processor, | |
| diff --git a/modeling_bailing_moe_v2.py b/modeling_bailing_moe_v2.py | |
| index a608f45..b62b66c 100644 | |
| --- a/modeling_bailing_moe_v2.py | |
| +++ b/modeling_bailing_moe_v2.py | |
| import torch.nn.functional as F | |
| import torch.utils.checkpoint | |
| from torch import nn | |
| from torch.nn import CrossEntropyLoss | |
| -import transformer_engine.pytorch as te | |
| from transformers.activations import ACT2FN | |
| from transformers.cache_utils import Cache, DynamicCache | |
| from transformers.modeling_attn_mask_utils import ( | |
| diff --git a/modeling_bailingmm2.py b/modeling_bailingmm2.py | |
| index fc0ecfa..27552f3 100644 | |
| --- a/modeling_bailingmm2.py | |
| +++ b/modeling_bailingmm2.py | |
| class BailingMM2NativeForConditionalGeneration(PreTrainedModel): | |
| self.diffusion_loss.to(device) | |
| self.loaded_image_gen_modules = True | |
| @classmethod | |
| + def _from_int8_checkpoint(cls, vlm_directory, device, **kwargs): | |
| + """Load an mllm/ component written by quant/quantize_stream.py (weight-only int8). | |
| + | |
| + The model is built with its parameters on the meta device, the Linear modules listed | |
| + in int8_manifest.json become Int8Linear shells, and every stored tensor is loaded | |
| + straight onto `device`, so BF16 weights for the quantized modules never exist in memory. | |
| + """ | |
| + from accelerate import init_empty_weights | |
| + from quant.load_int8 import load_int8_mllm_ | |
| + | |
| + device_map = kwargs.pop("device_map", None) | |
| + if device_map is not None: | |
| + # infer.py's default "balanced" plan on a single-GPU box maps every module to GPU 0; | |
| + # that is honoured. Splitting the int8 model across devices is not supported. | |
| + targets = set(device_map.values()) if isinstance(device_map, dict) else {device_map} | |
| + if len(targets) != 1 or not isinstance(next(iter(targets)), int): | |
| + raise ValueError( | |
| + "the int8 mllm checkpoint loads onto a single GPU; device_map targets " | |
| + f"{sorted(map(str, targets))} (use --device-map none)" | |
| + ) | |
| + device = torch.device("cuda", next(iter(targets))) | |
| + supported = ("torch_dtype", "dtype", "attn_implementation") | |
| + unsupported = sorted(key for key in kwargs if key not in supported) | |
| + if unsupported: | |
| + raise ValueError( | |
| + f"the int8 mllm checkpoint loads onto a single device; unsupported arguments: {unsupported}" | |
| + ) | |
| + device = torch.device(device) if device is not None else torch.device("cpu") | |
| + if device.type == "cuda" and device.index is None: | |
| + device = torch.device("cuda", torch.cuda.current_device()) | |
| + config = BailingMM2Config.from_pretrained(vlm_directory) | |
| + with init_empty_weights(): | |
| + model = cls._from_config(config, **kwargs) | |
| + report = load_int8_mllm_(model, vlm_directory, device) | |
| + # Buffers built in __init__ (rotary inv_freq) are not stored in the checkpoint; they follow | |
| + # the weights. Int8Linear keeps its fp32 scales through this and any later dtype cast. | |
| + model.to(device) | |
| + logger.info(f"int8 mllm loaded from {vlm_directory}: {report}") | |
| + model.tie_weights() | |
| + model.eval() | |
| + return model | |
| + | |
| + @classmethod | |
| def from_pretrained( | |
| cls, | |
| pretrained_model_name_or_path: Optional[Union[str, os.PathLike]], | |
| class BailingMM2NativeForConditionalGeneration(PreTrainedModel): | |
| f"{vlm_directory}. Migrate the package to the component " | |
| "layout before loading." | |
| ) | |
| - if load_vlm: | |
| + if load_vlm and os.path.exists(os.path.join(vlm_directory, "int8_manifest.json")): | |
| + model = cls._from_int8_checkpoint(vlm_directory, image_gen_device, **kwargs) | |
| + elif load_vlm: | |
| model = super().from_pretrained( | |
| vlm_directory, | |
| *model_args, | |
| diff --git a/pe_ling.py b/pe_ling.py | |
| new file mode 100644 | |
| index 0000000..88da86f | |
| --- /dev/null | |
| +++ b/pe_ling.py | |
| +#!/usr/bin/env python3 | |
| +"""Prompt enhancement (PE) for Ming-Image text-to-image via a Ling-3.0-flash-VL seat. | |
| + | |
| +Per the README, PE is a pre-processing step *outside* ``infer.py``: an | |
| +instruction-following VLM rewrites a short caption into the structured | |
| +Figma-style JSON prompt that the text-to-image pipeline consumes, and the | |
| +result is passed to ``infer.py --prompt`` as raw text or via a file. | |
| + | |
| +This module drives any OpenAI-compatible ``/chat/completions`` endpoint using | |
| +only the standard library (``urllib``): by default the local llama-server seat | |
| +serving Ling-3.0-flash-VL, optionally the LiteLLM lab gateway (Bearer auth via | |
| +``--api-key`` or the ``LITELLM_API_KEY`` environment variable). The rewriter | |
| +system prompt is read verbatim from ``assets/t2i_rewriter_system_prompt.txt``. | |
| + | |
| +The reply is parsed robustly (```json fences and surrounding prose are | |
| +tolerated), then validated against the schema the system prompt demands. On a | |
| +parse or validation failure the request is retried exactly once with the | |
| +errors appended to the user turn; if that still fails, PromptEnhancementError | |
| +is raised with the errors. Invalid JSON is never passed through silently. | |
| + | |
| +CLI: | |
| + python pe_ling.py "a caption" --out prompt.json \ | |
| + [--base-url http://127.0.0.1:8090/v1] \ | |
| + [--model ling-3.0-flash-vl-mtp-halo-STRIX_LEAN] | |
| +""" | |
| + | |
| +from __future__ import annotations | |
| + | |
| +import argparse | |
| +import json | |
| +import os | |
| +import re | |
| +import sys | |
| +import time | |
| +import urllib.error | |
| +import urllib.request | |
| +from pathlib import Path | |
| +from typing import Any, Dict, List, Optional, Tuple | |
| + | |
| +CODE_DIRECTORY = Path(__file__).resolve().parent | |
| +SYSTEM_PROMPT_PATH = CODE_DIRECTORY / "assets" / "t2i_rewriter_system_prompt.txt" | |
| + | |
| +# The Ling-3.0-flash-VL seat already served on the target box (llama-server, | |
| +# OpenAI-compatible, thinking disabled); both endpoints speak the same | |
| +# /chat/completions protocol. | |
| +DEFAULT_BASE_URL = "http://127.0.0.1:8090/v1" | |
| +DEFAULT_MODEL = "ling-3.0-flash-vl-mtp-halo-STRIX_LEAN" | |
| +API_KEY_ENV = "LITELLM_API_KEY" | |
| + | |
| +# Low temperature: the rewrite is a deterministic schema transformation, not | |
| +# creative sampling. | |
| +DEFAULT_TEMPERATURE = 0.2 | |
| +# The upstream example rewrite (assets/t2i_four_seasons_cabin_prompt.json) is | |
| +# ~5 KB (~2k tokens); dense multi-layer infographic rewrites run several times | |
| +# longer, so leave generous headroom for a complete JSON object. | |
| +DEFAULT_MAX_TOKENS = 16384 | |
| +# A multi-thousand-token completion on the local seat can take minutes. | |
| +DEFAULT_TIMEOUT_SECONDS = 600.0 | |
| + | |
| +REPAIR_INSTRUCTION = "Return only the corrected JSON object: no prose, no code fences." | |
| + | |
| +CANVAS_SETTINGS_KEYS = ("aspect_ratio", "ambient_lighting", "image_style") | |
| +LAYER_KEYS = ("description", "coordinates", "hierarchy_and_relation", "color_specs") | |
| +COORDINATE_FIELDS = ("cx", "cy", "w", "h") | |
| + | |
| +# `coordinates` must be ONE string of the form | |
| +# "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000". The upstream example also uses | |
| +# bare integers ("h: 1"), so accept any decimal spelling and enforce the | |
| +# [0, 1] range on the parsed value. Whitespace around ':' and ',' is | |
| +# tolerated; the key order is fixed. | |
| +_COORDINATE_NUMBER = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)" | |
| +COORDINATES_RE = re.compile( | |
| + rf"^\s*cx:\s*(?P<cx>{_COORDINATE_NUMBER})\s*," | |
| + rf"\s*cy:\s*(?P<cy>{_COORDINATE_NUMBER})\s*," | |
| + rf"\s*w:\s*(?P<w>{_COORDINATE_NUMBER})\s*," | |
| + rf"\s*h:\s*(?P<h>{_COORDINATE_NUMBER})\s*$" | |
| +) | |
| + | |
| +# Hex colors: #RGB, #RGBA, #RRGGBB, #RRGGBBAA (the upstream example uses | |
| +# #RRGGBB; the alpha forms keep RGBA-design outputs from failing validation). | |
| +HEX_COLOR_RE = re.compile( | |
| + r"^#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{4}|[0-9a-fA-F]{6}|[0-9a-fA-F]{8})$" | |
| +) | |
| + | |
| + | |
| +class PromptEnhancementError(RuntimeError): | |
| + """PE failed: transport/protocol error, or schema failure after the retry.""" | |
| + | |
| + def __init__( | |
| + self, | |
| + message: str, | |
| + errors: Optional[List[str]] = None, | |
| + reply: Optional[str] = None, | |
| + ): | |
| + super().__init__(message) | |
| + self.errors = list(errors or []) | |
| + self.reply = reply | |
| + | |
| + | |
| +def load_system_prompt(path: Path = SYSTEM_PROMPT_PATH) -> str: | |
| + """Return the released rewriter system prompt, verbatim.""" | |
| + return path.read_text(encoding="utf-8") | |
| + | |
| + | |
| +def extract_json_object(text: str) -> Dict[str, Any]: | |
| + """Return the first complete top-level JSON object found in ``text``. | |
| + | |
| + Models sometimes wrap JSON in ```json fences or add prose around it. | |
| + Scanning every ``{`` position with ``JSONDecoder.raw_decode`` (which | |
| + decodes a document at an offset and ignores trailing data) recovers the | |
| + object in all of those shapes. Raises ValueError when no complete JSON | |
| + object is present, e.g. a reply truncated mid-object. | |
| + """ | |
| + decoder = json.JSONDecoder() | |
| + position = text.find("{") | |
| + while position != -1: | |
| + try: | |
| + document, _ = decoder.raw_decode(text, position) | |
| + except ValueError: | |
| + position = text.find("{", position + 1) | |
| + continue | |
| + return document | |
| + snippet = text.strip() | |
| + if len(snippet) > 300: | |
| + snippet = snippet[:300] + "..." | |
| + raise ValueError( | |
| + f"reply contains no complete top-level JSON object " | |
| + f"({len(text)} characters); starts with: {snippet!r}" | |
| + ) | |
| + | |
| + | |
| +def _check_exact_keys( | |
| + mapping: Dict[str, Any], expected: Tuple[str, ...], path: str, errors: List[str] | |
| +) -> None: | |
| + missing = [key for key in expected if key not in mapping] | |
| + unexpected = [key for key in mapping if key not in expected] | |
| + if missing: | |
| + errors.append(f"{path}: missing required key(s): {', '.join(missing)}") | |
| + if unexpected: | |
| + errors.append( | |
| + f"{path}: unexpected key(s): {', '.join(unexpected)} " | |
| + f"(exactly {', '.join(expected)} are required)" | |
| + ) | |
| + | |
| + | |
| +def _check_non_empty_string(value: Any, path: str, errors: List[str]) -> None: | |
| + if not isinstance(value, str): | |
| + errors.append(f"{path}: expected a string, got {type(value).__name__}") | |
| + elif not value.strip(): | |
| + errors.append(f"{path}: string is empty") | |
| + | |
| + | |
| +def _check_coordinates(value: Any, path: str, errors: List[str]) -> None: | |
| + if not isinstance(value, str): | |
| + errors.append( | |
| + f"{path}: must be ONE string of the form " | |
| + f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000', got {type(value).__name__}" | |
| + ) | |
| + return | |
| + match = COORDINATES_RE.match(value) | |
| + if match is None: | |
| + errors.append( | |
| + f"{path}: {value!r} is not of the form " | |
| + f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000'" | |
| + ) | |
| + return | |
| + for field in COORDINATE_FIELDS: | |
| + number = float(match.group(field)) | |
| + if not 0.0 <= number <= 1.0: | |
| + errors.append(f"{path}: {field}={match.group(field)} is outside [0, 1]") | |
| + | |
| + | |
| +def _check_color_specs(value: Any, path: str, errors: List[str]) -> None: | |
| + if not isinstance(value, list): | |
| + errors.append( | |
| + f"{path}: expected a list of hex colors, got {type(value).__name__}" | |
| + ) | |
| + return | |
| + for index, color in enumerate(value): | |
| + if not isinstance(color, str) or HEX_COLOR_RE.match(color) is None: | |
| + errors.append( | |
| + f"{path}[{index}]: {color!r} is not a hex color " | |
| + f"(expected #RGB, #RGBA, #RRGGBB, or #RRGGBBAA)" | |
| + ) | |
| + | |
| + | |
| +def validate_enhanced_prompt(document: Any) -> List[str]: | |
| + """Return schema errors for a rewritten prompt; an empty list means valid. | |
| + | |
| + Schema demanded by assets/t2i_rewriter_system_prompt.txt: exactly two | |
| + top-level keys ``canvas_settings`` (exactly ``aspect_ratio``, | |
| + ``ambient_lighting``, ``image_style``) and ``layers`` (each layer exactly | |
| + ``description``, ``coordinates``, ``hierarchy_and_relation``, | |
| + ``color_specs``); ``coordinates`` is a string "cx: 0.500, cy: 0.500, | |
| + w: 1.000, h: 1.000" with values in [0, 1]; ``color_specs`` is a list of | |
| + hex colors. ``layers`` must hold at least one visible layer -- an empty | |
| + list means the rewrite failed even though it is type-correct. | |
| + """ | |
| + if not isinstance(document, dict): | |
| + return [f"top level: expected a JSON object, got {type(document).__name__}"] | |
| + errors: List[str] = [] | |
| + _check_exact_keys(document, ("canvas_settings", "layers"), "top level", errors) | |
| + | |
| + if "canvas_settings" in document: | |
| + canvas = document["canvas_settings"] | |
| + if not isinstance(canvas, dict): | |
| + errors.append( | |
| + f"canvas_settings: expected a JSON object, got {type(canvas).__name__}" | |
| + ) | |
| + else: | |
| + _check_exact_keys(canvas, CANVAS_SETTINGS_KEYS, "canvas_settings", errors) | |
| + for key in CANVAS_SETTINGS_KEYS: | |
| + if key in canvas: | |
| + _check_non_empty_string( | |
| + canvas[key], f"canvas_settings.{key}", errors | |
| + ) | |
| + | |
| + if "layers" in document: | |
| + layers = document["layers"] | |
| + if not isinstance(layers, list): | |
| + errors.append(f"layers: expected a list, got {type(layers).__name__}") | |
| + elif not layers: | |
| + errors.append("layers: expected at least one visible layer") | |
| + else: | |
| + for index, layer in enumerate(layers): | |
| + path = f"layers[{index}]" | |
| + if not isinstance(layer, dict): | |
| + errors.append( | |
| + f"{path}: expected a JSON object, got {type(layer).__name__}" | |
| + ) | |
| + continue | |
| + _check_exact_keys(layer, LAYER_KEYS, path, errors) | |
| + for key in ("description", "hierarchy_and_relation"): | |
| + if key in layer: | |
| + _check_non_empty_string(layer[key], f"{path}.{key}", errors) | |
| + if "coordinates" in layer: | |
| + _check_coordinates( | |
| + layer["coordinates"], f"{path}.coordinates", errors | |
| + ) | |
| + if "color_specs" in layer: | |
| + _check_color_specs(layer["color_specs"], f"{path}.color_specs", errors) | |
| + return errors | |
| + | |
| + | |
| +def _chat_completion( | |
| + base_url: str, | |
| + model: str, | |
| + messages: List[Dict[str, str]], | |
| + *, | |
| + temperature: float, | |
| + max_tokens: int, | |
| + api_key: Optional[str], | |
| + timeout: float, | |
| +) -> Tuple[str, Optional[str]]: | |
| + """POST one chat completion; return (content, finish_reason).""" | |
| + url = base_url.rstrip("/") + "/chat/completions" | |
| + payload = json.dumps( | |
| + { | |
| + "model": model, | |
| + "messages": messages, | |
| + "temperature": temperature, | |
| + "max_tokens": max_tokens, | |
| + "stream": False, | |
| + } | |
| + ).encode("utf-8") | |
| + headers = {"Content-Type": "application/json"} | |
| + if api_key: | |
| + headers["Authorization"] = f"Bearer {api_key}" | |
| + request = urllib.request.Request(url, data=payload, headers=headers, method="POST") | |
| + try: | |
| + with urllib.request.urlopen(request, timeout=timeout) as response: | |
| + body = response.read().decode("utf-8", errors="replace") | |
| + except urllib.error.HTTPError as error: | |
| + detail = error.read().decode("utf-8", errors="replace") | |
| + raise PromptEnhancementError( | |
| + f"HTTP {error.code} from {url}: {detail[:2000]}" | |
| + ) from error | |
| + except urllib.error.URLError as error: | |
| + raise PromptEnhancementError(f"cannot reach {url}: {error.reason}") from error | |
| + except OSError as error: # includes socket timeouts during the read | |
| + raise PromptEnhancementError(f"request to {url} failed: {error}") from error | |
| + | |
| + try: | |
| + envelope = json.loads(body) | |
| + choice = envelope["choices"][0] | |
| + content = choice["message"]["content"] | |
| + except (json.JSONDecodeError, KeyError, IndexError, TypeError) as error: | |
| + raise PromptEnhancementError( | |
| + f"malformed chat completion response from {url}: {body[:500]}" | |
| + ) from error | |
| + finish_reason = choice.get("finish_reason") | |
| + if not isinstance(content, str) or not content.strip(): | |
| + raise PromptEnhancementError( | |
| + f"empty completion content from {url} (finish_reason={finish_reason!r})" | |
| + ) | |
| + return content, finish_reason | |
| + | |
| + | |
| +def enhance( | |
| + caption: str, | |
| + base_url: str, | |
| + model: str, | |
| + api_key: Optional[str] = None, | |
| + timeout: float = DEFAULT_TIMEOUT_SECONDS, | |
| + temperature: float = DEFAULT_TEMPERATURE, | |
| + max_tokens: int = DEFAULT_MAX_TOKENS, | |
| +) -> Dict[str, Any]: | |
| + """Return the validated structured rewrite of ``caption``. | |
| + | |
| + Sends the verbatim rewriter system prompt plus the caption to | |
| + ``{base_url}/chat/completions``. On a parse or schema failure, retries | |
| + exactly once with the validation errors appended to the user turn; if | |
| + that also fails, raises PromptEnhancementError carrying the errors. | |
| + """ | |
| + system_prompt = load_system_prompt() | |
| + messages = [ | |
| + {"role": "system", "content": system_prompt}, | |
| + {"role": "user", "content": caption}, | |
| + ] | |
| + request_kwargs = { | |
| + "temperature": temperature, | |
| + "max_tokens": max_tokens, | |
| + "api_key": api_key, | |
| + "timeout": timeout, | |
| + } | |
| + errors: List[str] = [] | |
| + content = "" | |
| + for attempt in (1, 2): | |
| + content, finish_reason = _chat_completion( | |
| + base_url, model, messages, **request_kwargs | |
| + ) | |
| + document: Optional[Dict[str, Any]] = None | |
| + try: | |
| + document = extract_json_object(content) | |
| + except ValueError as error: | |
| + errors = [str(error)] | |
| + if document is not None: | |
| + errors = validate_enhanced_prompt(document) | |
| + if not errors: | |
| + assert document is not None # errors empty implies extraction succeeded | |
| + return document | |
| + if finish_reason == "length": | |
| + errors.append( | |
| + "the reply was cut off (finish_reason='length'): the complete " | |
| + f"JSON object must fit within max_tokens={max_tokens}" | |
| + ) | |
| + print(f"pe_ling: attempt {attempt}/2 failed validation:", file=sys.stderr) | |
| + for error in errors: | |
| + print(f"pe_ling: - {error}", file=sys.stderr) | |
| + if attempt == 1: | |
| + retry_content = ( | |
| + f"{caption}\n\n" | |
| + "Your previous reply failed schema validation:\n" | |
| + + "".join(f"- {error}\n" for error in errors) | |
| + + "\n" | |
| + + REPAIR_INSTRUCTION | |
| + ) | |
| + messages = [ | |
| + {"role": "system", "content": system_prompt}, | |
| + {"role": "user", "content": retry_content}, | |
| + ] | |
| + raise PromptEnhancementError( | |
| + "prompt enhancement failed schema validation after 2 attempts:\n" | |
| + + "".join(f" - {error}\n" for error in errors).rstrip(), | |
| + errors=errors, | |
| + reply=content, | |
| + ) | |
| + | |
| + | |
| +def main() -> None: | |
| + parser = argparse.ArgumentParser( | |
| + description=( | |
| + "Enhance a Ming-Image text-to-image caption into the validated " | |
| + "structured JSON prompt via an OpenAI-compatible Ling-3.0-flash-VL " | |
| + "endpoint." | |
| + ) | |
| + ) | |
| + parser.add_argument("caption", help="free-form design caption to enhance") | |
| + parser.add_argument( | |
| + "--out", | |
| + type=Path, | |
| + help="write the validated JSON here (default: stdout, summary on stderr)", | |
| + ) | |
| + parser.add_argument( | |
| + "--base-url", | |
| + default=DEFAULT_BASE_URL, | |
| + help=f"OpenAI-compatible base URL (default: {DEFAULT_BASE_URL})", | |
| + ) | |
| + parser.add_argument( | |
| + "--model", | |
| + default=DEFAULT_MODEL, | |
| + help=f"chat model id served at the endpoint (default: {DEFAULT_MODEL})", | |
| + ) | |
| + parser.add_argument( | |
| + "--api-key", | |
| + default=os.environ.get(API_KEY_ENV), | |
| + help=f"Bearer token for gated endpoints; defaults to ${API_KEY_ENV} when set", | |
| + ) | |
| + parser.add_argument( | |
| + "--timeout", | |
| + type=float, | |
| + default=DEFAULT_TIMEOUT_SECONDS, | |
| + help=f"per-request timeout in seconds (default: {DEFAULT_TIMEOUT_SECONDS})", | |
| + ) | |
| + parser.add_argument( | |
| + "--temperature", | |
| + type=float, | |
| + default=DEFAULT_TEMPERATURE, | |
| + help=f"sampling temperature (default: {DEFAULT_TEMPERATURE})", | |
| + ) | |
| + parser.add_argument( | |
| + "--max-tokens", | |
| + type=int, | |
| + default=DEFAULT_MAX_TOKENS, | |
| + help=f"completion token budget (default: {DEFAULT_MAX_TOKENS})", | |
| + ) | |
| + args = parser.parse_args() | |
| + | |
| + started = time.perf_counter() | |
| + try: | |
| + document = enhance( | |
| + args.caption, | |
| + args.base_url, | |
| + args.model, | |
| + api_key=args.api_key, | |
| + timeout=args.timeout, | |
| + temperature=args.temperature, | |
| + max_tokens=args.max_tokens, | |
| + ) | |
| + except PromptEnhancementError as error: | |
| + print(f"pe_ling: {error}", file=sys.stderr) | |
| + raise SystemExit(1) | |
| + elapsed = time.perf_counter() - started | |
| + layer_count = len(document["layers"]) | |
| + payload = json.dumps(document, indent=2, ensure_ascii=False) + "\n" | |
| + if args.out is not None: | |
| + args.out.parent.mkdir(parents=True, exist_ok=True) | |
| + args.out.write_text(payload, encoding="utf-8") | |
| + print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s) -> {args.out}") | |
| + else: | |
| + sys.stdout.write(payload) | |
| + print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s)", file=sys.stderr) | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |
| diff --git a/quant/__init__.py b/quant/__init__.py | |
| new file mode 100644 | |
| index 0000000..2e60ff4 | |
| --- /dev/null | |
| +++ b/quant/__init__.py | |
| @@ -0,0 +1 @@ | |
| +"""Weight-only INT8 for the Ming-Image MLLM: quantize_stream.py writes it, load_int8.py loads it.""" | |
| diff --git a/quant/int8_linear.py b/quant/int8_linear.py | |
| new file mode 100644 | |
| index 0000000..5939ea5 | |
| --- /dev/null | |
| +++ b/quant/int8_linear.py | |
| +"""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) | |
| diff --git a/quant/load_int8.py b/quant/load_int8.py | |
| new file mode 100644 | |
| index 0000000..1117892 | |
| --- /dev/null | |
| +++ b/quant/load_int8.py | |
| +"""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}") | |
| diff --git a/quant/quantize_stream.py b/quant/quantize_stream.py | |
| new file mode 100644 | |
| index 0000000..580c2f2 | |
| --- /dev/null | |
| +++ b/quant/quantize_stream.py | |
| +"""Stream a Ming MLLM directory to weight-only INT8 shards. | |
| + | |
| +Never builds the model: it buffers at most one output shard (<= 5 GB) of tensors at a time. Measured on | |
| +the real 34.0 GB checkpoint (AMD Strix Halo, 2026-09-23): 266 s wall, peak RSS 17.8 GiB. | |
| +CLI: ``python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]`` (matching modules stay BF16). | |
| +""" | |
| + | |
| +from __future__ import annotations | |
| + | |
| +import json | |
| +import math | |
| +import re | |
| +import os | |
| +import shutil | |
| +import sys | |
| +from dataclasses import dataclass | |
| +from pathlib import Path | |
| + | |
| +import torch | |
| +from safetensors import safe_open | |
| +from safetensors.torch import save_file | |
| + | |
| +try: # imported as the `quant` package | |
| + from .int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight | |
| +except ImportError: # run as a script: python quant/quantize_stream.py SRC DST | |
| + from int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight | |
| + | |
| +# Decimal GB, same unit Hugging Face uses for max_shard_size="5GB". | |
| +MAX_SHARD_BYTES = 5 * 10**9 | |
| + | |
| +_DTYPE_BYTES = { | |
| + "BOOL": 1, | |
| + "U8": 1, | |
| + "I8": 1, | |
| + "F8_E4M3": 1, | |
| + "F8_E5M2": 1, | |
| + "F8_E8M0": 1, | |
| + "U16": 2, | |
| + "I16": 2, | |
| + "F16": 2, | |
| + "BF16": 2, | |
| + "U32": 4, | |
| + "I32": 4, | |
| + "F32": 4, | |
| + "U64": 8, | |
| + "I64": 8, | |
| + "F64": 8, | |
| +} | |
| + | |
| +INDEX_NAME = "model.safetensors.index.json" | |
| +MANIFEST_NAME = "int8_manifest.json" | |
| + | |
| + | |
| +class QuantizeError(Exception): | |
| + """User-facing checkpoint error. main() prints it and returns 1.""" | |
| + | |
| + | |
| +def _die(msg: str) -> None: | |
| + raise QuantizeError(msg) | |
| + | |
| + | |
| +def _normalize_dtype(dtype_name) -> str: | |
| + text = str(dtype_name).upper() | |
| + if "." in text: | |
| + text = text.rsplit(".", 1)[-1] | |
| + aliases = { | |
| + "BFLOAT16": "BF16", | |
| + "FLOAT16": "F16", | |
| + "FLOAT32": "F32", | |
| + "FLOAT64": "F64", | |
| + "FLOAT8_E4M3FN": "F8_E4M3", | |
| + "FLOAT8_E5M2": "F8_E5M2", | |
| + "INT8": "I8", | |
| + "INT16": "I16", | |
| + "INT32": "I32", | |
| + "INT64": "I64", | |
| + "UINT8": "U8", | |
| + } | |
| + return aliases.get(text, text) | |
| + | |
| + | |
| +def _dtype_nbytes(dtype_name: str) -> int: | |
| + try: | |
| + return _DTYPE_BYTES[dtype_name] | |
| + except KeyError: | |
| + _die(f"unsupported safetensors dtype {dtype_name!r}") | |
| + raise # unreachable; satisfies type checkers | |
| + | |
| + | |
| +def _numel(shape: tuple[int, ...]) -> int: | |
| + n = 1 | |
| + for d in shape: | |
| + n *= int(d) | |
| + return n | |
| + | |
| + | |
| +def _load_index(path: Path) -> dict: | |
| + if not path.is_file(): | |
| + _die(f"missing index: {path}") | |
| + | |
| + def _pairs(pairs): | |
| + keys = [k for k, _ in pairs] | |
| + dupes = sorted({k for k in keys if keys.count(k) > 1}) | |
| + if dupes: | |
| + _die(f"duplicate key(s) in {path}: {dupes}") | |
| + return dict(pairs) | |
| + | |
| + try: | |
| + raw = path.read_text(encoding="utf-8") | |
| + index = json.loads(raw, object_pairs_hook=_pairs) | |
| + except QuantizeError: | |
| + raise | |
| + except (OSError, json.JSONDecodeError) as exc: | |
| + _die(f"cannot read index {path}: {exc}") | |
| + if not isinstance(index, dict) or not isinstance(index.get("weight_map"), dict): | |
| + _die(f"index {path} has no weight_map object") | |
| + if not index["weight_map"]: | |
| + _die(f"index {path} weight_map is empty") | |
| + return index | |
| + | |
| + | |
| +def _check_dst_clean(dst: Path) -> None: | |
| + if not dst.exists(): | |
| + return | |
| + if not dst.is_dir(): | |
| + _die(f"destination is not a directory: {dst}") | |
| + found = sorted(p.relative_to(dst).as_posix() for p in dst.rglob("*.safetensors")) | |
| + if found: | |
| + _die(f"destination already contains safetensors: {found}") | |
| + | |
| + | |
| +def _reject_nested(src: Path, dst: Path) -> None: | |
| + src_r = src.resolve() | |
| + dst_r = dst.resolve() | |
| + if src_r == dst_r or src_r in dst_r.parents or dst_r in src_r.parents: | |
| + _die(f"SRC and DST must be distinct and not nested: {src} vs {dst}") | |
| + | |
| + | |
| +def _shard_path(src: Path, shard_name: str) -> Path: | |
| + rel = Path(shard_name) | |
| + if rel.is_absolute() or ".." in rel.parts: | |
| + _die(f"unsafe shard path in index: {shard_name}") | |
| + path = src / rel | |
| + if not path.is_file(): | |
| + _die(f"index lists missing shard: {shard_name}") | |
| + return path | |
| + | |
| + | |
| +@dataclass | |
| +class Item: | |
| + src_shard: str | |
| + name: str | |
| + kind: str # "copy" or "quant" | |
| + shape: tuple[int, ...] | |
| + src_dtype: str | |
| + src_bytes: int | |
| + out_bytes: int | |
| + group: int = -1 | |
| + | |
| + | |
| +def _scale_name(weight_name: str) -> str: | |
| + return weight_name[: -len("weight")] + "scale" | |
| + | |
| + | |
| +def _plan(src: Path, index: dict, exclude: str | None = None) -> list[Item]: | |
| + """Metadata-only pass. Reads shapes and dtypes, not tensor bodies.""" | |
| + weight_map: dict[str, str] = index["weight_map"] | |
| + shard_order: list[str] = [] | |
| + seen_shards: set[str] = set() | |
| + for shard in weight_map.values(): | |
| + if shard not in seen_shards: | |
| + seen_shards.add(shard) | |
| + shard_order.append(shard) | |
| + | |
| + index_names_by_shard: dict[str, set[str]] = {s: set() for s in shard_order} | |
| + for name, shard in weight_map.items(): | |
| + if shard not in index_names_by_shard: | |
| + _die(f"weight_map value {shard!r} for {name} was not collected") | |
| + index_names_by_shard[shard].add(name) | |
| + | |
| + items: list[Item] = [] | |
| + seen_names: dict[str, str] = {} | |
| + for shard in shard_order: | |
| + path = _shard_path(src, shard) | |
| + with safe_open(str(path), framework="pt", device="cpu") as handle: | |
| + file_names = list(handle.keys()) | |
| + file_set = set(file_names) | |
| + if len(file_set) != len(file_names): | |
| + _die(f"shard {shard} header lists a tensor name twice") | |
| + missing = sorted(index_names_by_shard[shard] - file_set) | |
| + extra = sorted(file_set - index_names_by_shard[shard]) | |
| + if missing: | |
| + _die(f"index lists tensors missing from {shard}: {missing}") | |
| + if extra: | |
| + _die(f"{shard} contains tensors absent from the index: {extra}") | |
| + for name in file_names: | |
| + if name in seen_names: | |
| + _die( | |
| + f"tensor name appears twice: {name} " | |
| + f"({seen_names[name]} and {shard})" | |
| + ) | |
| + seen_names[name] = shard | |
| + sl = handle.get_slice(name) | |
| + if not hasattr(sl, "get_dtype") or not hasattr(sl, "get_shape"): | |
| + _die( | |
| + "safetensors safe_open slice is missing get_shape/get_dtype; " | |
| + "cannot plan shards without loading tensor bodies" | |
| + ) | |
| + shape = tuple(int(d) for d in sl.get_shape()) | |
| + dtype_name = _normalize_dtype(sl.get_dtype()) | |
| + src_bytes = _numel(shape) * _dtype_nbytes(dtype_name) | |
| + leaf = quant_rule_leaf(name) | |
| + if leaf is not None and exclude and re.search(exclude, name[: -len(".weight")]): | |
| + leaf = None # kept BF16 by --exclude | |
| + if leaf is not None and len(shape) != 2: | |
| + _die( | |
| + f"tensor {name} matches the quantize rule but is not 2-D " | |
| + f"(shape={list(shape)}, dtype={dtype_name})" | |
| + ) | |
| + if leaf is not None: | |
| + out_bytes = _numel(shape) * 1 + shape[0] * 4 # int8 weight + fp32 scale | |
| + items.append( | |
| + Item(shard, name, "quant", shape, dtype_name, src_bytes, out_bytes) | |
| + ) | |
| + else: | |
| + items.append( | |
| + Item(shard, name, "copy", shape, dtype_name, src_bytes, src_bytes) | |
| + ) | |
| + | |
| + index_names = set(weight_map) | |
| + planned = {it.name for it in items} | |
| + if planned != index_names: | |
| + _die( | |
| + "index / shard mismatch after scan: " | |
| + f"only_in_index={sorted(index_names - planned)[:8]} " | |
| + f"only_in_shards={sorted(planned - index_names)[:8]}" | |
| + ) | |
| + | |
| + produced = set(planned) | |
| + for it in items: | |
| + if it.kind != "quant": | |
| + continue | |
| + sname = _scale_name(it.name) | |
| + if sname in produced: | |
| + _die(f"scale name collides with an existing tensor: {sname}") | |
| + produced.add(sname) | |
| + return items | |
| + | |
| + | |
| +def _assign_groups(items: list[Item], max_shard_bytes: int) -> list[list[Item]]: | |
| + if max_shard_bytes <= 0: | |
| + _die(f"max_shard_bytes must be positive, got {max_shard_bytes}") | |
| + groups: list[list[Item]] = [] | |
| + cur: list[Item] = [] | |
| + cur_bytes = 0 | |
| + for it in items: | |
| + if cur and cur_bytes + it.out_bytes > max_shard_bytes: | |
| + groups.append(cur) | |
| + cur = [] | |
| + cur_bytes = 0 | |
| + if cur_bytes == 0 and it.out_bytes > max_shard_bytes: | |
| + print( | |
| + f"warning: {it.name} contributes {it.out_bytes} bytes, " | |
| + f"over the {max_shard_bytes}-byte shard target; writing it alone", | |
| + file=sys.stderr, | |
| + flush=True, | |
| + ) | |
| + it.group = len(groups) | |
| + cur.append(it) | |
| + cur_bytes += it.out_bytes | |
| + if cur: | |
| + groups.append(cur) | |
| + return groups | |
| + | |
| + | |
| +def _relative_frobenius(weight: torch.Tensor, q: torch.Tensor, scale: torch.Tensor) -> float: | |
| + w = weight.detach().to(dtype=torch.float64) | |
| + deq = q.detach().to(dtype=torch.float64) * scale.detach().to(dtype=torch.float64)[:, None] | |
| + denom = torch.linalg.matrix_norm(w, ord="fro") | |
| + numer = torch.linalg.matrix_norm(w - deq, ord="fro") | |
| + d = denom.item() | |
| + n = numer.item() | |
| + if d == 0.0: | |
| + return 0.0 if n == 0.0 else math.inf | |
| + return n / d | |
| + | |
| + | |
| +def _percentile_linear(values: list[float], pct: float) -> float: | |
| + """NumPy-style linear percentile. Empty → 0.""" | |
| + if not values: | |
| + return 0.0 | |
| + ordered = sorted(values) | |
| + if len(ordered) == 1: | |
| + return ordered[0] | |
| + rank = (len(ordered) - 1) * (pct / 100.0) | |
| + lo = math.floor(rank) | |
| + hi = math.ceil(rank) | |
| + if lo == hi: | |
| + return ordered[lo] | |
| + w = rank - lo | |
| + return ordered[lo] * (1.0 - w) + ordered[hi] * w | |
| + | |
| + | |
| +def _copy_sidecars(src: Path, dst: Path) -> list[str]: | |
| + copied: list[str] = [] | |
| + for dirpath, _dirnames, filenames in os.walk(src): | |
| + rel = Path(dirpath).relative_to(src) | |
| + out_dir = dst / rel | |
| + out_dir.mkdir(parents=True, exist_ok=True) | |
| + for filename in filenames: | |
| + if filename.endswith(".safetensors"): | |
| + continue | |
| + if filename == INDEX_NAME and rel == Path("."): | |
| + continue | |
| + src_file = Path(dirpath) / filename | |
| + dst_file = out_dir / filename | |
| + shutil.copy2(src_file, dst_file) | |
| + copied.append((rel / filename).as_posix()) | |
| + return copied | |
| + | |
| + | |
| +def _write_shards( | |
| + src: Path, | |
| + dst: Path, | |
| + items: list[Item], | |
| + groups: list[list[Item]], | |
| +) -> tuple[dict[str, str], int, int, list[tuple[str, float]], list[Path]]: | |
| + n_out = len(groups) | |
| + weight_map: dict[str, str] = {} | |
| + bytes_in = 0 | |
| + bytes_out = 0 | |
| + errors: list[tuple[str, float]] = [] | |
| + written: list[Path] = [] | |
| + | |
| + n_src = len({it.src_shard for it in items}) | |
| + src_seen = 0 | |
| + open_name: str | None = None | |
| + handle = None | |
| + buf: dict[str, torch.Tensor] = {} | |
| + buf_q = 0 | |
| + buf_c = 0 | |
| + current_group = 0 | |
| + | |
| + def flush() -> None: | |
| + nonlocal buf, buf_q, buf_c, current_group | |
| + if not buf: | |
| + return | |
| + fname = f"model-{current_group + 1:05d}-of-{n_out:05d}.safetensors" | |
| + path = dst / fname | |
| + for key, tensor in buf.items(): | |
| + if not tensor.is_contiguous(): | |
| + buf[key] = tensor.contiguous() | |
| + save_file(buf, str(path)) | |
| + shard_bytes = 0 | |
| + for key, tensor in buf.items(): | |
| + weight_map[key] = fname | |
| + shard_bytes += tensor.numel() * tensor.element_size() | |
| + written.append(path) | |
| + print( | |
| + f"wrote {fname}: tensors={len(buf)} quantized={buf_q} copied={buf_c} " | |
| + f"bytes={shard_bytes}", | |
| + flush=True, | |
| + ) | |
| + buf = {} | |
| + buf_q = 0 | |
| + buf_c = 0 | |
| + current_group += 1 | |
| + | |
| + try: | |
| + for it in items: | |
| + if it.src_shard != open_name: | |
| + if handle is not None: | |
| + handle.__exit__(None, None, None) | |
| + handle = None | |
| + path = _shard_path(src, it.src_shard) | |
| + handle = safe_open(str(path), framework="pt", device="cpu") | |
| + handle.__enter__() | |
| + open_name = it.src_shard | |
| + src_seen += 1 | |
| + n_here = sum(1 for x in items if x.src_shard == it.src_shard) | |
| + print( | |
| + f"reading source shard {src_seen}/{n_src} {it.src_shard} ({n_here} tensors)", | |
| + flush=True, | |
| + ) | |
| + assert handle is not None | |
| + tensor = handle.get_tensor(it.name) | |
| + got = tensor.numel() * tensor.element_size() | |
| + if got != it.src_bytes: | |
| + _die( | |
| + f"{it.name} byte size {got} != planned {it.src_bytes} " | |
| + f"(dtype={tensor.dtype}, shape={tuple(tensor.shape)})" | |
| + ) | |
| + bytes_in += got | |
| + if it.kind == "quant": | |
| + if not tensor.is_floating_point(): | |
| + _die( | |
| + f"{it.name} matches the quantize rule but dtype is {tensor.dtype}, " | |
| + "expected a floating dtype" | |
| + ) | |
| + if tuple(tensor.shape) != it.shape: | |
| + _die(f"{it.name} shape changed between passes: {tuple(tensor.shape)} vs {it.shape}") | |
| + q, scale = quantize_weight(tensor) | |
| + err = _relative_frobenius(tensor, q, scale) | |
| + if math.isnan(err) or math.isinf(err): | |
| + _die(f"non-finite relative error for {it.name}: {err}") | |
| + errors.append((it.name, err)) | |
| + del tensor | |
| + sname = _scale_name(it.name) | |
| + buf[it.name] = q | |
| + buf[sname] = scale | |
| + produced = q.numel() * q.element_size() + scale.numel() * scale.element_size() | |
| + if produced != it.out_bytes: | |
| + _die(f"{it.name} output bytes {produced} != planned {it.out_bytes}") | |
| + buf_q += 1 | |
| + else: | |
| + if not tensor.is_contiguous(): | |
| + tensor = tensor.contiguous() | |
| + buf[it.name] = tensor | |
| + buf_c += 1 | |
| + bytes_out += it.out_bytes | |
| + # Flush when this item closes its planned output shard. | |
| + group_items = groups[it.group] | |
| + if it is group_items[-1]: | |
| + flush() | |
| + finally: | |
| + if handle is not None: | |
| + handle.__exit__(None, None, None) | |
| + | |
| + if buf: | |
| + _die("internal error: output buffer not flushed") | |
| + if current_group != n_out: | |
| + _die(f"internal error: wrote {current_group} shards, planned {n_out}") | |
| + return weight_map, bytes_in, bytes_out, errors, written | |
| + | |
| + | |
| +def _summary( | |
| + errors: list[tuple[str, float]], | |
| + n_quant: int, | |
| + n_copy: int, | |
| + bytes_in: int, | |
| + bytes_out: int, | |
| +) -> dict: | |
| + vals = [e for _, e in errors] | |
| + if errors: | |
| + worst_name, worst_err = min( | |
| + errors, | |
| + key=lambda pair: (-pair[1], pair[0]), | |
| + ) | |
| + else: | |
| + worst_name, worst_err = None, 0.0 | |
| + mean = (sum(vals) / len(vals)) if vals else 0.0 | |
| + return { | |
| + "tensors_quantized": n_quant, | |
| + "tensors_copied": n_copy, | |
| + "bytes_in": bytes_in, | |
| + "bytes_out": bytes_out, | |
| + "mean_relative_error": mean, | |
| + "p99_relative_error": _percentile_linear(vals, 99.0), | |
| + "max_relative_error": worst_err if vals else 0.0, | |
| + "worst_tensor": worst_name, | |
| + } | |
| + | |
| + | |
| +def run(src: Path, dst: Path, max_shard_bytes: int = MAX_SHARD_BYTES, exclude: str | None = None) -> dict: | |
| + src = src.resolve() | |
| + dst = dst.resolve() | |
| + if not src.is_dir(): | |
| + _die(f"SRC is not a directory: {src}") | |
| + _reject_nested(src, dst) | |
| + _check_dst_clean(dst) | |
| + index = _load_index(src / INDEX_NAME) | |
| + items = _plan(src, index, exclude) | |
| + groups = _assign_groups(items, max_shard_bytes) | |
| + dst.mkdir(parents=True, exist_ok=True) | |
| + | |
| + written: list[Path] = [] | |
| + try: | |
| + weight_map, bytes_in, bytes_out, errors, written = _write_shards(src, dst, items, groups) | |
| + copied = _copy_sidecars(src, dst) | |
| + n_quant = sum(1 for it in items if it.kind == "quant") | |
| + n_copy = sum(1 for it in items if it.kind == "copy") | |
| + measured = _summary(errors, n_quant, n_copy, bytes_in, bytes_out) | |
| + if measured["bytes_in"] != bytes_in or measured["bytes_out"] != bytes_out: | |
| + _die("internal error: summary byte counters diverged") | |
| + # Recompute the on-disk total from the tensors we recorded. weight_map | |
| + # values are what we just saved; bytes_out is that sum. | |
| + out_index = {"metadata": {"total_size": bytes_out}, "weight_map": weight_map} | |
| + (dst / INDEX_NAME).write_text( | |
| + json.dumps(out_index, indent=2) + "\n", encoding="utf-8" | |
| + ) | |
| + modules = sorted( | |
| + it.name[: -len(".weight")] for it in items if it.kind == "quant" | |
| + ) | |
| + manifest = { | |
| + "format": "ming-int8-wo-v1", | |
| + "scheme": "weight-only int8, per-output-channel symmetric, fp32 scales", | |
| + "rule": QUANT_RULE + (f" Additionally kept BF16: modules matching /{exclude}/." if exclude else ""), | |
| + "exclude": exclude, | |
| + "quantized_modules": modules, | |
| + "source_total_size": bytes_in, | |
| + "total_size": bytes_out, | |
| + "measured": measured, | |
| + } | |
| + (dst / MANIFEST_NAME).write_text( | |
| + json.dumps(manifest, indent=2, allow_nan=False) + "\n", encoding="utf-8" | |
| + ) | |
| + except Exception: | |
| + for path in written: | |
| + try: | |
| + path.unlink() | |
| + except OSError: | |
| + pass | |
| + raise | |
| + | |
| + print(f"copied {len(copied)} non-safetensors file(s)", flush=True) | |
| + m = measured | |
| + print( | |
| + "summary: " | |
| + f"quantized={m['tensors_quantized']} copied={m['tensors_copied']} " | |
| + f"bytes_in={m['bytes_in']} bytes_out={m['bytes_out']} " | |
| + f"mean_rel={m['mean_relative_error']:.8g} " | |
| + f"p99_rel={m['p99_relative_error']:.8g} " | |
| + f"max_rel={m['max_relative_error']:.8g} " | |
| + f"worst={m['worst_tensor']}", | |
| + flush=True, | |
| + ) | |
| + return manifest | |
| + | |
| + | |
| +def main(argv: list[str] | None = None) -> int: | |
| + args = list(sys.argv[1:] if argv is None else argv) | |
| + exclude = None | |
| + if "--exclude" in args: | |
| + i = args.index("--exclude") | |
| + if i + 1 >= len(args): | |
| + print("--exclude needs a regex", file=sys.stderr) | |
| + return 2 | |
| + exclude = args[i + 1] | |
| + re.compile(exclude) | |
| + del args[i : i + 2] | |
| + if len(args) != 2: | |
| + print( | |
| + "usage: python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]", | |
| + file=sys.stderr, | |
| + ) | |
| + return 2 | |
| + try: | |
| + run(Path(args[0]), Path(args[1]), max_shard_bytes=MAX_SHARD_BYTES, exclude=exclude) | |
| + except QuantizeError as exc: | |
| + print(f"error: {exc}", file=sys.stderr) | |
| + return 1 | |
| + return 0 | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + sys.exit(main()) | |
| diff --git a/quant/test_int8.py b/quant/test_int8.py | |
| new file mode 100644 | |
| index 0000000..03baa56 | |
| --- /dev/null | |
| +++ b/quant/test_int8.py | |
| +"""CPU tests for weight-only INT8 Ming MLLM quantize + load. | |
| + | |
| +Run: HIP_VISIBLE_DEVICES=-1 python test_int8.py | |
| +""" | |
| + | |
| +from __future__ import annotations | |
| + | |
| +import json | |
| +import sys | |
| +import tempfile | |
| +import traceback | |
| +from pathlib import Path | |
| + | |
| +import torch | |
| +import torch.nn.functional as F | |
| +from safetensors.torch import load_file, save_file | |
| +from torch import nn | |
| + | |
| +import quantize_stream | |
| +from int8_linear import Int8Linear, is_quantizable, quantize_weight | |
| +from load_int8 import load_int8_mllm_ | |
| + | |
| +# Tiny stand-in for Ming's MLLM names. Not the real model. | |
| +HIDDEN = 32 | |
| +INTER = 48 | |
| +VOCAB = 64 | |
| +N_EXPERTS = 2 | |
| + | |
| + | |
| +class RMSNorm(nn.Module): | |
| + def __init__(self, dim: int, eps: float = 1e-6): | |
| + super().__init__() | |
| + self.weight = nn.Parameter(torch.ones(dim)) | |
| + self.eps = eps | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + var = x.float().pow(2).mean(dim=-1, keepdim=True) | |
| + y = x * torch.rsqrt(var + self.eps) | |
| + return (y * self.weight).to(dtype=x.dtype) | |
| + | |
| + | |
| +class Attention(nn.Module): | |
| + def __init__(self, hidden: int): | |
| + super().__init__() | |
| + self.hidden = hidden | |
| + self.query_key_value = nn.Linear(hidden, hidden * 3, bias=True) | |
| + self.dense = nn.Linear(hidden, hidden, bias=False) | |
| + self.q_norm = RMSNorm(hidden) | |
| + self.k_norm = RMSNorm(hidden) | |
| + # Non-persistent, like BailingMoeV2RotaryEmbedding.inv_freq. | |
| + self.register_buffer( | |
| + "inv_freq", torch.arange(hidden // 2, dtype=torch.float32), persistent=False | |
| + ) | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + qkv = self.query_key_value(x) | |
| + h = self.hidden | |
| + q = self.q_norm(qkv[..., :h]) | |
| + k = self.k_norm(qkv[..., h : 2 * h]) | |
| + v = qkv[..., 2 * h :] | |
| + return self.dense(q + k + v) | |
| + | |
| + | |
| +class DenseMLP(nn.Module): | |
| + def __init__(self, hidden: int, inter: int): | |
| + super().__init__() | |
| + self.gate_proj = nn.Linear(hidden, inter, bias=False) | |
| + self.up_proj = nn.Linear(hidden, inter, bias=True) | |
| + self.down_proj = nn.Linear(inter, hidden, bias=False) | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) | |
| + | |
| + | |
| +class Expert(nn.Module): | |
| + def __init__(self, hidden: int, inter: int): | |
| + super().__init__() | |
| + self.gate_proj = nn.Linear(hidden, inter, bias=False) | |
| + self.up_proj = nn.Linear(hidden, inter, bias=True) | |
| + self.down_proj = nn.Linear(inter, hidden, bias=False) | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) | |
| + | |
| + | |
| +class Router(nn.Module): | |
| + """Not an nn.Linear. Leaf name is gate / image_gate / audio_gate.""" | |
| + | |
| + def __init__(self, hidden: int, n_experts: int): | |
| + super().__init__() | |
| + self.weight = nn.Parameter(torch.empty(n_experts, hidden)) | |
| + self.expert_bias = nn.Parameter(torch.zeros(n_experts), requires_grad=False) | |
| + nn.init.kaiming_uniform_(self.weight, a=5**0.5) | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + return F.linear(x, self.weight, self.expert_bias) | |
| + | |
| + | |
| +class MoeMLP(nn.Module): | |
| + def __init__(self, hidden: int, inter: int, n_experts: int): | |
| + super().__init__() | |
| + self.gate = Router(hidden, n_experts) | |
| + self.image_gate = Router(hidden, n_experts) | |
| + self.audio_gate = Router(hidden, n_experts) | |
| + self.experts = nn.ModuleList(Expert(hidden, inter) for _ in range(n_experts)) | |
| + self.shared_experts = Expert(hidden, inter) | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + scores = self.gate(x) + self.image_gate(x) + self.audio_gate(x) | |
| + weights = torch.softmax(scores, dim=-1) | |
| + mixed = self.shared_experts(x) | |
| + for i, expert in enumerate(self.experts): | |
| + mixed = mixed + expert(x) * weights[..., i : i + 1] | |
| + return mixed | |
| + | |
| + | |
| +class DecoderLayer(nn.Module): | |
| + def __init__(self, hidden: int, mlp: nn.Module): | |
| + super().__init__() | |
| + self.input_layernorm = RMSNorm(hidden) | |
| + self.post_attention_layernorm = RMSNorm(hidden) | |
| + self.attention = Attention(hidden) | |
| + self.mlp = mlp | |
| + | |
| + def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| + x = x + self.attention(self.input_layernorm(x)) | |
| + x = x + self.mlp(self.post_attention_layernorm(x)) | |
| + return x | |
| + | |
| + | |
| +class TinyMing(nn.Module): | |
| + """Names match the real checkpoint: model.model.layers.*, model.lm_head, vision.*.""" | |
| + | |
| + def __init__(self): | |
| + super().__init__() | |
| + self.model = nn.Module() | |
| + self.model.model = nn.Module() | |
| + self.model.model.word_embeddings = nn.Embedding(VOCAB, HIDDEN) | |
| + self.model.model.layers = nn.ModuleList( | |
| + [ | |
| + DecoderLayer(HIDDEN, DenseMLP(HIDDEN, INTER)), | |
| + DecoderLayer(HIDDEN, MoeMLP(HIDDEN, INTER, N_EXPERTS)), | |
| + ] | |
| + ) | |
| + self.model.model.norm = RMSNorm(HIDDEN) | |
| + self.model.lm_head = nn.Linear(HIDDEN, VOCAB, bias=False) | |
| + block = nn.Module() | |
| + block.attn = nn.Module() | |
| + block.attn.qkv = nn.Linear(HIDDEN, HIDDEN, bias=False) | |
| + self.vision = nn.Module() | |
| + self.vision.blocks = nn.ModuleList([block]) | |
| + self.linear_proj = nn.ModuleList([nn.Linear(HIDDEN, HIDDEN, bias=True)]) | |
| + | |
| + def forward(self, input_ids: torch.Tensor) -> torch.Tensor: | |
| + h = self.model.model.word_embeddings(input_ids) | |
| + for layer in self.model.model.layers: | |
| + h = layer(h) | |
| + h = self.model.model.norm(h) | |
| + return self.model.lm_head(h) | |
| + | |
| + | |
| +# Modules the rule must select for TinyMing. Hardcoded — not derived from is_quantizable. | |
| +EXPECTED_QUANT_MODULES = [ | |
| + "model.model.layers.0.attention.dense", | |
| + "model.model.layers.0.attention.query_key_value", | |
| + "model.model.layers.0.mlp.down_proj", | |
| + "model.model.layers.0.mlp.gate_proj", | |
| + "model.model.layers.0.mlp.up_proj", | |
| + "model.model.layers.1.attention.dense", | |
| + "model.model.layers.1.attention.query_key_value", | |
| + "model.model.layers.1.mlp.experts.0.down_proj", | |
| + "model.model.layers.1.mlp.experts.0.gate_proj", | |
| + "model.model.layers.1.mlp.experts.0.up_proj", | |
| + "model.model.layers.1.mlp.experts.1.down_proj", | |
| + "model.model.layers.1.mlp.experts.1.gate_proj", | |
| + "model.model.layers.1.mlp.experts.1.up_proj", | |
| + "model.model.layers.1.mlp.shared_experts.down_proj", | |
| + "model.model.layers.1.mlp.shared_experts.gate_proj", | |
| + "model.model.layers.1.mlp.shared_experts.up_proj", | |
| +] | |
| + | |
| +MUST_NOT_QUANTIZE = [ | |
| + "model.model.layers.1.mlp.gate", | |
| + "model.model.layers.1.mlp.image_gate", | |
| + "model.model.layers.1.mlp.audio_gate", | |
| + "model.model.word_embeddings", | |
| + "model.model.norm", | |
| + "model.lm_head", | |
| + "vision.blocks.0.attn.qkv", | |
| + "linear_proj.0", | |
| + "model.model.layers.0.attention.q_norm", | |
| + "model.model.layers.0.input_layernorm", | |
| +] | |
| + | |
| + | |
| +def _move_parameters_to_meta(model: nn.Module) -> nn.Module: | |
| + """Parameters → meta, buffers stay where they are (CPU). Matches accelerate include_buffers=False.""" | |
| + for mod in model.modules(): | |
| + for name, param in list(mod._parameters.items()): | |
| + if param is None: | |
| + continue | |
| + mod._parameters[name] = nn.Parameter( | |
| + param.detach().to(device="meta"), | |
| + requires_grad=param.requires_grad, | |
| + ) | |
| + return model | |
| + | |
| + | |
| +def _save_bf16_checkpoint(model: nn.Module, src: Path) -> None: | |
| + src.mkdir(parents=True, exist_ok=True) | |
| + sd = {k: v.detach().contiguous() for k, v in model.state_dict().items()} | |
| + if not sd: | |
| + raise AssertionError("empty state_dict") | |
| + for tensor in sd.values(): | |
| + if tensor.is_floating_point(): | |
| + assert tensor.dtype == torch.bfloat16, tensor.dtype | |
| + keys = list(sd) | |
| + mid = max(1, len(keys) // 2) | |
| + shards = { | |
| + "bf16-00001.safetensors": {k: sd[k] for k in keys[:mid]}, | |
| + "bf16-00002.safetensors": {k: sd[k] for k in keys[mid:]}, | |
| + } | |
| + weight_map = {} | |
| + total = 0 | |
| + for filename, tensors in shards.items(): | |
| + save_file(tensors, str(src / filename)) | |
| + for name, tensor in tensors.items(): | |
| + weight_map[name] = filename | |
| + total += tensor.numel() * tensor.element_size() | |
| + index = {"metadata": {"total_size": total}, "weight_map": weight_map} | |
| + (src / "model.safetensors.index.json").write_text( | |
| + json.dumps(index, indent=2) + "\n", encoding="utf-8" | |
| + ) | |
| + (src / "config.json").write_bytes(b'{"model_type":"tiny-ming","hidden":32}\n') | |
| + extra = src / "extra" | |
| + extra.mkdir() | |
| + (extra / "chat_template.jinja").write_text("{{ messages }}\n", encoding="utf-8") | |
| + | |
| + | |
| +def _load_all(folder: Path) -> dict[str, torch.Tensor]: | |
| + index = json.loads((folder / "model.safetensors.index.json").read_text(encoding="utf-8")) | |
| + order: list[str] = [] | |
| + seen: set[str] = set() | |
| + for shard in index["weight_map"].values(): | |
| + if shard not in seen: | |
| + seen.add(shard) | |
| + order.append(shard) | |
| + sd: dict[str, torch.Tensor] = {} | |
| + for shard in order: | |
| + sd.update(load_file(str(folder / shard))) | |
| + return sd | |
| + | |
| + | |
| +def _apply_int8_(model: nn.Module) -> None: | |
| + names = [] | |
| + for name, mod in model.named_modules(): | |
| + if isinstance(mod, nn.Linear) and is_quantizable( | |
| + f"{name}.weight", tuple(mod.weight.shape) | |
| + ): | |
| + names.append(name) | |
| + for name in names: | |
| + parent_name, _, leaf = name.rpartition(".") | |
| + parent = model.get_submodule(parent_name) if parent_name else model | |
| + setattr(parent, leaf, Int8Linear.from_linear(getattr(parent, leaf))) | |
| + | |
| + | |
| +def _assert_no_meta(model: nn.Module) -> None: | |
| + for name, param in model.named_parameters(): | |
| + assert param.device.type != "meta", name | |
| + for mod_name, mod in model.named_modules(): | |
| + 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 | |
| + assert buf.device.type != "meta", full | |
| + | |
| + | |
| +def test_from_linear_roundtrip() -> None: | |
| + torch.manual_seed(0) | |
| + out_f, in_f = 5, 7 | |
| + lin = nn.Linear(in_f, out_f, bias=True) | |
| + scales = torch.tensor([0.5, 0.25, 0.125, 2.0, 4.0], dtype=torch.float32) | |
| + q = torch.randint(-127, 128, (out_f, in_f), dtype=torch.int8) | |
| + q[:, 0] = 127 | |
| + q[2, :] = 0 # all-zero row; must not NaN | |
| + weight = q.float() * scales[:, None] | |
| + with torch.no_grad(): | |
| + lin.weight.copy_(weight) | |
| + lin.bias.copy_(torch.tensor([0.1, -0.2, 0.3, -0.4, 0.5])) | |
| + mod = Int8Linear.from_linear(lin) | |
| + deq = mod.weight.float() * mod.scale[:, None] | |
| + for row in range(out_f): | |
| + if row == 2: | |
| + assert torch.equal(mod.weight[row], torch.zeros(in_f, dtype=torch.int8)) | |
| + assert float(mod.scale[row]) == 1.0 | |
| + assert torch.equal(deq[row], torch.zeros(in_f)) | |
| + else: | |
| + assert torch.equal(deq[row], weight[row]), (deq[row] - weight[row]).abs().max().item() | |
| + assert mod.bias is not None and torch.equal(mod.bias, lin.bias) | |
| + assert mod.bias.dtype == lin.bias.dtype | |
| + assert torch.isfinite(mod.scale).all() | |
| + | |
| + # Random weights: per-element error stays within half a bin (+ float slack). | |
| + lin_r = nn.Linear(13, 9, bias=False) | |
| + mod_r = Int8Linear.from_linear(lin_r) | |
| + w = lin_r.weight.detach().float() | |
| + deq_r = (mod_r.weight.double() * mod_r.scale.double()[:, None]).float() | |
| + err = (w.double() - deq_r.double()).abs() | |
| + half = mod_r.scale.double()[:, None] * 0.5 | |
| + slip = (err - half).max().item() | |
| + assert slip <= 1e-4, slip | |
| + assert torch.isfinite(mod_r.scale).all() | |
| + | |
| + # Entirely zero weight: finite forward, zero codes, scale 1. | |
| + lin_z = nn.Linear(4, 3, bias=True) | |
| + with torch.no_grad(): | |
| + lin_z.weight.zero_() | |
| + mod_z = Int8Linear.from_linear(lin_z) | |
| + assert torch.equal(mod_z.weight, torch.zeros_like(mod_z.weight)) | |
| + assert torch.equal(mod_z.scale, torch.ones(3)) | |
| + y = mod_z(torch.randn(8, 4)) | |
| + assert torch.isfinite(y).all() | |
| + assert torch.allclose(y, mod_z.bias.expand_as(y)) | |
| + | |
| + # Zero row contributes only its bias. | |
| + x = torch.randn(6, in_f) | |
| + y_mix = mod(x) | |
| + assert torch.isfinite(y_mix).all() | |
| + assert torch.allclose(y_mix[:, 2], mod.bias[2].expand(6)) | |
| + | |
| + # bf16 source linear: codes int8, scale fp32, bias stays bf16. | |
| + lin_b = nn.Linear(8, 4, bias=True).to(dtype=torch.bfloat16) | |
| + mod_b = Int8Linear.from_linear(lin_b) | |
| + assert mod_b.weight.dtype == torch.int8 | |
| + assert mod_b.scale.dtype == torch.float32 | |
| + assert mod_b.bias is not None and mod_b.bias.dtype == torch.bfloat16 | |
| + w_b = lin_b.weight.detach().float() | |
| + deq_b = mod_b.weight.float() * mod_b.scale[:, None] | |
| + err_b = (w_b.double() - deq_b.double()).abs() | |
| + half_b = mod_b.scale.double()[:, None] * 0.5 | |
| + assert (err_b - half_b).max().item() <= 1e-2, (err_b - half_b).max().item() | |
| + | |
| + | |
| +def _assert_quant_dtypes(mod: Int8Linear, scale: torch.Tensor, weight: torch.Tensor, bias_dtype: torch.dtype) -> None: | |
| + assert mod.weight.dtype == torch.int8 | |
| + assert mod.scale.dtype == torch.float32 | |
| + assert torch.equal(mod.weight, weight) | |
| + assert torch.equal(mod.scale, scale) | |
| + assert mod.bias is not None and mod.bias.dtype == bias_dtype | |
| + | |
| + | |
| +def test_dtype_cast_keeps_scale_fp32() -> None: | |
| + torch.manual_seed(1) | |
| + lin = nn.Linear(5, 3, bias=True) | |
| + fresh = Int8Linear.from_linear(lin) | |
| + scale = fresh.scale.detach().clone() | |
| + weight = fresh.weight.detach().clone() | |
| + bias = fresh.bias.detach().clone() | |
| + assert scale.dtype == torch.float32 and weight.dtype == torch.int8 and bias.dtype == torch.float32 | |
| + | |
| + # Each cast starts from fp32 so "bias follows the cast" is the single cast of the source bias. | |
| + mod = Int8Linear.from_linear(lin) | |
| + mod.bfloat16() | |
| + _assert_quant_dtypes(mod, scale, weight, torch.bfloat16) | |
| + assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16)) | |
| + | |
| + mod = Int8Linear.from_linear(lin) | |
| + mod.half() | |
| + _assert_quant_dtypes(mod, scale, weight, torch.float16) | |
| + assert torch.equal(mod.bias, bias.to(dtype=torch.float16)) | |
| + | |
| + mod = Int8Linear.from_linear(lin) | |
| + mod.to(torch.bfloat16) | |
| + _assert_quant_dtypes(mod, scale, weight, torch.bfloat16) | |
| + assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16)) | |
| + | |
| + mod = Int8Linear.from_linear(lin) | |
| + mod.to(dtype=torch.float16) | |
| + _assert_quant_dtypes(mod, scale, weight, torch.float16) | |
| + assert torch.equal(mod.bias, bias.to(dtype=torch.float16)) | |
| + | |
| + # A second cast applies to the bias's current dtype, not the original fp32 value. | |
| + mod = Int8Linear.from_linear(lin) | |
| + mod.to(torch.bfloat16) | |
| + mod.to(dtype=torch.float16) | |
| + _assert_quant_dtypes(mod, scale, weight, torch.float16) | |
| + assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16).to(dtype=torch.float16)) | |
| + | |
| + # What the caller actually does: parent.to(device=..., dtype=bf16). | |
| + parent = nn.Sequential(Int8Linear.from_linear(lin)) | |
| + parent.to(device="cpu", dtype=torch.bfloat16) | |
| + _assert_quant_dtypes(parent[0], scale, weight, torch.bfloat16) | |
| + assert torch.equal(parent[0].bias, bias.to(dtype=torch.bfloat16)) | |
| + | |
| + shell = Int8Linear.shell(4, 3, bias=True, bias_dtype=torch.bfloat16, device="meta") | |
| + assert shell.weight.dtype == torch.int8 and shell.weight.device.type == "meta" | |
| + assert shell.scale.dtype == torch.float32 and shell.scale.device.type == "meta" | |
| + assert shell.bias is not None | |
| + assert shell.bias.dtype == torch.bfloat16 and shell.bias.device.type == "meta" | |
| + shell_nb = Int8Linear.shell(4, 3, bias=False, bias_dtype=torch.float32, device="meta") | |
| + assert shell_nb.bias is None | |
| + | |
| + | |
| +def test_forward_matches_reference() -> None: | |
| + torch.manual_seed(2) | |
| + for bias in (True, False): | |
| + lin = nn.Linear(6, 4, bias=bias) | |
| + # Bias is passed through unchanged, so it has to already match x's dtype | |
| + # (the caller does model.to(dtype=...) before the prefill). | |
| + modules = [ | |
| + (Int8Linear.from_linear(lin), torch.float32), | |
| + (Int8Linear.from_linear(lin).to(torch.bfloat16), torch.bfloat16), | |
| + (Int8Linear.from_linear(lin).to(dtype=torch.float16), torch.float16), | |
| + ] | |
| + for mod, dtype in modules: | |
| + if mod.bias is not None: | |
| + assert mod.bias.dtype == dtype | |
| + x = torch.randn(3, 5, 6, dtype=dtype) | |
| + ref_w = (mod.weight.float() * mod.scale[:, None]).to(dtype=x.dtype) | |
| + y = mod(x) | |
| + y_ref = F.linear(x, ref_w, mod.bias) | |
| + assert torch.equal(y, y_ref), (bias, dtype) | |
| + | |
| + | |
| +def test_is_quantizable_rule() -> None: | |
| + false_cases = [ | |
| + ("model.model.layers.1.mlp.gate.weight", (256, 2048)), | |
| + ("model.model.layers.1.mlp.image_gate.weight", (256, 2048)), | |
| + ("model.model.layers.1.mlp.audio_gate.weight", (256, 2048)), | |
| + ("model.model.layers.1.mlp.gate.expert_bias", (256,)), | |
| + ("model.lm_head.weight", (151936, 2048)), | |
| + ("model.model.word_embeddings.weight", (151936, 2048)), | |
| + ("vision.blocks.0.attn.qkv.weight", (3072, 1280)), | |
| + ("model.model.layers.0.input_layernorm.weight", (2048,)), | |
| + ("model.model.layers.0.post_attention_layernorm.weight", (2048,)), | |
| + ("model.model.layers.0.attention.q_norm.weight", (128,)), | |
| + ("model.model.layers.0.attention.k_norm.weight", (128,)), | |
| + ("model.model.norm.weight", (2048,)), | |
| + ("linear_proj.0.weight", (2048, 2048)), | |
| + ("model.model.layers.0.attention.query_key_value.bias", (3072,)), | |
| + ("model.model.layers.0.mlp.experts.0.gate_proj.bias", (512,)), | |
| + # Right leaf, wrong rank: not quantizable (the stream must reject it). | |
| + ("model.model.layers.0.attention.query_key_value.weight", (3072,)), | |
| + ("model.model.layers.0.mlp.gate_proj.weight", (1024, 2048, 1)), | |
| + ] | |
| + true_cases = [ | |
| + ("model.model.layers.3.mlp.experts.3.gate_proj.weight", (512, 2048)), | |
| + ("model.model.layers.3.mlp.shared_experts.down_proj.weight", (2048, 512)), | |
| + ("model.model.layers.0.mlp.up_proj.weight", (512, 2048)), | |
| + ("layers.0.mlp.up_proj.weight", (512, 2048)), | |
| + ("model.model.layers.0.attention.query_key_value.weight", (3072, 2048)), | |
| + ("model.model.layers.0.attention.dense.weight", (2048, 2048)), | |
| + ("model.model.layers.0.mlp.gate_proj.weight", (512, 2048)), | |
| + ("model.model.layers.0.mlp.down_proj.weight", (2048, 512)), | |
| + ("model.model.layers.19.mlp.experts.255.up_proj.weight", (512, 2048)), | |
| + ] | |
| + for name, shape in false_cases: | |
| + assert is_quantizable(name, shape) is False, name | |
| + for name, shape in true_cases: | |
| + assert is_quantizable(name, shape) is True, name | |
| + | |
| + | |
| +def _shard_groups(sd: dict[str, torch.Tensor]) -> set[str]: | |
| + """One copy-tensor, or one weight+scale pair, is one unsplittable group.""" | |
| + names = set(sd) | |
| + groups: set[str] = set() | |
| + for name in names: | |
| + if name.endswith(".scale") and name[: -len(".scale")] + ".weight" in names: | |
| + groups.add(name[: -len(".scale")]) | |
| + elif name.endswith(".weight") and name[: -len(".weight")] + ".scale" in names: | |
| + groups.add(name[: -len(".weight")]) | |
| + else: | |
| + groups.add(name) | |
| + return groups | |
| + | |
| + | |
| +def test_end_to_end_stream_and_load() -> None: | |
| + assert quantize_stream.MAX_SHARD_BYTES == 5 * 10**9 | |
| + torch.manual_seed(3) | |
| + src_model = TinyMing().to(dtype=torch.bfloat16) | |
| + # Non-persistent rotary buffer is not part of the checkpoint. | |
| + assert "model.model.layers.0.attention.inv_freq" not in src_model.state_dict() | |
| + | |
| + with tempfile.TemporaryDirectory(prefix="ming-int8-") as tmp: | |
| + root = Path(tmp) | |
| + src = root / "src" | |
| + dst = root / "dst" | |
| + _save_bf16_checkpoint(src_model, src) | |
| + limit = 2048 | |
| + old = quantize_stream.MAX_SHARD_BYTES | |
| + quantize_stream.MAX_SHARD_BYTES = limit | |
| + try: | |
| + rc = quantize_stream.main([str(src), str(dst)]) | |
| + finally: | |
| + quantize_stream.MAX_SHARD_BYTES = old | |
| + assert rc == 0, rc | |
| + assert quantize_stream.MAX_SHARD_BYTES == 5 * 10**9 | |
| + | |
| + # Sidecars copied verbatim; original index replaced. | |
| + assert (dst / "config.json").read_bytes() == (src / "config.json").read_bytes() | |
| + assert (dst / "extra" / "chat_template.jinja").read_bytes() == ( | |
| + src / "extra" / "chat_template.jinja" | |
| + ).read_bytes() | |
| + assert not (dst / "bf16-00001.safetensors").exists() | |
| + | |
| + manifest = json.loads((dst / "int8_manifest.json").read_text(encoding="utf-8")) | |
| + assert manifest["format"] == "ming-int8-wo-v1" | |
| + assert manifest["scheme"] == ( | |
| + "weight-only int8, per-output-channel symmetric, fp32 scales" | |
| + ) | |
| + assert manifest["quantized_modules"] == sorted(EXPECTED_QUANT_MODULES) | |
| + for banned in MUST_NOT_QUANTIZE: | |
| + assert banned not in manifest["quantized_modules"], banned | |
| + | |
| + index = json.loads((dst / "model.safetensors.index.json").read_text(encoding="utf-8")) | |
| + assert index["metadata"]["total_size"] == manifest["total_size"] | |
| + measured = manifest["measured"] | |
| + assert measured["tensors_quantized"] == len(EXPECTED_QUANT_MODULES) | |
| + assert measured["bytes_in"] == manifest["source_total_size"] | |
| + assert measured["bytes_out"] == manifest["total_size"] | |
| + assert measured["bytes_out"] < measured["bytes_in"] | |
| + n_out_keys = measured["tensors_copied"] + 2 * measured["tensors_quantized"] | |
| + assert len(index["weight_map"]) == n_out_keys | |
| + | |
| + src_sd = _load_all(src) | |
| + dst_sd = _load_all(dst) | |
| + assert manifest["source_total_size"] == sum( | |
| + t.numel() * t.element_size() for t in src_sd.values() | |
| + ) | |
| + assert manifest["total_size"] == sum(t.numel() * t.element_size() for t in dst_sd.values()) | |
| + | |
| + shard_names = sorted({*index["weight_map"].values()}) | |
| + assert len(shard_names) >= 2, shard_names | |
| + for shard in shard_names: | |
| + shard_sd = load_file(str(dst / shard)) | |
| + total = sum(t.numel() * t.element_size() for t in shard_sd.values()) | |
| + if total > limit: | |
| + assert len(_shard_groups(shard_sd)) == 1, (shard, total, list(shard_sd)) | |
| + | |
| + errors = [] | |
| + for name, src_t in src_sd.items(): | |
| + if is_quantizable(name, tuple(src_t.shape)): | |
| + q = dst_sd[name] | |
| + scale_key = name[: -len("weight")] + "scale" | |
| + scale = dst_sd[scale_key] | |
| + assert q.dtype == torch.int8, name | |
| + assert scale.dtype == torch.float32, scale_key | |
| + q_ref, scale_ref = quantize_weight(src_t) | |
| + assert torch.equal(q, q_ref), name | |
| + assert torch.equal(scale, scale_ref), scale_key | |
| + errors.append((name, quantize_stream._relative_frobenius(src_t, q, scale))) | |
| + else: | |
| + assert name in dst_sd, name | |
| + assert dst_sd[name].dtype == src_t.dtype, (name, dst_sd[name].dtype, src_t.dtype) | |
| + assert torch.equal(dst_sd[name], src_t), name | |
| + # Router weights stayed BF16 and byte-identical (the gate vs gate_proj trap). | |
| + router = "model.model.layers.1.mlp.gate.weight" | |
| + assert dst_sd[router].dtype == torch.bfloat16 | |
| + assert torch.equal(dst_sd[router], src_sd[router]) | |
| + for suffix in ("image_gate.weight", "audio_gate.weight", "gate.expert_bias"): | |
| + key = f"model.model.layers.1.mlp.{suffix}" | |
| + assert torch.equal(dst_sd[key], src_sd[key]), key | |
| + | |
| + vals = [e for _, e in errors] | |
| + assert measured["max_relative_error"] == max(vals) | |
| + assert measured["mean_relative_error"] == sum(vals) / len(vals) | |
| + assert measured["worst_tensor"] in dict(errors) | |
| + assert measured["max_relative_error"] == dict(errors)[measured["worst_tensor"]] | |
| + assert 0.0 <= measured["mean_relative_error"] <= measured["p99_relative_error"] | |
| + assert measured["p99_relative_error"] <= measured["max_relative_error"] | |
| + assert measured["max_relative_error"] < 0.05, measured | |
| + | |
| + # Eager quant of the same BF16 bytes. | |
| + eager = TinyMing().to(dtype=torch.bfloat16) | |
| + incompatible = eager.load_state_dict(src_sd, strict=True) | |
| + assert not incompatible.missing_keys and not incompatible.unexpected_keys | |
| + _apply_int8_(eager) | |
| + | |
| + loaded = _move_parameters_to_meta(TinyMing()) | |
| + for layer in loaded.model.model.layers: | |
| + assert layer.attention.inv_freq.device.type == "cpu" | |
| + assert layer.attention.query_key_value.weight.device.type == "meta" | |
| + report = load_int8_mllm_(loaded, dst, "cpu") | |
| + assert report["modules_swapped"] == len(EXPECTED_QUANT_MODULES) | |
| + assert report["tensors_loaded"] == len(dst_sd) | |
| + assert report["bytes_loaded"] == manifest["total_size"] | |
| + _assert_no_meta(loaded) | |
| + for layer in loaded.model.model.layers: | |
| + assert layer.attention.inv_freq.device.type == "cpu" | |
| + assert layer.attention.inv_freq.dtype == torch.float32 | |
| + for name in EXPECTED_QUANT_MODULES: | |
| + mod = loaded.get_submodule(name) | |
| + assert isinstance(mod, Int8Linear), name | |
| + assert mod.weight.dtype == torch.int8 | |
| + assert mod.scale.dtype == torch.float32 | |
| + | |
| + eager.eval() | |
| + loaded.eval() | |
| + ids = torch.randint(0, VOCAB, (2, 6)) | |
| + with torch.no_grad(): | |
| + y_eager = eager(ids) | |
| + y_loaded = loaded(ids) | |
| + assert y_eager.dtype == y_loaded.dtype | |
| + assert torch.equal(y_eager, y_loaded), (y_eager - y_loaded).abs().max().item() | |
| + | |
| + # A second run into a non-empty safetensors dir must fail loudly. | |
| + print(" re-running into a non-empty dst (expect error on stderr)", flush=True) | |
| + rc_again = quantize_stream.main([str(src), str(dst)]) | |
| + assert rc_again == 1 | |
| + | |
| + | |
| +def test_unknown_key_fails_loudly() -> None: | |
| + torch.manual_seed(4) | |
| + model = TinyMing().to(dtype=torch.bfloat16) | |
| + with tempfile.TemporaryDirectory(prefix="ming-int8-bad-") as tmp: | |
| + root = Path(tmp) | |
| + src = root / "src" | |
| + dst = root / "dst" | |
| + _save_bf16_checkpoint(model, src) | |
| + rc = quantize_stream.main([str(src), str(dst)]) | |
| + assert rc == 0, rc | |
| + shard = next(dst.glob("*.safetensors")) | |
| + sd = load_file(str(shard)) | |
| + sd["not.a.real.key"] = torch.zeros(4, dtype=torch.float32) | |
| + save_file(sd, str(shard)) | |
| + loaded = _move_parameters_to_meta(TinyMing()) | |
| + try: | |
| + load_int8_mllm_(loaded, dst, "cpu") | |
| + except RuntimeError as exc: | |
| + text = str(exc) | |
| + assert "unexpected" in text.lower(), text | |
| + assert "not.a.real.key" in text, text | |
| + print(f" caught RuntimeError: {text.splitlines()[0]}") | |
| + else: | |
| + raise AssertionError("load_int8_mllm_ returned instead of failing on an unknown key") | |
| + | |
| + | |
| +def main() -> int: | |
| + import safetensors | |
| + | |
| + print(f"torch={torch.__version__} safetensors={safetensors.__version__}", flush=True) | |
| + tests = [ | |
| + test_from_linear_roundtrip, | |
| + test_dtype_cast_keeps_scale_fp32, | |
| + test_forward_matches_reference, | |
| + test_is_quantizable_rule, | |
| + test_end_to_end_stream_and_load, | |
| + test_unknown_key_fails_loudly, | |
| + ] | |
| + failed = 0 | |
| + for fn in tests: | |
| + try: | |
| + fn() | |
| + except Exception: | |
| + failed += 1 | |
| + print(f"FAIL {fn.__name__}", flush=True) | |
| + traceback.print_exc() | |
| + else: | |
| + print(f"PASS {fn.__name__}", flush=True) | |
| + print(f"{len(tests) - failed} passed, {failed} failed", flush=True) | |
| + return 1 if failed else 0 | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + sys.exit(main()) | |
| diff --git a/qwen2_5_vit.py b/qwen2_5_vit.py | |
| index 3de0e73..efa8255 100644 | |
| --- a/qwen2_5_vit.py | |
| +++ b/qwen2_5_vit.py | |
| from transformers.utils import ( | |
| from typing import Union | |
| from transformers.configuration_utils import PretrainedConfig | |
| -import transformer_engine.pytorch as te | |
| if is_flash_attn_2_available(): | |
| from flash_attn import flash_attn_varlen_func | |
| class Qwen2_5_VisionRotaryEmbedding(nn.Module): | |
| new_inv_freq = 1.0 / (self.theta ** (torch.arange(0, self.dim, 2, dtype=torch.float) / self.dim)) | |
| self.inv_freq.copy_(new_inv_freq) | |
| -class Qwen2RMSNorm(te.RMSNorm): | |
| +class Qwen2RMSNorm(nn.Module): | |
| def __init__(self, hidden_size, eps=1e-6): | |
| """ | |
| - Qwen2RMSNorm is equivalent to T5LayerNorm | |
| + Qwen2RMSNorm is equivalent to T5LayerNorm. | |
| + | |
| + Replaces transformer_engine.pytorch.RMSNorm: ROCm has no transformer-engine. | |
| + te.RMSNorm defaults (zero_centered_gamma=False) are standard RMSNorm, and the | |
| + checkpoint stores this affine as `weight`. | |
| """ | |
| - super().__init__(hidden_size, eps=eps) | |
| + super().__init__() | |
| + self.weight = nn.Parameter(torch.ones(hidden_size)) | |
| + self.variance_epsilon = eps | |
| + | |
| + def forward(self, hidden_states): | |
| + input_dtype = hidden_states.dtype | |
| + hidden_states = hidden_states.to(torch.float32) | |
| + variance = hidden_states.pow(2).mean(-1, keepdim=True) | |
| + hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) | |
| + return self.weight * hidden_states.to(input_dtype) | |
| + | |
| + def extra_repr(self): | |
| + return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" | |
| class Qwen2_5_VLPatchMerger(nn.Module): | |
| def __init__(self, dim: int, context_dim: int, spatial_merge_size: int = 2) -> None: | |
| diff --git a/requirements-rocm.txt b/requirements-rocm.txt | |
| new file mode 100644 | |
| index 0000000..d8cda21 | |
| --- /dev/null | |
| +++ b/requirements-rocm.txt | |
| +# ROCm port of requirements.txt for AMD gfx1151 (ROCm 7.13). | |
| +# Omitted on purpose — do not add them back: | |
| +# torch, torchvision: the target interpreter already has a working ROCm | |
| +# build (torch 2.10.0, torch.version.hip 7.13.99004). Reinstalling the | |
| +# upstream CUDA pins would replace it. | |
| +# transformer-engine: NVIDIA CUDA-only; ROCm has no TE. Qwen2RMSNorm in | |
| +# qwen2_5_vit.py is pure PyTorch, and the unused TE import is gone. | |
| +# Create the venv with system site packages so that ROCm torch is inherited: | |
| +# python3 -m venv --system-site-packages .venv | |
| +# .venv/bin/pip install -r requirements-rocm.txt | |
| +# | |
| +# Also omitted / relaxed versus upstream, because the ROCm interpreter is Python 3.13 | |
| +# and its torch is built against numpy 2.x: | |
| +# numpy upstream 1.23.1 has no Python 3.13 wheels, and downgrading would | |
| +# break the inherited torch. Inherit the system numpy (validated 2.2.4). | |
| +# Pillow upstream 10.4.0 has no Python 3.13 wheels. Inherit (validated 11.1.0). | |
| +# safetensors inherit the system build (validated 0.8.0). | |
| +# | |
| +# Validated on halo (gfx1151, ROCm 7.13, Python 3.13.5) on 2026-09-22: | |
| +# torch 2.10.0 (hip 7.13.99004) | numpy 2.2.4 | Pillow 11.1.0 | safetensors 0.8.0 | |
| +# transformers 4.57.1 | diffusers 0.36.0 | accelerate 1.13.0 | tokenizers 0.22.2 | |
| +# huggingface-hub 0.34.0 | peft 0.17.0 | |
| +transformers==4.57.1 | |
| +diffusers==0.36.0 | |
| +accelerate==1.13.0 | |
| +tokenizers==0.22.2 | |
| +huggingface-hub==0.34.0 | |
| +peft==0.17.0 | |
| +requests==2.32.3 | |
| +tqdm==4.67.1 | |
| +typing-extensions==4.15.0 | |
| + | |
| +# Optional FlashAttention 2 backend (validated: flash-attn==2.7.3). The CLI | |
| +# default is eager attention; --attn-implementation flash_attention_2 needs | |
| +# this package. | |
| +# flash-attn==2.7.3 | |
| diff --git a/tools/convert_connector.py b/tools/convert_connector.py | |
| new file mode 100644 | |
| index 0000000..8f3733c | |
| --- /dev/null | |
| +++ b/tools/convert_connector.py | |
| +#!/usr/bin/env python3 | |
| +"""Store the connector component (Qwen2 1.5B, shipped as float32) as bfloat16. | |
| + | |
| +infer.py loads the connector with torch_dtype=bfloat16, so the tensors it runs with are the | |
| +fp32 values rounded to bf16 at load time. This does the same rounding once, offline, and proves | |
| +every converted tensor equals `fp32_tensor.to(torch.bfloat16)` exactly — the runtime model is | |
| +unchanged; only the download halves. | |
| + | |
| + usage: convert_connector.py SRC_CONNECTOR_DIR DST_CONNECTOR_DIR | |
| +""" | |
| +import json | |
| +import shutil | |
| +import sys | |
| +from pathlib import Path | |
| + | |
| +import torch | |
| +from safetensors import safe_open | |
| +from safetensors.torch import load_file, save_file | |
| + | |
| + | |
| +def main(): | |
| + src, dst = Path(sys.argv[1]), Path(sys.argv[2]) | |
| + dst.mkdir(parents=True, exist_ok=True) | |
| + if any(dst.glob("*.safetensors")): | |
| + sys.exit(f"refusing: {dst} already contains safetensors") | |
| + index = json.loads((src / "model.safetensors.index.json").read_text()) | |
| + shards = sorted(set(index["weight_map"].values())) | |
| + | |
| + def converted(tensor): | |
| + return tensor.to(torch.bfloat16) if tensor.is_floating_point() else tensor | |
| + | |
| + out = {} | |
| + for shard in shards: | |
| + with safe_open(str(src / shard), "pt") as handle: | |
| + for key in handle.keys(): | |
| + if key in out: | |
| + sys.exit(f"duplicate tensor {key}") | |
| + out[key] = converted(handle.get_tensor(key)) | |
| + if set(out) != set(index["weight_map"]): | |
| + sys.exit("tensor set does not match the index weight_map") | |
| + target = dst / "model.safetensors" | |
| + save_file(out, str(target), metadata={"format": "pt"}) | |
| + del out | |
| + | |
| + back = load_file(str(target)) | |
| + checked = 0 | |
| + for shard in shards: | |
| + with safe_open(str(src / shard), "pt") as handle: | |
| + for key in handle.keys(): | |
| + reference = converted(handle.get_tensor(key)) | |
| + if back[key].dtype != reference.dtype or not torch.equal(back[key], reference): | |
| + sys.exit(f"MISMATCH {key}") | |
| + checked += 1 | |
| + if checked != len(back): | |
| + sys.exit(f"checked {checked} tensors but the output holds {len(back)}") | |
| + | |
| + for path in src.iterdir(): | |
| + if path.suffix == ".safetensors" or path.name == "model.safetensors.index.json": | |
| + continue | |
| + shutil.copy2(path, dst / path.name) | |
| + config = json.loads((dst / "config.json").read_text()) | |
| + key = "dtype" if "dtype" in config else "torch_dtype" | |
| + previous = config.get(key) | |
| + config[key] = "bfloat16" | |
| + (dst / "config.json").write_text(json.dumps(config, indent=2) + "\n") | |
| + dtypes = sorted({str(t.dtype) for t in back.values()}) | |
| + print(f"CONNECTOR_OK tensors={checked} exact=all dtypes={dtypes} bytes={target.stat().st_size} " | |
| + f"config.{key}: {previous} -> bfloat16") | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |
| diff --git a/tools/fidelity_compare.py b/tools/fidelity_compare.py | |
| new file mode 100644 | |
| index 0000000..99736d6 | |
| --- /dev/null | |
| +++ b/tools/fidelity_compare.py | |
| +#!/usr/bin/env python3 | |
| +"""Compare two ming_bench.py output dirs (reference vs candidate), stem by stem. | |
| + | |
| +Conditioning (what the DiT receives): cosine similarity over the whole tensor, the | |
| +per-token cosine (mean and worst token), and relative L2 = |a - b| / |a|. | |
| +Images: MAE, PSNR, windowed 7x7 SSIM on luminance, and alpha MAE for RGBA. | |
| + | |
| + usage: fidelity_compare.py <reference_dir> <candidate_dir> [--json out.json] | |
| +""" | |
| +import json | |
| +import sys | |
| +from pathlib import Path | |
| + | |
| +import numpy as np | |
| +from PIL import Image | |
| +from safetensors.numpy import load_file | |
| + | |
| + | |
| +def load_image(path): | |
| + im = Image.open(path) | |
| + rgb = np.asarray(im.convert("RGB"), dtype=np.float64) | |
| + alpha = np.asarray(im.convert("RGBA"), dtype=np.float64)[..., 3] if im.mode in ("RGBA", "LA") else None | |
| + return rgb, alpha, im.size | |
| + | |
| + | |
| +def box(x, k): | |
| + c = np.cumsum(np.cumsum(np.pad(x, ((1, 0), (1, 0))), 0), 1) | |
| + return (c[k:, k:] - c[:-k, k:] - c[k:, :-k] + c[:-k, :-k]) / (k * k) | |
| + | |
| + | |
| +def ssim(a, b, k=7, L=255.0): | |
| + c1, c2 = (0.01 * L) ** 2, (0.03 * L) ** 2 | |
| + mu_a, mu_b = box(a, k), box(b, k) | |
| + va, vb = box(a * a, k) - mu_a ** 2, box(b * b, k) - mu_b ** 2 | |
| + cov = box(a * b, k) - mu_a * mu_b | |
| + s = ((2 * mu_a * mu_b + c1) * (2 * cov + c2)) / ((mu_a ** 2 + mu_b ** 2 + c1) * (va + vb + c2)) | |
| + return float(s.mean()) | |
| + | |
| + | |
| +def image_metrics(ref_path, cand_path): | |
| + ra, aa, sa = load_image(ref_path) | |
| + rb, ab, sb = load_image(cand_path) | |
| + if sa != sb: | |
| + raise SystemExit(f"size mismatch {ref_path} {sa} vs {cand_path} {sb}") | |
| + lum = lambda x: 0.299 * x[..., 0] + 0.587 * x[..., 1] + 0.114 * x[..., 2] | |
| + mse = float(((ra - rb) ** 2).mean()) | |
| + out = { | |
| + "mae": round(float(np.abs(ra - rb).mean()), 3), | |
| + "psnr_db": None if mse == 0 else round(10 * np.log10(255.0 ** 2 / mse), 2), | |
| + "ssim_lum": round(ssim(lum(ra), lum(rb)), 4), | |
| + } | |
| + if aa is not None and ab is not None: | |
| + out["alpha_mae"] = round(float(np.abs(aa - ab).mean()), 3) | |
| + return out | |
| + | |
| + | |
| +def cond_metrics(ref_path, cand_path): | |
| + ref, cand = load_file(str(ref_path)), load_file(str(cand_path)) | |
| + out = {} | |
| + for key in sorted(set(ref) & set(cand)): | |
| + a, b = ref[key].astype(np.float64), cand[key].astype(np.float64) | |
| + if a.shape != b.shape: | |
| + raise SystemExit(f"{key}: shape mismatch {a.shape} vs {b.shape}") | |
| + fa, fb = a.ravel(), b.ravel() | |
| + tok_a, tok_b = a.reshape(-1, a.shape[-1]), b.reshape(-1, b.shape[-1]) | |
| + tok_cos = (tok_a * tok_b).sum(-1) / (np.linalg.norm(tok_a, axis=-1) * np.linalg.norm(tok_b, axis=-1)) | |
| + out[key] = { | |
| + "shape": list(a.shape), | |
| + "cosine": round(float(fa @ fb / (np.linalg.norm(fa) * np.linalg.norm(fb))), 6), | |
| + "token_cos_mean": round(float(tok_cos.mean()), 6), | |
| + "token_cos_min": round(float(tok_cos.min()), 6), | |
| + "rel_l2": round(float(np.linalg.norm(fa - fb) / np.linalg.norm(fa)), 6), | |
| + } | |
| + missing = sorted(set(ref) ^ set(cand)) | |
| + if missing: | |
| + raise SystemExit(f"conditioning keys present on one side only: {missing}") | |
| + return out | |
| + | |
| + | |
| +def main(): | |
| + ref_dir, cand_dir = Path(sys.argv[1]), Path(sys.argv[2]) | |
| + stems = sorted(p.stem for p in ref_dir.glob("*.png") if (cand_dir / p.name).exists()) | |
| + if not stems: | |
| + raise SystemExit(f"no common images between {ref_dir} and {cand_dir}") | |
| + rows = [] | |
| + for stem in stems: | |
| + row = {"stem": stem, "image": image_metrics(ref_dir / f"{stem}.png", cand_dir / f"{stem}.png")} | |
| + rc, cc = ref_dir / f"{stem}.cond.safetensors", cand_dir / f"{stem}.cond.safetensors" | |
| + if rc.exists() and cc.exists(): | |
| + row["cond"] = cond_metrics(rc, cc) | |
| + rows.append(row) | |
| + im = row["image"] | |
| + line = f"{stem:32s} SSIM {im['ssim_lum']:.4f} PSNR {im['psnr_db']} MAE {im['mae']:.2f}" | |
| + if "alpha_mae" in im: | |
| + line += f" aMAE {im['alpha_mae']:.2f}" | |
| + for key, c in row.get("cond", {}).items(): | |
| + line += f" | {key[:3]} cos {c['cosine']:.6f} tokmin {c['token_cos_min']:.4f} relL2 {c['rel_l2']:.4f}" | |
| + print(line) | |
| + if "--json" in sys.argv: | |
| + Path(sys.argv[sys.argv.index("--json") + 1]).write_text(json.dumps(rows, indent=2)) | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |
| diff --git a/tools/ming_bench.py b/tools/ming_bench.py | |
| new file mode 100644 | |
| index 0000000..c92bb3f | |
| --- /dev/null | |
| +++ b/tools/ming_bench.py | |
| +#!/usr/bin/env python3 | |
| +"""Ming-Image speed + fidelity harness: one model load, N prompts. | |
| + | |
| +Reuses infer.py's own loader and generation path unchanged. The only addition is a | |
| +wrapper around model.diffusion_loss.sample that records the conditioning tensors the | |
| +DiT receives (encoder_hidden_states / directvlm_hidden_states) and times the sampling | |
| +stage (DiT steps + VAE decode) separately from the MLLM stage. | |
| + | |
| + usage: ming_bench.py --prompts a.json b.json --out DIR [--repeat-first N] -- <infer.py args> | |
| + | |
| + <infer.py args> are passed to infer.parse_args() as-is (e.g. --model, --resolution, | |
| + --steps, --seed, --device, --device-map none, --attn-implementation eager, --int8-mllm). | |
| + --repeat-first N re-runs the first prompt N more times at the same seed: the images | |
| + measure the platform's run-to-run noise floor and the timings are warm timings. | |
| +""" | |
| +import argparse | |
| +import json | |
| +import sys | |
| +import time | |
| +from pathlib import Path | |
| + | |
| + | |
| +def main(): | |
| + ap = argparse.ArgumentParser() | |
| + ap.add_argument("--prompts", nargs="+", required=True) | |
| + ap.add_argument("--out", required=True) | |
| + ap.add_argument("--repeat-first", type=int, default=0) | |
| + own, rest = ap.parse_known_args() | |
| + if rest and rest[0] == "--": | |
| + rest = rest[1:] | |
| + sys.argv = [sys.argv[0], "--prompt", own.prompts[0]] + rest | |
| + | |
| + import torch | |
| + from safetensors.torch import save_file | |
| + import infer | |
| + | |
| + args = infer.parse_args() | |
| + model_directory = infer.resolve_model_directory( | |
| + args.model, revision=args.revision, cache_dir=args.cache_dir, | |
| + local_files_only=args.local_files_only, | |
| + ) | |
| + profile = infer.load_checkpoint_capabilities(model_directory) | |
| + resolution = infer.resolve_task_resolution(args.task, args.resolution) | |
| + sampling = profile.resolve_sampling_parameters(steps=args.steps, cfg=args.cfg) | |
| + dtype = infer._dtype(args.dtype) | |
| + out = Path(own.out) | |
| + out.mkdir(parents=True, exist_ok=True) | |
| + | |
| + def sync(): | |
| + if torch.cuda.is_available(): | |
| + torch.cuda.synchronize() | |
| + | |
| + sync() | |
| + t0 = time.perf_counter() | |
| + model, processor = infer.load_model_and_processor(model_directory, args) | |
| + sync() | |
| + load_s = time.perf_counter() - t0 | |
| + print(f"LOAD_S {load_s:.1f}", flush=True) | |
| + | |
| + captured = {} | |
| + original_sample = model.diffusion_loss.sample | |
| + | |
| + def recording_sample(*a, **kw): | |
| + for key in ("encoder_hidden_states", "directvlm_hidden_states"): | |
| + value = kw.get(key) | |
| + if isinstance(value, (list, tuple)): | |
| + value = torch.stack(list(value), dim=0) | |
| + if isinstance(value, torch.Tensor): | |
| + captured[key] = value.detach().float().cpu().contiguous() | |
| + sync() | |
| + ts = time.perf_counter() | |
| + result = original_sample(*a, **kw) | |
| + sync() | |
| + captured["_sample_s"] = time.perf_counter() - ts | |
| + return result | |
| + | |
| + model.diffusion_loss.sample = recording_sample | |
| + | |
| + runs = [(p, 0) for p in own.prompts] + [(own.prompts[0], i + 1) for i in range(own.repeat_first)] | |
| + results = [] | |
| + for prompt_path, rep in runs: | |
| + stem = Path(prompt_path).stem + (f"_rep{rep}" if rep else "") | |
| + prompt = infer._load_prompt(prompt_path) | |
| + captured.clear() | |
| + if torch.cuda.is_available(): | |
| + torch.cuda.reset_peak_memory_stats() | |
| + sync() | |
| + t1 = time.perf_counter() | |
| + images = infer.run_generation( | |
| + model, processor, profile, task=args.task, prompt=prompt, input_image=None, | |
| + resolution=resolution, sampling=sampling, seed=args.seed, num_layers=args.num_layers, | |
| + dtype=dtype, | |
| + ) | |
| + sync() | |
| + total_s = time.perf_counter() - t1 | |
| + if len(images) != 1: | |
| + raise RuntimeError(f"{stem}: expected 1 image, got {len(images)}") | |
| + image_path = out / f"{stem}.png" | |
| + images[0].save(image_path) | |
| + cond = {k: v for k, v in captured.items() if not k.startswith("_")} | |
| + if "encoder_hidden_states" not in cond: | |
| + raise RuntimeError(f"{stem}: conditioning was not captured") | |
| + save_file(cond, str(out / f"{stem}.cond.safetensors")) | |
| + sample_s = captured["_sample_s"] | |
| + row = { | |
| + "load_s": round(load_s, 1), | |
| + "prompt": str(prompt_path), "stem": stem, "seed": args.seed, "resolution": resolution, | |
| + "steps": sampling.steps, "cfg": sampling.cfg, "mode": images[0].mode, | |
| + "size": list(images[0].size), "total_s": round(total_s, 2), | |
| + "sample_s": round(sample_s, 2), "mllm_s": round(total_s - sample_s, 2), | |
| + "peak_alloc_gib": round(torch.cuda.max_memory_allocated() / 2**30, 2) | |
| + if torch.cuda.is_available() else None, | |
| + "cond_shapes": {k: list(v.shape) for k, v in cond.items()}, | |
| + } | |
| + results.append(row) | |
| + print("RUN " + json.dumps(row), flush=True) | |
| + with open(out / "runs.jsonl", "a") as fh: # accumulates across one-prompt-per-process runs | |
| + fh.write(json.dumps(row) + "\n") | |
| + | |
| + manifest = {"load_s": round(load_s, 1), "args": {k: str(v) for k, v in vars(args).items()}, | |
| + "runs": results} | |
| + (out / "manifest.json").write_text(json.dumps(manifest, indent=2)) | |
| + print("BENCH_DONE", out, flush=True) | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |
| diff --git a/tools/sdpa_layout.py b/tools/sdpa_layout.py | |
| new file mode 100644 | |
| index 0000000..cfb3720 | |
| --- /dev/null | |
| +++ b/tools/sdpa_layout.py | |
| +#!/usr/bin/env python3 | |
| +"""Math SDPA with the DiT's real input layout: [B, L, H, D] permuted to [B, H, L, D] (non-contiguous, | |
| +exactly what diffusers' native attention backend passes) vs the same tensors made contiguous. | |
| +Speed and error vs an fp32 reference, masked, at the cabin prompt's real length. | |
| + | |
| + usage: sdpa_layout.py [L] | |
| +""" | |
| +import sys | |
| +import time | |
| + | |
| +import torch | |
| +import torch.nn.functional as F | |
| +from torch.nn.attention import SDPBackend, sdpa_kernel | |
| + | |
| +L = int(sys.argv[1]) if len(sys.argv) > 1 else 5759 | |
| +H, D, dev = 30, 128, "cuda" | |
| +g = torch.Generator(device=dev).manual_seed(0) | |
| +blhd = [torch.randn(1, L, H, D, device=dev, dtype=torch.bfloat16, generator=g) for _ in range(3)] | |
| +q, k, v = (x.permute(0, 2, 1, 3) for x in blhd) # views, as diffusers passes them | |
| +qc, kc, vc = (x.contiguous() for x in (q, k, v)) | |
| +mask = torch.ones(1, 1, 1, L, dtype=torch.bool, device=dev) | |
| +mask[..., L - 64:] = False | |
| +with sdpa_kernel(SDPBackend.MATH): | |
| + ref = F.scaled_dot_product_attention(qc.float(), kc.float(), vc.float(), attn_mask=mask) | |
| + | |
| + | |
| +def run(tag, a, b, c, bf16_reduction): | |
| + torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(bf16_reduction) | |
| + with sdpa_kernel(SDPBackend.MATH): | |
| + out = F.scaled_dot_product_attention(a, b, c, attn_mask=mask) | |
| + torch.cuda.synchronize() | |
| + t0 = time.perf_counter() | |
| + for _ in range(3): | |
| + out = F.scaled_dot_product_attention(a, b, c, attn_mask=mask) | |
| + torch.cuda.synchronize() | |
| + ms = (time.perf_counter() - t0) / 3 * 1000 | |
| + rel = ((out.float() - ref).norm() / ref.norm()).item() | |
| + exact = torch.equal(out, base) if base is not None else None | |
| + print(f" {tag:34s} {ms:8.2f} ms rel_l2 {rel:.3e} identical_to_default: {exact}") | |
| + return out | |
| + | |
| + | |
| +base = None | |
| +print(f"torch {torch.__version__} | L={L} | q strides {tuple(q.stride())} contiguous={q.is_contiguous()}") | |
| +base = run("permuted views (DiT today), fp32", q, k, v, False) | |
| +run("contiguous, fp32 (math unchanged)", qc, kc, vc, False) | |
| +run("permuted views, bf16 reduction", q, k, v, True) | |
| +run("contiguous, bf16 reduction", qc, kc, vc, True) | |
| +torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(False) | |
| diff --git a/tools/step_probe.py b/tools/step_probe.py | |
| new file mode 100644 | |
| index 0000000..dfd8af2 | |
| --- /dev/null | |
| +++ b/tools/step_probe.py | |
| +#!/usr/bin/env python3 | |
| +"""Per-step timing of Ming-Image's DiT with allocator stats and a GPU clock/power sampler. | |
| + | |
| +Diagnoses step time that grows within one generation. For every DiT call it records the | |
| +synchronized wall time, the caching allocator's reserved/allocated bytes, how many device | |
| +mallocs and malloc retries (fragmentation) have happened so far; a sampler thread reads the | |
| +GPU sclk, power and temperature twice a second. | |
| + | |
| + usage: PYTHONPATH=<code_dir> step_probe.py --prompt P.json [--runs N] -- <infer.py args> | |
| +""" | |
| +import argparse | |
| +import glob | |
| +import json | |
| +import sys | |
| +import threading | |
| +import time | |
| + | |
| + | |
| +def read_gpu(): | |
| + base = "/sys/class/drm/card0/device" | |
| + sclk = next((l.split(":")[1].strip().rstrip("*").strip() for l in open(f"{base}/pp_dpm_sclk") if "*" in l), "?") | |
| + hw = sorted(glob.glob(f"{base}/hwmon/hwmon*"))[0] | |
| + power = int(open(f"{hw}/power1_average").read()) / 1e6 | |
| + temp = int(open(f"{hw}/temp1_input").read()) / 1e3 | |
| + busy = int(open(f"{base}/gpu_busy_percent").read()) | |
| + return sclk, power, temp, busy | |
| + | |
| + | |
| +def main(): | |
| + ap = argparse.ArgumentParser() | |
| + ap.add_argument("--prompt", required=True) | |
| + ap.add_argument("--runs", type=int, default=1) | |
| + own, rest = ap.parse_known_args() | |
| + if rest and rest[0] == "--": | |
| + rest = rest[1:] | |
| + sys.argv = [sys.argv[0], "--prompt", own.prompt] + rest | |
| + | |
| + import torch | |
| + import infer | |
| + | |
| + args = infer.parse_args() | |
| + model_directory = infer.resolve_model_directory(args.model, local_files_only=True) | |
| + caps = infer.load_checkpoint_capabilities(model_directory) | |
| + resolution = infer.resolve_task_resolution(args.task, args.resolution) | |
| + sampling = caps.resolve_sampling_parameters(steps=args.steps, cfg=args.cfg) | |
| + dtype = infer._dtype(args.dtype) | |
| + model, processor = infer.load_model_and_processor(model_directory, args) | |
| + prompt = infer._load_prompt(own.prompt) | |
| + | |
| + samples, stop = [], threading.Event() | |
| + | |
| + def sampler(): | |
| + t0 = time.perf_counter() | |
| + while not stop.is_set(): | |
| + samples.append((round(time.perf_counter() - t0, 1),) + read_gpu()) | |
| + time.sleep(0.5) | |
| + | |
| + dit = model.diffusion_loss.train_model | |
| + marks = {} | |
| + | |
| + def pre(_module, _args, _kwargs): | |
| + torch.cuda.synchronize() | |
| + marks["t"] = time.perf_counter() | |
| + | |
| + def post(_module, _args, _kwargs, _out): | |
| + torch.cuda.synchronize() | |
| + st = torch.cuda.memory_stats() | |
| + step_log.append({ | |
| + "step_s": round(time.perf_counter() - marks["t"], 2), | |
| + "reserved_gib": round(torch.cuda.memory_reserved() / 2**30, 2), | |
| + "allocated_gib": round(torch.cuda.memory_allocated() / 2**30, 2), | |
| + "device_mallocs": st.get("num_device_alloc", 0), | |
| + "device_frees": st.get("num_device_free", 0), | |
| + "alloc_retries": st.get("num_alloc_retries", 0), | |
| + }) | |
| + | |
| + dit.register_forward_pre_hook(pre, with_kwargs=True) | |
| + dit.register_forward_hook(post, with_kwargs=True) | |
| + thread = threading.Thread(target=sampler, daemon=True) | |
| + thread.start() | |
| + for run in range(own.runs): | |
| + step_log = [] | |
| + torch.cuda.synchronize() | |
| + t0 = time.perf_counter() | |
| + infer.run_generation(model, processor, caps, task=args.task, prompt=prompt, input_image=None, | |
| + resolution=resolution, sampling=sampling, seed=args.seed, | |
| + num_layers=args.num_layers, dtype=dtype) | |
| + torch.cuda.synchronize() | |
| + print(f"RUN {run} total_s {time.perf_counter() - t0:.1f}", flush=True) | |
| + for i, row in enumerate(step_log): | |
| + print("STEP " + json.dumps({"run": run, "i": i, **row}), flush=True) | |
| + stop.set() | |
| + thread.join() | |
| + for s in samples[:: max(1, len(samples) // 60)]: | |
| + print("GPU t=%6.1fs sclk=%s power=%.0fW temp=%.0fC busy=%d%%" % s) | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |
| diff --git a/tools/verify_package.py b/tools/verify_package.py | |
| new file mode 100644 | |
| index 0000000..cfb91f7 | |
| --- /dev/null | |
| +++ b/tools/verify_package.py | |
| +#!/usr/bin/env python3 | |
| +"""Pre-upload verification of the INT8 package against the upstream download. Read-only on both trees | |
| +except for writing SHA256SUMS into the package. Exits non-zero on the first failed check. | |
| + | |
| + usage: verify_package.py UPSTREAM_DIR PACKAGE_DIR | |
| +""" | |
| +import hashlib | |
| +import json | |
| +import os | |
| +import sys | |
| +from pathlib import Path | |
| + | |
| +import torch | |
| +from safetensors import safe_open | |
| + | |
| + | |
| +def fail(msg): | |
| + sys.exit(f"VERIFY_FAIL {msg}") | |
| + | |
| + | |
| +def shard_map(d): | |
| + index = json.loads((d / "model.safetensors.index.json").read_text()) | |
| + return index["weight_map"] | |
| + | |
| + | |
| +def main(): | |
| + up, pkg = Path(sys.argv[1]), Path(sys.argv[2]) | |
| + | |
| + # 1. connector: bf16 file == upstream fp32 cast to bf16, tensor by tensor | |
| + up_map = shard_map(up / "connector") | |
| + with safe_open(str(pkg / "connector/model.safetensors"), "pt") as new: | |
| + if set(new.keys()) != set(up_map): | |
| + fail("connector tensor set differs from upstream") | |
| + n = 0 | |
| + for shard in sorted(set(up_map.values())): | |
| + with safe_open(str(up / "connector" / shard), "pt") as old: | |
| + for key in old.keys(): | |
| + ref = old.get_tensor(key) | |
| + ref = ref.to(torch.bfloat16) if ref.is_floating_point() else ref | |
| + got = new.get_tensor(key) | |
| + if got.dtype != ref.dtype or not torch.equal(got, ref): | |
| + fail(f"connector {key} != fp32->bf16") | |
| + n += 1 | |
| + print(f"OK connector: {n} tensors equal upstream fp32 -> bf16", flush=True) | |
| + | |
| + # 2. unchanged components: hardlink (same inode) or identical bytes | |
| + same = 0 | |
| + for comp in ("transformer", "vae", "mlp", "scheduler"): | |
| + for f in sorted((up / comp).rglob("*")): | |
| + if f.is_dir(): | |
| + continue | |
| + g = pkg / f.relative_to(up) | |
| + if not g.is_file(): | |
| + fail(f"missing {g}") | |
| + if os.stat(f).st_ino != os.stat(g).st_ino and f.read_bytes() != g.read_bytes(): | |
| + fail(f"{g} differs from upstream") | |
| + same += 1 | |
| + if (up / "LICENSE").read_bytes() != (pkg / "LICENSE").read_bytes(): | |
| + fail("LICENSE differs from upstream") | |
| + print(f"OK unchanged components: {same} files identical to upstream (+ LICENSE)", flush=True) | |
| + | |
| + # 3. mllm: copied tensors byte-identical, quantized ones present as int8 + fp32 scale | |
| + manifest = json.loads((pkg / "mllm/int8_manifest.json").read_text()) | |
| + quant = set(manifest["quantized_modules"]) | |
| + old_map, new_map = shard_map(up / "mllm"), shard_map(pkg / "mllm") | |
| + expect_new = {k for k in old_map if k[: -len(".weight")] not in quant or not k.endswith(".weight")} | |
| + expect_new |= {m + ".weight" for m in quant} | {m + ".scale" for m in quant} | |
| + if set(new_map) != expect_new: | |
| + fail(f"mllm index: {len(set(new_map) ^ expect_new)} names differ from the expected set") | |
| + handles = {} | |
| + | |
| + def tensor(tree, mapping, key): | |
| + path = str(tree / mapping[key]) | |
| + if path not in handles: | |
| + handles[path] = safe_open(path, "pt") | |
| + return handles[path].get_tensor(key) | |
| + | |
| + copied = quantized = 0 | |
| + for key in sorted(old_map): | |
| + module = key[: -len(".weight")] if key.endswith(".weight") else None | |
| + ref = tensor(up / "mllm", old_map, key) | |
| + if module in quant: | |
| + w, s = tensor(pkg / "mllm", new_map, key), tensor(pkg / "mllm", new_map, module + ".scale") | |
| + if w.dtype != torch.int8 or s.dtype != torch.float32 or w.shape != ref.shape or s.shape != (ref.shape[0],): | |
| + fail(f"{key}: int8/scale dtype or shape wrong") | |
| + quantized += 1 | |
| + else: | |
| + got = tensor(pkg / "mllm", new_map, key) | |
| + if got.dtype != ref.dtype or not torch.equal(got, ref): | |
| + fail(f"{key}: copied tensor differs from upstream") | |
| + copied += 1 | |
| + if len(handles) > 4: | |
| + handles.clear() | |
| + print(f"OK mllm: {copied} tensors byte-identical to upstream, {quantized} quantized (int8 + fp32 scale)", flush=True) | |
| + | |
| + # 4. sha256 of every file in the package | |
| + lines = [] | |
| + for f in sorted(p for p in pkg.rglob("*") if p.is_file() and p.name != "SHA256SUMS" and ".cache" not in p.parts): | |
| + h = hashlib.sha256() | |
| + with open(f, "rb") as fh: | |
| + for chunk in iter(lambda: fh.read(1 << 24), b""): | |
| + h.update(chunk) | |
| + lines.append(f"{h.hexdigest()} {f.relative_to(pkg).as_posix()}") | |
| + (pkg / "SHA256SUMS").write_text("\n".join(lines) + "\n") | |
| + print(f"OK sha256: {len(lines)} files -> SHA256SUMS", flush=True) | |
| + print("VERIFY_OK", flush=True) | |
| + | |
| + | |
| +if __name__ == "__main__": | |
| + main() | |