Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
142 changes: 142 additions & 0 deletions docs/en/model_normalize.md
Original file line number Diff line number Diff line change
@@ -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 <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
<OUTPUT_ROOT>/
├── 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.
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
191 changes: 191 additions & 0 deletions tests/tools/test_model_normalize_cli.py
Original file line number Diff line number Diff line change
@@ -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")
1 change: 1 addition & 0 deletions xtuner/tools/model_normalize/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
"""HF checkpoint normalization utilities for XTuner."""
5 changes: 5 additions & 0 deletions xtuner/tools/model_normalize/__main__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
from .cli import main


if __name__ == "__main__":
main()
Loading
Loading