Files
realAsmaandClaude Opus 4.8 196c091027 [1/N] Refactor llm_qat example: YAML configs + ModelOptArgParser (#1172)
### What does this PR do?

Type of change: new example

Refactors `examples/llm_qat` from a monolithic launch/script flow into a
modular, config-driven Hugging Face QAT/QAD workflow. The high-level
user flow is now:

1. Quantize a base model with a ModelOpt PTQ recipe.
2. Train or evaluate the quantized checkpoint with QAT, QAD, LoRA QAT,
QLoRA, or fine-tuning configs.
3. Export the trained checkpoint for deployment.

Highlights:

- Replaces the legacy `examples/llm_qat/launch.sh` +
`examples/llm_qat/main.py` path with separate `quantize.py` and
`train.py` entrypoints.
- Adds YAML-driven argument parsing through `ModelOptArgParser`,
including `--config <yaml>` defaults, CLI overrides, and generated
`examples/llm_qat/ARGUMENTS.md`.
- Adds declarative configs under `examples/llm_qat/configs/` for
training modes, dataset blends, and Accelerate backends.
- Adds `dataset_utils.py` for weighted multi-source dataset blending,
Hugging Face streaming, local dataset loading, distributed rank-aware
loading, tokenization caching, pre-tokenization, chat templating, and
assistant-token label masking.
- Moves the Hugging Face QAD flow into the shared trainer path through
`DistillArguments`, `QADTrainer`, and teacher-model distillation kwargs.
- Updates `QuantizationArguments` to prefer recipe paths via `--recipe`,
while keeping legacy `--quant_cfg` available with deprecation warnings
for in-trainer quantization.
- Updates `QATTrainer` handling for pre-quantized checkpoints, FSDP2
TensorQuantizer buffers, and recipe-resolved PTQ configs.
- Adds the `general/ptq/int4_blockwise_weight_only` recipe and refreshes
docs for NVFP4, FP8, INT4, FSDP2, DDP, DeepSpeed, QLoRA, and
LLaMA-Factory integration.
- Adds focused parser, dataset tokenization, assistant-mask, and example
workflow coverage.

### Usage

From the repo root:

```sh
cd examples/llm_qat

# 1. Quantize
python quantize.py \
  --model_name_or_path Qwen/Qwen3-8B \
  --dataset_config configs/dataset/blend.yaml \
  --recipe general/ptq/nvfp4_default-kv_fp8 \
  --output_dir qwen3-8b-quantized

# 2. QAT train
accelerate launch --config-file configs/accelerate/fsdp2.yaml train.py \
  --config configs/train/qat_nvfp4.yaml \
  --model_name_or_path qwen3-8b-quantized \
  --output_dir qwen3-8b-qat-nvfp4

# 3. QAD train
accelerate launch --config-file configs/accelerate/fsdp2.yaml train.py \
  --config configs/train/qad_nvfp4.yaml \
  --model_name_or_path qwen3-8b-quantized \
  --teacher_model Qwen/Qwen3-8B \
  --output_dir qwen3-8b-qad-nvfp4
```

Dataset blends can be pre-tokenized and cached before training:

```sh
cd examples/llm_qat
python dataset_utils.py \
  --dataset_config configs/dataset/blend.yaml \
  --model_name_or_path Qwen/Qwen3-8B
```

### Testing

Focused coverage added or updated:

- `tests/unit/torch/opt/plugins/test_modelopt_arg_parser.py`
- `tests/examples/llm_qat/test_dataset_tokenization.py`
- `tests/examples/llm_qat/test_assistant_mask.py`
- `tests/examples/llm_qat/test_llm_qat.py`

Recorded validation:

- [x] `pytest tests/unit/torch/opt/plugins/test_modelopt_arg_parser.py`
- [x] `pytest
tests/examples/llm_qat/test_llm_qat.py::test_dataset_utils_pretokenize`
- [x] `pre-commit run --all-files`
- [ ] `pytest tests/examples/llm_qat/test_llm_qat.py` full GPU/backend
suite

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: No for the legacy
`examples/llm_qat` CLI/file layout (`main.py`, `launch.sh`, and FSDP1
config are removed); library `quant_cfg` usage remains available but is
deprecated for this workflow.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: Yes.
`train.py`/`simple_qat_train.py` retain the upstream Alpaca attribution
where applicable; `examples/llm_qat/requirements.txt` switches
`tensorboardX` to `tensorboard`.
- Did you write any new necessary tests?: Yes. Parser, dataset
tokenization, assistant masking, pre-tokenization, and QAT/QAD workflow
tests were added or updated.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
Yes.
- Did you get Claude approval on this PR?: N/A in this description
update.

### Additional Information

This PR intentionally changes the `llm_qat` example surface. Existing
users should move from `examples/llm_qat/main.py` and
`examples/llm_qat/launch.sh` to `examples/llm_qat/quantize.py`,
`examples/llm_qat/train.py`, and the YAML configs under
`examples/llm_qat/configs/`.

---------

Signed-off-by: realAsma <akuriparambi@nvidia.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-03 01:25:02 +00:00

149 lines
5.6 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import argparse
import torch
import torch.nn as nn
from dataset_utils import build_blend_dataset, load_blend_config
from torch.optim import AdamW
from torch.utils.data import DataLoader
from tqdm import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer
import modelopt.torch.opt as mto
import modelopt.torch.quantization as mtq
from modelopt.recipe import ModelOptPTQRecipe, load_recipe
def get_dataloader(args, tokenizer):
config = load_blend_config("configs/dataset/blend.yaml")
ds = build_blend_dataset(config, tokenizer, args.max_length)
train_dataset = ds["train"]
if 0 < args.train_size < len(train_dataset):
train_dataset = train_dataset.select(range(args.train_size))
calib_dataset = ds["eval"]
if 0 < args.calib_size < len(calib_dataset):
calib_dataset = calib_dataset.select(range(args.calib_size))
def collate_fn(batch):
return {
"input_ids": torch.tensor([item["input_ids"] for item in batch]),
"attention_mask": torch.tensor([item["attention_mask"] for item in batch]),
"labels": torch.tensor([item["labels"] for item in batch]),
}
train_dataloader = DataLoader(
train_dataset, batch_size=args.batch_size, shuffle=True, collate_fn=collate_fn
)
calib_dataloader = DataLoader(
calib_dataset, batch_size=args.batch_size, shuffle=False, collate_fn=collate_fn
)
return train_dataloader, calib_dataloader
def train(model, optimizer, train_dataloader, tokenizer, epochs, output_dir, device):
for epoch in (pbar := tqdm(range(epochs))):
pbar.set_description(f"Epoch {epoch + 1}/{epochs}")
for batch in (pbar_batch := tqdm(train_dataloader)):
inputs = batch["input_ids"].to(device)
attention_mask = batch["attention_mask"].to(device)
outputs = model(
input_ids=inputs, attention_mask=attention_mask, labels=batch["labels"].to(device)
)
loss = outputs.loss
optimizer.zero_grad()
loss.backward()
optimizer.step()
pbar_batch.set_description(f"loss: {loss.item():.4f}")
print(f"Epoch {epoch + 1} completed | Loss: {loss.item():.4f}")
if output_dir:
tokenizer.save_pretrained(output_dir)
model.save_pretrained(output_dir)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="QAT Training Script")
# Data paths
parser.add_argument("--model-path", type=str, required=True, help="Path to the model")
parser.add_argument("--train-size", type=int, default=2048, help="Train size")
parser.add_argument("--calib-size", type=int, default=512, help="Calib size")
parser.add_argument("--max-length", type=int, default=2048, help="Max length")
# Hyperparameters
parser.add_argument("--batch-size", type=int, default=4, help="Batch size")
parser.add_argument("--epochs", type=int, default=2, help="Number of epochs")
parser.add_argument("--lr", type=float, default=1e-5, help="Learning rate")
parser.add_argument(
"--recipe",
type=str,
default="general/ptq/nvfp4_default-kv_fp8",
help="Path to a quantization recipe YAML (built-in or custom)",
)
# Reproducibility
parser.add_argument("--seed", type=int, default=42, help="Random seed")
parser.add_argument("--print-freq", type=int, default=100, help="Print frequency")
parser.add_argument(
"--output-dir", type=str, default="qat_model", help="Directory to save the checkpoints"
)
return parser.parse_args()
def main() -> None:
args = parse_args()
# Enable automatic save/load of modelopt state huggingface checkpointing
# modelopt state will be saved automatically to "modelopt_state.pt"
mto.enable_huggingface_checkpointing()
# Load model and initialize loss
model = AutoModelForCausalLM.from_pretrained(args.model_path).cuda()
tokenizer = AutoTokenizer.from_pretrained(args.model_path)
# Get dataloaders
train_dataloader, calib_dataloader = get_dataloader(args, tokenizer)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Calibrate the model
def calibrate(m: nn.Module):
for batch in calib_dataloader:
m(
input_ids=batch["input_ids"].to(device),
attention_mask=batch["attention_mask"].to(device),
)
# Load recipe and quantize the model
recipe = load_recipe(args.recipe)
if not isinstance(recipe, ModelOptPTQRecipe):
raise ValueError(f"Expected PTQ recipe, but got {type(recipe).__name__} from {args.recipe}")
model = mtq.quantize(model, recipe.quantize, calibrate)
# Initialize optimizer
optimizer = AdamW(model.parameters(), lr=args.lr)
# Train the model
model.train()
model.to(device)
train(model, optimizer, train_dataloader, tokenizer, args.epochs, args.output_dir, device)
if __name__ == "__main__":
main()