diff --git a/docs/en/model_normalize.md b/docs/en/model_normalize.md new file mode 100644 index 0000000000..69e0a77103 --- /dev/null +++ b/docs/en/model_normalize.md @@ -0,0 +1,142 @@ +# HF model normalization + +XTuner provides a standalone conversion entry point for turning training +checkpoints that already have HF tensor names into standard HF shard layouts. +It does not upload models or run inference validation. + +## BF16/FP16 repack + +```bash +bash xtuner/tools/model_normalize/run_model_normalize.sh repack \ + --source /path/to/source \ + --output /path/to/output \ + --shard-size-gb 4 +``` + +All tensors, including MTP tensors, are retained. Non-weight files such as +`config.json`, tokenizer files, chat templates, and an existing +`generation_config.json` are copied from the source. + +## Base-model assets (LICENSE + generation_config) + +Both `repack` and `to-fp8` accept `--base-model-dir `, a directory of +user-supplied release assets applied *after* conversion completes: + +- `generation_config.json` (if present) is copied verbatim into the output, + overwriting any existing file (the overwrite is logged). +- `LICENSE` (if present) is copied into the output with its Copyright line + rewritten to the fixed string `Copyright 2025-2026 Shanghai AI Laboratory`; + the rest of the license text (e.g. the MIT permission grant) is left + untouched. An existing output `LICENSE` is overwritten (logged). If no + Copyright line is found, the file is written unchanged with a warning. + +When `--base-model-dir` is omitted, conversion behavior is unchanged. A +source-supplied `generation_config.json` remains in the output as-is. + +## FP8 conversion + +For a reference-guided conversion: + +```bash +bash xtuner/tools/model_normalize/run_model_normalize.sh to-fp8 \ + --source /path/to/bf16 \ + --output /path/to/fp8 \ + --reference /path/to/reference-fp8 \ + --max-save-workers 4 +``` + +Without a reference, the heuristic policy must be explicit: + +```bash +python -m xtuner.tools.model_normalize to-fp8 \ + --source /path/to/bf16 --output /path/to/fp8 --policy heuristic +``` + +FP8 conversion requires CUDA. Save workers are bounded so conversion and disk +writes can overlap without submitting an unbounded number of shard writes. + +## One-command GLM-5.2 example + +`examples/run_glm52.sh` is a small wrapper around the generic CLI. It writes +the two release variants below one level under `OUTPUT_ROOT`. + +For the BF16, standard-shard variant: + +```bash +SOURCE_DIR=/path/to/glm52-hf-source \ +OUTPUT_ROOT=/path/to/release \ +bash xtuner/tools/model_normalize/examples/run_glm52.sh bf16 +``` + +For the reference-guided FP8 variant: + +```bash +SOURCE_DIR=/path/to/glm52-hf-source \ +OUTPUT_ROOT=/path/to/release \ +REFERENCE_DIR=/path/to/glm52-fp8-reference \ +SHARD_SIZE_GB=4 \ +MAX_SAVE_WORKERS=4 \ +bash xtuner/tools/model_normalize/examples/run_glm52.sh fp8 +``` + +To stamp a user-supplied `LICENSE` and `generation_config.json` into the +product, set `BASE_MODEL_DIR` (optional; applies to both variants): + +```bash +SOURCE_DIR=/path/to/glm52-hf-source \ +OUTPUT_ROOT=/path/to/release \ +BASE_MODEL_DIR=/path/to/base-model/glm5-2 \ +bash xtuner/tools/model_normalize/examples/run_glm52.sh bf16 +``` + +`run_glm52.sh` environment variables: + +| Variable | BF16 | FP8 | Default | Description | +|---|---|---|---|---| +| `SOURCE_DIR` | required | required | — | Input HF model directory | +| `OUTPUT_ROOT` | required | required | — | Output root directory | +| `SHARD_SIZE_GB` | optional | optional | `4` | Target shard size in GiB | +| `REFERENCE_DIR` | not used | required | — | FP8 reference model directory | +| `MAX_SAVE_WORKERS` | not used | optional | `4` | Parallel FP8 shard-save workers | +| `BASE_MODEL_DIR` | optional | optional | unset | Directory with user-supplied `LICENSE` and `generation_config.json` | + +The wrapper produces this shape (the actual shard count depends on the input +and `SHARD_SIZE_GB`): + +```text +/ +├── 20_hf_bf16_mtp/ +│ ├── model-00001-of-00NNN.safetensors +│ ├── ... +│ ├── model.safetensors.index.json +│ ├── config.json +│ ├── tokenizer.json / tokenizer_config.json +│ ├── chat_template.jinja # when supplied by the source +│ ├── generation_config.json # from source or base-model-dir +│ └── LICENSE # when supplied via --base-model-dir (Copyright rewritten) +└── 20_hf_fp8_mtp/ + ├── model-00001-of-00NNN.safetensors + ├── ... + ├── model.safetensors.index.json + ├── config.json # includes the FP8 quantization metadata + ├── tokenizer/chat-template files + ├── generation_config.json # from source or base-model-dir + └── LICENSE # when supplied via --base-model-dir (Copyright rewritten) +``` + +The BF16 output keeps the original tensor keys, including MTP keys, and only +repackages them into standard HF shards. The FP8 output also keeps the complete +key set, while converting selected weights to FP8 and writing their matching +`*_scale_inv` tensors. Both variants regenerate a consistent +`model.safetensors.index.json`; neither variant uploads to the Hub or performs a +full-model validation/SHA256 scan. + +## MTP and output safety + +The tool has no option that drops MTP. Repack and FP8 conversion preserve the +complete tensor key set and the source MTP configuration. Source and output +directories must differ, and a non-empty output is not overwritten unless +`--overwrite` is passed. + +The tool intentionally does not run a full-model validation or SHA256 scan. +Those checks remain optional release-side operations. diff --git a/pyproject.toml b/pyproject.toml index a941493dfe..7e2244db79 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,8 @@ dependencies = [ "tilelang==0.1.11", "apache-tvm-ffi==0.1.11", "torch>=2.6.0", + "safetensors", + "tqdm", "torchvision", "transformers==5.14.1", "cyclopts", diff --git a/tests/tools/test_model_normalize_cli.py b/tests/tools/test_model_normalize_cli.py new file mode 100644 index 0000000000..c0b520f20d --- /dev/null +++ b/tests/tools/test_model_normalize_cli.py @@ -0,0 +1,191 @@ +import json + +import pytest + +from xtuner.tools.model_normalize.cli import _reference_predicate + + +def test_reference_policy_selects_scale_companions(tmp_path): + reference = tmp_path / "reference" + reference.mkdir() + (reference / "model.safetensors.index.json").write_text( + json.dumps( + { + "weight_map": { + "model.layers.0.mlp.up_proj.weight": "model-00001-of-00001.safetensors", + "model.layers.0.mlp.up_proj.weight_scale_inv": "model-00001-of-00001.safetensors", + "model.layers.0.mtp.fc.weight": "model-00001-of-00001.safetensors", + } + } + ) + ) + predicate = _reference_predicate(reference) + assert predicate("model.layers.0.mlp.up_proj.weight") + assert not predicate("model.layers.0.mtp.fc.weight") + + +def test_mtp_keys_are_not_filtered_by_reference_policy(tmp_path): + reference = tmp_path / "reference" + reference.mkdir() + (reference / "model.safetensors.index.json").write_text( + json.dumps( + { + "weight_map": { + "model.layers.0.mtp.fc.weight": "model-00001-of-00001.safetensors", + "model.layers.0.mtp.fc.weight_scale_inv": "model-00001-of-00001.safetensors", + } + } + ) + ) + predicate = _reference_predicate(reference) + assert predicate("model.layers.0.mtp.fc.weight") + + +def test_repack_preserves_mtp_keys(tmp_path): + save_file = pytest.importorskip("safetensors.torch").save_file + import torch + + from xtuner.tools.model_normalize.repack import repack + + source = tmp_path / "source" + output = tmp_path / "output" + source.mkdir() + tensors = { + "model.layers.0.mlp.up_proj.weight": torch.ones((2, 2), dtype=torch.bfloat16), + "model.layers.0.mtp.fc.weight": torch.zeros((2, 2), dtype=torch.bfloat16), + } + save_file(tensors, source / "engine-rank0.safetensors") + (source / "model.safetensors.index.json").write_text( + json.dumps( + { + "weight_map": dict.fromkeys(tensors, "engine-rank0.safetensors"), + } + ) + ) + repack(source, output, shard_size_bytes=1024) + index = json.loads((output / "model.safetensors.index.json").read_text()) + assert set(index["weight_map"]) == set(tensors) + assert (output / "model-00001-of-00001.safetensors").is_file() + + +MIT_LICENSE = """\ +MIT License + +Copyright (c) 2026 Zhipu AI + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +""" + + +def test_rewrite_license_holder_replaces_copyright_line(): + from xtuner.tools.model_normalize.cli import _rewrite_license_holder + + rewritten, replaced = _rewrite_license_holder(MIT_LICENSE) + assert replaced + assert "Copyright 2025-2026 Shanghai AI Laboratory" in rewritten + # The original holder is gone. + assert "Zhipu AI" not in rewritten + # The MIT permission body is preserved verbatim. + assert "Permission is hereby granted, free of charge" in rewritten + assert "THE SOFTWARE IS PROVIDED" in rewritten + + +def test_rewrite_license_holder_no_copyright_line(): + from xtuner.tools.model_normalize.cli import _rewrite_license_holder + + text = "Some license body\nwithout a copyright line\n" + rewritten, replaced = _rewrite_license_holder(text) + assert not replaced + assert rewritten == text + + +def test_rewrite_license_holder_only_first_copyright_line(): + from xtuner.tools.model_normalize.cli import _rewrite_license_holder + + text = "Copyright (c) 2026 Zhipu AI\nsome line\nCopyright (c) 2030 Other\n" + rewritten, replaced = _rewrite_license_holder(text) + assert replaced + # Only the first Copyright line is rewritten; the second is left as-is. + assert rewritten == "Copyright 2025-2026 Shanghai AI Laboratory\nsome line\nCopyright (c) 2030 Other\n" + + +def test_apply_base_model_assets_copies_generation_config(tmp_path): + from xtuner.tools.model_normalize.cli import _apply_base_model_assets + + base_dir = tmp_path / "base" + base_dir.mkdir() + (base_dir / "generation_config.json").write_text(json.dumps({"a": 1})) + output = tmp_path / "output" + output.mkdir() + # A pre-existing generation_config is overwritten. + (output / "generation_config.json").write_text(json.dumps({"old": True})) + + _apply_base_model_assets(base_dir, output) + + assert json.loads((output / "generation_config.json").read_text()) == {"a": 1} + + +def test_apply_base_model_assets_rewrites_license(tmp_path): + from xtuner.tools.model_normalize.cli import _apply_base_model_assets + + base_dir = tmp_path / "base" + base_dir.mkdir() + (base_dir / "LICENSE").write_text(MIT_LICENSE) + output = tmp_path / "output" + output.mkdir() + + _apply_base_model_assets(base_dir, output) + + license_text = (output / "LICENSE").read_text() + assert "Copyright 2025-2026 Shanghai AI Laboratory" in license_text + assert "Zhipu AI" not in license_text + assert "Permission is hereby granted" in license_text + + +def test_apply_base_model_assets_none_is_noop(tmp_path): + from xtuner.tools.model_normalize.cli import _apply_base_model_assets + + output = tmp_path / "output" + output.mkdir() + (output / "generation_config.json").write_text(json.dumps({"keep": True})) + + # No exception, no change to output. + _apply_base_model_assets(None, output) + assert json.loads((output / "generation_config.json").read_text()) == {"keep": True} + assert not (output / "LICENSE").exists() + + +def test_apply_base_model_assets_missing_files_is_noop(tmp_path): + from xtuner.tools.model_normalize.cli import _apply_base_model_assets + + base_dir = tmp_path / "base" + base_dir.mkdir() + output = tmp_path / "output" + output.mkdir() + + _apply_base_model_assets(base_dir, output) + assert not (output / "generation_config.json").exists() + assert not (output / "LICENSE").exists() + + +def test_apply_base_model_assets_missing_dir_raises(tmp_path): + from xtuner.tools.model_normalize.cli import _apply_base_model_assets + + with pytest.raises(FileNotFoundError): + _apply_base_model_assets(tmp_path / "does-not-exist", tmp_path / "output") diff --git a/xtuner/tools/model_normalize/__init__.py b/xtuner/tools/model_normalize/__init__.py new file mode 100644 index 0000000000..315cda1b04 --- /dev/null +++ b/xtuner/tools/model_normalize/__init__.py @@ -0,0 +1 @@ +"""HF checkpoint normalization utilities for XTuner.""" diff --git a/xtuner/tools/model_normalize/__main__.py b/xtuner/tools/model_normalize/__main__.py new file mode 100644 index 0000000000..2f05ddc225 --- /dev/null +++ b/xtuner/tools/model_normalize/__main__.py @@ -0,0 +1,5 @@ +from .cli import main + + +if __name__ == "__main__": + main() diff --git a/xtuner/tools/model_normalize/cli.py b/xtuner/tools/model_normalize/cli.py new file mode 100644 index 0000000000..2a2760cc60 --- /dev/null +++ b/xtuner/tools/model_normalize/cli.py @@ -0,0 +1,196 @@ +"""Command-line entry point for HF checkpoint normalization. + +This tool intentionally does conversion/repacking only. Expensive full-model +validation is a separate release concern and is not run automatically. +""" + +from __future__ import annotations + +import argparse +import json +import os +import shutil +import tempfile +from pathlib import Path + + +def _add_common(parser: argparse.ArgumentParser) -> None: + parser.add_argument("--source", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--shard-size-gb", type=float, default=float(os.getenv("MODEL_NORMALIZE_SHARD_SIZE_GB", "4"))) + parser.add_argument("--overwrite", action="store_true") + parser.add_argument( + "--base-model-dir", + type=Path, + default=None, + help="directory with user-supplied LICENSE and generation_config.json " + "to apply after conversion; the Copyright line of LICENSE is rewritten " + "to 'Shanghai AI Laboratory'", + ) + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description="Normalize HF checkpoints for XTuner/HF inference") + sub = parser.add_subparsers(dest="command", required=True) + + rep = sub.add_parser("repack", help="re-shard an existing HF checkpoint") + _add_common(rep) + rep.set_defaults(handler=_run_repack) + + fp8 = sub.add_parser("to-fp8", help="convert BF16/FP16 weights to per-block FP8") + _add_common(fp8) + policy = fp8.add_mutually_exclusive_group(required=True) + policy.add_argument("--reference", type=Path, help="FP8 reference checkpoint") + policy.add_argument("--policy", choices=["heuristic"], help="explicit fallback quantization policy") + fp8.add_argument( + "--max-save-workers", + type=int, + default=int(os.getenv("MODEL_NORMALIZE_MAX_SAVE_WORKERS", "4")), + ) + fp8.add_argument("--block-size", type=int, default=128) + fp8.set_defaults(handler=_run_fp8) + return parser + + +def _prepare_output(source: Path, output: Path, overwrite: bool) -> None: + source = source.resolve() + output = output.resolve() + if not source.is_dir(): + raise FileNotFoundError(f"source directory does not exist: {source}") + if source == output: + raise ValueError("source and output must be different directories") + if output.exists(): + if not overwrite and any(output.iterdir()): + raise FileExistsError(f"output directory is not empty; pass --overwrite: {output}") + if overwrite: + shutil.rmtree(output) + output.mkdir(parents=True, exist_ok=True) + + +# Fixed copyright line stamped onto a user-supplied LICENSE. The year range is +# a literal, not derived from the input, so the release copyright is stable. +_LICENSE_COPYRIGHT_LINE = "Copyright 2025-2026 Shanghai AI Laboratory" + + +def _rewrite_license_holder(text: str) -> tuple[str, bool]: + """Rewrite the Copyright holder line in a LICENSE text. + + Only the first line that starts with ``Copyright`` (after stripping leading + whitespace) is replaced with the fixed Shanghai AI Laboratory copyright line; + the rest of the text (e.g. the MIT permission grant) is left untouched. + + Args: + text (str): Original LICENSE content. + + Returns: + tuple[str, bool]: The rewritten text and whether a Copyright line was + found and replaced. + """ + lines = text.splitlines(keepends=True) + replaced = False + for i, line in enumerate(lines): + if line.lstrip().startswith("Copyright"): + # Preserve the original line ending style of the replaced line. + newline = "\r\n" if line.endswith("\r\n") else ("\n" if line.endswith("\n") else "") + lines[i] = f"{_LICENSE_COPYRIGHT_LINE}{newline}" if newline else _LICENSE_COPYRIGHT_LINE + replaced = True + break + return "".join(lines), replaced + + +def _apply_base_model_assets(base_dir: Path | None, output: Path) -> None: + """Apply user-supplied assets from ``base_dir`` into ``output``. + + ``base_dir`` may contain a ``generation_config.json`` (copied verbatim, + overwriting any existing one) and a ``LICENSE`` (copied with its Copyright + line rewritten to Shanghai AI Laboratory). Missing files are skipped + silently; a ``None`` directory is a no-op so existing behavior is preserved. + + Args: + base_dir (Path | None): User-supplied base-model assets directory. + output (Path): Conversion product directory. + """ + if base_dir is None: + return + base_dir = base_dir.resolve() + if not base_dir.is_dir(): + raise FileNotFoundError(f"base-model directory does not exist: {base_dir}") + + gen_src = base_dir / "generation_config.json" + if gen_src.is_file(): + gen_dst = output / "generation_config.json" + if gen_dst.exists(): + print("[model_normalize] overwriting existing generation_config.json from base-model-dir") + shutil.copy2(gen_src, gen_dst) + + license_src = base_dir / "LICENSE" + if license_src.is_file(): + license_dst = output / "LICENSE" + if license_dst.exists(): + print("[model_normalize] overwriting existing LICENSE from base-model-dir") + text = license_src.read_text(encoding="utf-8") + rewritten, replaced = _rewrite_license_holder(text) + if not replaced: + print("[model_normalize] warning: no Copyright line found in base-model LICENSE; writing unchanged") + license_dst.write_text(rewritten, encoding="utf-8") + + +def _run_repack(args: argparse.Namespace) -> None: + try: + from .repack import repack + except ModuleNotFoundError as exc: + raise RuntimeError("repack requires the 'safetensors' package") from exc + _prepare_output(args.source, args.output, args.overwrite) + repack(args.source, args.output, shard_size_bytes=int(args.shard_size_gb * 1024**3)) + _apply_base_model_assets(args.base_model_dir, args.output) + + +def _reference_predicate(reference: Path): + with open(reference / "model.safetensors.index.json") as f: + keys = set(json.load(f)["weight_map"]) + suffix = "_scale_inv" + names = {key[: -len(suffix)] for key in keys if key.endswith(suffix)} + return names.__contains__ + + +def _run_fp8(args: argparse.Namespace) -> None: + try: + import torch + + from .fp8 import convert + from .heuristics import DEFAULT_QUANTIZE_PATTERNS, build_heuristic_predicate + from .repack import repack + except ModuleNotFoundError as exc: + raise RuntimeError("FP8 conversion requires torch, safetensors, and tqdm") from exc + if not torch.cuda.is_available(): + raise RuntimeError("FP8 conversion requires CUDA; use repack for CPU-only HF sharding") + if args.max_save_workers < 1: + raise ValueError("--max-save-workers must be positive") + if args.block_size < 1: + raise ValueError("--block-size must be positive") + _prepare_output(args.source, args.output, args.overwrite) + if args.reference is not None: + predicate = _reference_predicate(args.reference) + else: + predicate = build_heuristic_predicate() + print(f"[model_normalize] using heuristic policy ({len(DEFAULT_QUANTIZE_PATTERNS)} patterns)") + with tempfile.TemporaryDirectory(prefix="model_normalize_fp8_", dir=args.output.parent) as tmp: + staging = Path(tmp) + convert( + args.source, + staging, + should_quantize=predicate, + block_size=args.block_size, + max_workers=args.max_save_workers, + ) + repack(staging, args.output, shard_size_bytes=int(args.shard_size_gb * 1024**3)) + _apply_base_model_assets(args.base_model_dir, args.output) + + +def main(argv: list[str] | None = None) -> None: + args = build_parser().parse_args(argv) + args.handler(args) + + +if __name__ == "__main__": + main() diff --git a/xtuner/tools/model_normalize/examples/run_glm52.sh b/xtuner/tools/model_normalize/examples/run_glm52.sh new file mode 100644 index 0000000000..5952634985 --- /dev/null +++ b/xtuner/tools/model_normalize/examples/run_glm52.sh @@ -0,0 +1,44 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +if (($# < 1)); then + echo "usage: $0 {bf16|fp8}" >&2 + exit 2 +fi +MODE="$1" +SOURCE_DIR="${SOURCE_DIR:?set SOURCE_DIR to the source HF directory}" +OUTPUT_ROOT="${OUTPUT_ROOT:?set OUTPUT_ROOT to the output root}" +SHARD_SIZE_GB="${SHARD_SIZE_GB:-4}" +MAX_SAVE_WORKERS="${MAX_SAVE_WORKERS:-4}" +BASE_MODEL_DIR="${BASE_MODEL_DIR:-}" +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +RUNNER="${SCRIPT_DIR}/../run_model_normalize.sh" + +EXTRA_ARGS=() +if [[ -n "${BASE_MODEL_DIR}" ]]; then + EXTRA_ARGS+=(--base-model-dir "${BASE_MODEL_DIR}") +fi + +case "${MODE}" in + bf16) + exec bash "${RUNNER}" repack \ + --source "${SOURCE_DIR}" \ + --output "${OUTPUT_ROOT}/20_hf_bf16_mtp" \ + --shard-size-gb "${SHARD_SIZE_GB}" \ + "${EXTRA_ARGS[@]}" + ;; + fp8) + : "${REFERENCE_DIR:?set REFERENCE_DIR for FP8, or use the generic CLI with --policy heuristic}" + exec bash "${RUNNER}" to-fp8 \ + --source "${SOURCE_DIR}" \ + --output "${OUTPUT_ROOT}/20_hf_fp8_mtp" \ + --reference "${REFERENCE_DIR}" \ + --shard-size-gb "${SHARD_SIZE_GB}" \ + --max-save-workers "${MAX_SAVE_WORKERS}" \ + "${EXTRA_ARGS[@]}" + ;; + *) + echo "unsupported mode: ${MODE}; expected bf16 or fp8" >&2 + exit 2 + ;; +esac diff --git a/xtuner/tools/model_normalize/fp8.py b/xtuner/tools/model_normalize/fp8.py new file mode 100644 index 0000000000..c75f617ba7 --- /dev/null +++ b/xtuner/tools/model_normalize/fp8.py @@ -0,0 +1,285 @@ +r"""Per-block FP8 quantization of an HF-format safetensors checkpoint. + +The module exposes both a library API and a small CLI. The library form is +preferred when you want to drive the quantization policy from another script +(e.g. ``model_normalize.py``); the CLI form is kept for ad-hoc use. + +Library: + from hf_to_fp8 import convert + convert(source, target, should_quantize=lambda name: ...) + +CLI (single pattern): + python hf_to_fp8.py \ + 'model\.language_model\.layers\.\d+\.mlp\.experts\.\d+\.(gate|down|up)_proj\.weight$' + +CLI (multiple patterns, OR'd): + python hf_to_fp8.py \ + -p 'model\.language_model\.layers\.\d+\.self_attn\.[qkvo]_proj\.weight$' \ + -p 'model\.language_model\.layers\.\d+\.mlp\.experts\.\d+\.(gate|up|down)_proj\.weight$' +""" + +from __future__ import annotations + +import argparse +import json +import re +import shutil +from concurrent.futures import ProcessPoolExecutor +from pathlib import Path +from typing import Callable + +import torch +from safetensors import safe_open +from safetensors.torch import save_file +from tqdm import tqdm + + +FP8_TYPES = { + torch.float8_e4m3fn, + torch.float8_e5m2, + torch.float8_e4m3fnuz, + torch.float8_e5m2fnuz, +} + + +def convert( + source: Path, + target: Path, + should_quantize: Callable[[str], bool], + *, + block_size: int = 128, + float8_dtype: torch.dtype = torch.float8_e4m3fn, + max_workers: int = 4, +) -> None: + """Quantize an HF safetensors checkpoint to FP8 in-place into ``target``. + + Args: + source (Path): Source checkpoint directory containing + ``model.safetensors.index.json`` and the referenced shards. + target (Path): Destination directory. Will be created if missing. + Shard filenames are preserved from the source — call the repacker + afterwards if the source uses non-standard naming. + should_quantize (Callable[[str], bool]): Predicate that, given a + tensor name, returns whether that tensor should be quantized. + block_size (int): Block size for per-block scaling (default 128). + float8_dtype (torch.dtype): Target FP8 dtype (default ``e4m3fn``). + max_workers (int): Max parallel save workers. + """ + with open(source / "model.safetensors.index.json") as f: + index = json.load(f) + + target.mkdir(parents=True, exist_ok=True) + + original_weight_map = index.pop("weight_map") + + # Decide quantization status per module up front so we can emit a complete + # ``modules_to_not_convert`` list in the HF quantization_config. + modules_to_not_convert: set[str] = set() + for param_name in original_weight_map: + module_name = param_name.rsplit(".", 1)[0] + if not should_quantize(param_name): + modules_to_not_convert.add(module_name) + + quantization_config = { + "activation_scheme": "dynamic", + "fmt": "e4m3", + "quant_method": "fp8", + "scale_fmt": "ue8m0", + "weight_block_size": [block_size, block_size], + "modules_to_not_convert": sorted(modules_to_not_convert), + } + + executor = ProcessPoolExecutor(max_workers=max_workers) + pending = [] + new_weight_map: dict[str, str] = {} + + for filename in tqdm(sorted(set(original_weight_map.values()))): + filepath = source / filename + safetensor_fh = safe_open(filepath, framework="pt") + new_shard: dict[str, torch.Tensor] = {} + + for key in safetensor_fh.keys(): + new_weight_map[key] = filename + tensor = safetensor_fh.get_tensor(key) + if not should_quantize(key): + new_shard[key] = tensor + continue + + fp8_tensor, scale = per_block_quant_torch( + tensor.cuda(), block_size=block_size, float8_dtype=float8_dtype + ) + scale_key = f"{key}_scale_inv" + new_shard[key] = fp8_tensor.cpu() + new_shard[scale_key] = scale.cpu() + new_weight_map[scale_key] = filename + + pending.append(executor.submit(save_file, new_shard, target / filename)) + if len(pending) >= max_workers: + pending.pop(0).result() + + for future in pending: + future.result() + executor.shutdown() + _copy_others(source, target) + + index["weight_map"] = new_weight_map + with open(target / "model.safetensors.index.json", "w") as f: + json.dump(index, f, indent=2) + + with open(source / "config.json") as f: + hf_config = json.load(f) + hf_config["quantization_config"] = quantization_config + with open(target / "config.json", "w") as f: + json.dump(hf_config, f, indent=2) + + +def compile_union(patterns: list[str]) -> re.Pattern: + """Compile a list of regexes into a single union pattern. + + Each input is wrapped in a non-capturing group so caller-side groups don't + bleed across patterns, then OR'd together. + + Args: + patterns (list[str]): Regex strings to union. + + Returns: + re.Pattern: A compiled regex equivalent to ``(?:p1)|(?:p2)|...``. + """ + if len(patterns) == 1: + return re.compile(patterns[0]) + return re.compile("|".join(f"(?:{p})" for p in patterns)) + + +@torch.no_grad() +def per_block_quant_torch( + tensor: torch.Tensor, + block_size: int = 128, + float8_dtype: torch.dtype = torch.float8_e4m3fn, +) -> tuple[torch.Tensor, torch.Tensor]: + """Per-block FP8 quantization for 2D weights and 3D fused-expert weights. + + For 3D inputs of shape ``(num_experts, dim0, dim1)``, each expert slice is + quantized independently and the per-slice scale is stacked along a leading + dim, matching the layout vLLM / SGLang expect for fused MoE expert + weights. + + Args: + tensor (torch.Tensor): Input weight tensor (bf16/fp16/fp32), 2D or 3D. + block_size (int): Block size on each quantized dim. + float8_dtype (torch.dtype): Target FP8 dtype. + + Returns: + tuple[torch.Tensor, torch.Tensor]: ``(fp8_tensor, scale_inv)``. + """ + if tensor.dim() == 2: + return _per_block_quant_2d(tensor, block_size, float8_dtype) + if tensor.dim() == 3: + fp8_slices: list[torch.Tensor] = [] + scale_slices: list[torch.Tensor] = [] + for expert_idx in range(tensor.shape[0]): + t, s = _per_block_quant_2d(tensor[expert_idx], block_size, float8_dtype) + fp8_slices.append(t) + scale_slices.append(s) + return torch.stack(fp8_slices, dim=0), torch.stack(scale_slices, dim=0) + raise ValueError(f"per_block_quant_torch only supports 2D or 3D tensors, got shape {tuple(tensor.shape)}") + + +def _per_block_quant_2d(tensor: torch.Tensor, block_size: int, float8_dtype: torch.dtype): + dim0, dim1 = tensor.shape + tensor_pad = _pad_for_block(tensor, (0, 1), block_size) + dim0_pad, dim1_pad = tensor_pad.shape + tensor_pad = ( + tensor_pad.view(dim0_pad // block_size, block_size, dim1_pad // block_size, block_size) + .transpose(1, 2) + .reshape(-1, block_size * block_size) + ) + amax = tensor_pad.abs().amax(-1, True).to(torch.float64) + scales = (amax / torch.finfo(float8_dtype).max).to(torch.float32) + tensor_pad_scaled = tensor_pad.float() / scales + fp8 = _to_fp8_saturated(tensor_pad_scaled, float8_dtype) + fp8 = ( + fp8.view(dim0_pad // block_size, dim1_pad // block_size, block_size, block_size) + .transpose(1, 2) + .reshape(dim0_pad, dim1_pad) + ) + scales = scales.view(dim0_pad // block_size, dim1_pad // block_size) + return fp8[:dim0, :dim1], scales + + +def _pad_for_block(tensor: torch.Tensor, dims, block_size: int) -> torch.Tensor: + assert tensor.dim() == 2 + if isinstance(dims, int): + dims = (dims,) + dim1, dim2 = tensor.shape + dim1_aligned = _align(dim1, block_size) if 0 in dims else dim1 + dim2_aligned = _align(dim2, block_size) if 1 in dims else dim2 + return torch.nn.functional.pad(tensor, (0, dim2_aligned - dim2, 0, dim1_aligned - dim1)) + + +def _align(size: int, alignment: int) -> int: + return (1 + ((size - 1) // alignment)) * alignment + + +def _to_fp8_saturated(x: torch.Tensor, float8_dtype: torch.dtype) -> torch.Tensor: + # PyTorch's default cast to float8_e4m3fn / e5m2 does not saturate; we + # clamp first so that out-of-range values become +/- max instead of NaN. + if float8_dtype not in FP8_TYPES: + raise ValueError(f"Unsupported float8_dtype: {float8_dtype}") + max_value = torch.finfo(float8_dtype).max + return x.clamp(min=-max_value, max=max_value).to(float8_dtype) + + +def _copy_others(source: Path, target: Path) -> None: + # Tensor shards, index, and config.json are produced by the conversion + # itself, so they are excluded. Everything else at the top level is + # mirrored over — files via copy2 to preserve mtime, directories + # recursively (e.g. ``figs/`` in HF model dirs). Existing entries in + # target are overwritten so reruns are idempotent. + for entry in source.iterdir(): + if entry.name.endswith("safetensors"): + continue + if entry.name.startswith("."): + continue + if entry.name == "model.safetensors.index.json": + continue + if entry.name == "config.json": + continue + dst = target / entry.name + if entry.is_dir(): + shutil.copytree(entry, dst, dirs_exist_ok=True) + else: + shutil.copy2(entry, dst) + + +def _get_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="HF bf16 -> FP8 per-block quantization") + parser.add_argument("source", type=Path, help="source HF model directory") + parser.add_argument("target", type=Path, help="target HF model directory") + parser.add_argument( + "regex", + nargs="*", + type=str, + help="one or more regex patterns; tensor is quantized iff any pattern matches", + ) + parser.add_argument( + "--pattern", + "-p", + action="append", + default=[], + help="alternative to positional regex; can be repeated", + ) + args = parser.parse_args() + args.regex = list(args.regex) + list(args.pattern) + if not args.regex: + parser.error("at least one regex pattern is required (positional or via --pattern)") + return args + + +def main() -> None: + args = _get_args() + pattern = compile_union(args.regex) + convert(args.source, args.target, should_quantize=lambda k: pattern.search(k) is not None) + + +if __name__ == "__main__": + main() diff --git a/xtuner/tools/model_normalize/heuristics.py b/xtuner/tools/model_normalize/heuristics.py new file mode 100644 index 0000000000..443d7ae881 --- /dev/null +++ b/xtuner/tools/model_normalize/heuristics.py @@ -0,0 +1,77 @@ +"""Default heuristic rule set for FP8 quantization without a reference model. + +Rule of thumb +------------- +Quantize the *large matrix multiplications* in the transformer block; keep +everything else in the original dtype. Concretely this means: + +QUANTIZE + * ``self_attn.{q,k,v,o}_proj.weight`` (standard attention) + * ``mlp.{gate,up,down}_proj.weight`` (dense MLP) + * ``mlp.experts..{gate,up,down}_proj.weight`` (per-expert MoE) + * ``mlp.experts.gate_up_proj`` / ``mlp.experts.down_proj`` + (fused 3D MoE-expert weights) + * ``mlp.shared_expert.{gate,up,down}_proj.weight`` (MoE shared expert) + * ``linear_attn.{in_proj_qkv, in_proj_z, out_proj}.weight`` + (the wide projections in hybrid / Mamba-style blocks) + + These patterns apply both under ``model.language_model.layers..`` and + ``mtp.layers..`` (multi-token-prediction blocks share the language + model's block layout). + +KEEP IN ORIGINAL DTYPE (excluded by construction — no negative regex needed) + * all norms (``*norm*.weight``, ``*_layernorm.weight``, ``norm.bias``) + * MoE routers (``mlp.gate.weight``, ``mlp.shared_expert_gate.weight``) + * embeddings (``embed_tokens``, ``lm_head``, ``pos_embed``, ``patch_embed``, ``mtp.fc``) + * vision tower (``model.visual.*``) + * all ``*.bias`` tensors + * ``linear_attn`` control-flow tensors: ``A_log``, ``conv1d.weight``, + ``dt_bias``, ``in_proj_a.weight``, ``in_proj_b.weight``, ``norm.weight`` + + These are kept because they are either small (negligible memory win), + dtype-sensitive (norms / control-flow projections), or shape-incompatible + with block-128 FP8 quantization (1D biases, embeddings). + +This rule set was validated against the InternS2 / Qwen3-MoE FP8 reference +layouts produced by SGLang / vLLM tooling. If the source model uses different +naming for any of these concepts, fall back to the ``--reference`` mode. +""" + +from __future__ import annotations + +from typing import Callable + +from .fp8 import compile_union + + +DEFAULT_QUANTIZE_PATTERNS: list[str] = [ + # ---- Standard self-attention projections --------------------------------- + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.self_attn\.[qkvo]_proj\.weight$", + r"^mtp\.layers\.\d+\.self_attn\.[qkvo]_proj\.weight$", + # ---- Dense MLP (non-MoE) ------------------------------------------------- + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.mlp\.(?:gate|up|down)_proj\.weight$", + # ---- MoE per-expert linears (unfused) ----------------------------------- + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.mlp\.experts\.\d+\.(?:gate|up|down)_proj\.weight$", + r"^mtp\.layers\.\d+\.mlp\.experts\.\d+\.(?:gate|up|down)_proj\.weight$", + # ---- MoE fused-expert linears (3D, single tensor per layer) ------------- + # Note: no trailing ``.weight`` — these names come from the fused + # representation used by vLLM / SGLang fused MoE kernels. + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.mlp\.experts\.(?:gate_up_proj|down_proj)$", + # ---- MoE shared expert (always-on dense path) --------------------------- + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.mlp\.shared_expert\.(?:gate|up|down)_proj\.weight$", + r"^mtp\.layers\.\d+\.mlp\.shared_expert\.(?:gate|up|down)_proj\.weight$", + # ---- Linear-attention wide projections ---------------------------------- + r"^(?:model\.)?(?:language_model\.)?layers\.\d+\.linear_attn\.(?:in_proj_qkv|in_proj_z|out_proj)\.weight$", + r"^mtp\.layers\.\d+\.linear_attn\.(?:in_proj_qkv|in_proj_z|out_proj)\.weight$", +] + + +def build_heuristic_predicate() -> Callable[[str], bool]: + """Return a predicate matching ``DEFAULT_QUANTIZE_PATTERNS``. + + Returns: + Callable[[str], bool]: ``True`` iff the tensor name should be FP8-quantized. + """ + pattern = compile_union(DEFAULT_QUANTIZE_PATTERNS) + return lambda name: pattern.search(name) is not None + diff --git a/xtuner/tools/model_normalize/repack.py b/xtuner/tools/model_normalize/repack.py new file mode 100644 index 0000000000..0053105471 --- /dev/null +++ b/xtuner/tools/model_normalize/repack.py @@ -0,0 +1,197 @@ +"""Repack a safetensors model directory into HF-standard ~4GB shards. + +Training engines may save safetensors with engine-specific filenames +(e.g. ``model-language-0001-fused-save_rank0.safetensors``). HF inference +backends expect the standard ``model-{i:05d}-of-{n:05d}.safetensors`` +layout. This script rewrites shards in place into the standard layout, +preserves all non-shard files (config, tokenizer, etc.), and emits a fresh +``model.safetensors.index.json``. + +Library: + from repack_hf import repack + repack(source, target, shard_size_bytes=4 * 1024**3) + +CLI: + python repack_hf.py [--shard-size-gb 4] +""" + +from __future__ import annotations + +import argparse +import json +import shutil +from dataclasses import dataclass +from pathlib import Path + +from safetensors import safe_open +from safetensors.torch import save_file +from tqdm import tqdm + + +# safetensors dtype string -> bytes per element. Covers the dtypes we emit. +_DTYPE_BYTES = { + "BOOL": 1, + "U8": 1, + "I8": 1, + "F8_E4M3": 1, + "F8_E5M2": 1, + "I16": 2, + "U16": 2, + "F16": 2, + "BF16": 2, + "I32": 4, + "U32": 4, + "F32": 4, + "I64": 8, + "U64": 8, + "F64": 8, +} + + +@dataclass +class _TensorEntry: + name: str + source_file: str + nbytes: int + + +def repack(source: Path, target: Path, shard_size_bytes: int = 4 * 1024**3) -> None: + """Repack ``source`` safetensors into ~``shard_size_bytes`` standard shards. + + Args: + source (Path): Source model directory. + target (Path): Output directory; created if missing. Must differ from + ``source`` (we never rewrite shards in place to keep the source + recoverable on failure). + shard_size_bytes (int): Soft upper bound for each output shard. A + tensor larger than this bound goes into its own shard. + """ + source = source.resolve() + target = target.resolve() + if source == target: + raise ValueError("source and target must differ") + target.mkdir(parents=True, exist_ok=True) + + entries = _scan_tensors(source) + plan = _plan_shards(entries, shard_size_bytes) + num_shards = len(plan) + + new_weight_map: dict[str, str] = {} + total_size = 0 + + for shard_idx, shard in enumerate(tqdm(plan, desc="writing shards")): + shard_name = f"model-{shard_idx + 1:05d}-of-{num_shards:05d}.safetensors" + # Group tensors in this shard by source file so we open each source + # shard at most once per output shard. + by_source: dict[str, list[str]] = {} + for entry in shard: + by_source.setdefault(entry.source_file, []).append(entry.name) + + tensors = {} + for src_file, names in by_source.items(): + with safe_open(source / src_file, framework="pt") as f: + for name in names: + tensors[name] = f.get_tensor(name) + + save_file(tensors, target / shard_name) + + for entry in shard: + new_weight_map[entry.name] = shard_name + total_size += entry.nbytes + + _copy_non_shard_files(source, target) + + index = { + "metadata": {"total_size": total_size}, + "weight_map": new_weight_map, + } + with open(target / "model.safetensors.index.json", "w") as f: + json.dump(index, f, indent=2) + + +def _scan_tensors(source: Path) -> list[_TensorEntry]: + # Prefer the index when available — it gives a stable ordering and avoids + # scanning files we never intended to read. + index_path = source / "model.safetensors.index.json" + if index_path.exists(): + with open(index_path) as f: + weight_map: dict[str, str] = json.load(f)["weight_map"] + files_to_keys: dict[str, list[str]] = {} + for name, file in weight_map.items(): + files_to_keys.setdefault(file, []).append(name) + else: + files_to_keys = {p.name: None for p in sorted(source.glob("*.safetensors"))} # type: ignore[assignment] + + entries: list[_TensorEntry] = [] + for src_file, names in files_to_keys.items(): + with safe_open(source / src_file, framework="pt") as f: + keys = names if names is not None else list(f.keys()) + for key in keys: + slice_ = f.get_slice(key) + shape = slice_.get_shape() + dtype = slice_.get_dtype() + numel = 1 + for d in shape: + numel *= d + nbytes = numel * _DTYPE_BYTES[dtype] + entries.append(_TensorEntry(name=key, source_file=src_file, nbytes=nbytes)) + return entries + + +def _plan_shards(entries: list[_TensorEntry], shard_size_bytes: int) -> list[list[_TensorEntry]]: + # Greedy bin-packing. We preserve the input order (which mirrors the + # source index order) so that produced shards stay roughly contiguous in + # the original layer order — friendlier for streamed loading. + shards: list[list[_TensorEntry]] = [] + current: list[_TensorEntry] = [] + current_size = 0 + for entry in entries: + if current and current_size + entry.nbytes > shard_size_bytes: + shards.append(current) + current = [] + current_size = 0 + current.append(entry) + current_size += entry.nbytes + if current: + shards.append(current) + return shards + + +def _copy_non_shard_files(source: Path, target: Path) -> None: + for file in source.iterdir(): + if file.name.endswith(".safetensors"): + continue + if file.name == "model.safetensors.index.json": + continue + if file.name.startswith("."): + continue + dst = target / file.name + if file.is_dir(): + if not dst.exists(): + shutil.copytree(file, dst) + else: + shutil.copy(file, dst) + + +def _get_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Repack safetensors into 4GB-shard HF layout") + parser.add_argument("source", type=Path, help="source model directory") + parser.add_argument("target", type=Path, help="target model directory") + parser.add_argument( + "--shard-size-gb", + type=float, + default=4.0, + help="soft upper bound per output shard, in GiB (default: 4)", + ) + return parser.parse_args() + + +def main() -> None: + args = _get_args() + shard_size_bytes = int(args.shard_size_gb * 1024**3) + repack(args.source, args.target, shard_size_bytes=shard_size_bytes) + + +if __name__ == "__main__": + main() + diff --git a/xtuner/tools/model_normalize/run_model_normalize.sh b/xtuner/tools/model_normalize/run_model_normalize.sh new file mode 100644 index 0000000000..a3384eaf2c --- /dev/null +++ b/xtuner/tools/model_normalize/run_model_normalize.sh @@ -0,0 +1,9 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../" && pwd)" +export PYTHONPATH="${REPO_ROOT}:${PYTHONPATH:-}" + +command -v python3 >/dev/null || { echo "python3 is required" >&2; exit 2; } +exec python3 -m xtuner.tools.model_normalize "$@"