[2/N] Simplify KDTrainer and enhance ModelOptHFTrainer (#1191)

## Summary

This PR simplifies the HuggingFace knowledge distillation trainer and
enhances the base `ModelOptHFTrainer` with Liger fused loss,
per-parameter learning rates, and training utilities.

### Model-agnostic Liger kernel fused loss

Adds custom Liger kernel integration in `ModelOptHFTrainer` that extends
HuggingFace's built-in support in three ways:

1. **Model-agnostic**: Works with any causal LM that has an `lm_head`,
unlike HF's Liger which only supports [a fixed set of model
architectures](https://github.com/linkedin/Liger-Kernel/blob/main/src/liger_kernel/transformers/monkey_patch.py).
2. **DeepSpeed ZeRO-3 support**: HF's Liger integration only works with
FSDP. ModelOpt adds distributed param gathering for DeepSpeed ZeRO-3 and
DDP as well.
3. **KD loss support**: `KDTrainer` extends fused loss to knowledge
distillation via `LigerFusedLinearJSD` for fused lm_head +
Jensen-Shannon divergence.

#### Liger kernel memory sweep (Qwen3-1.7B, 2×H100 FSDP2, NVFP4+FP8_KV)

Max per-GPU batch size before OOM at each sequence length:

**QAT (no teacher)**

| Seq Length | 512 | 1024 | 2048 | 4096 | 8192 | 16384 |
|------------|-----|------|------|------|------|-------|
| **Liger**  | 16  | 16   | 16   | 16   | 8    | 4     |
| **No Liger** | 16 | 16  | 8    | 4    | 2    | OOM   |

**QAD (with teacher)**

| Seq Length | 512 | 1024 | 2048 | 4096 | 8192 | 16384 |
|------------|-----|------|------|------|------|-------|
| **Liger**  | 16  | 16   | 8    | 4    | 2    | 1     |
| **No Liger** | 8 | 4   | 2    | 1    | OOM  | OOM   |

Liger fused loss enables **2-4× larger batch sizes** at long context
lengths by avoiding the materialization of the full logit tensor.

### ModelOptHFTrainer enhancements

- `ModelOptTrainerArguments` with `--trainable_params`,
`--frozen_params`, `--lr_config`, `--save_dtype`, and `--manual_gc`
flags
- Per-parameter learning rate support via YAML config (`lr_config`)
- `_prepare_model` and `_update_config_json_dtype` promoted to base
class

### KDTrainer simplification + fix

Removes `mtd.convert()` and the `DistillationModel` in-place class-swap
for the HF path. The teacher model now lives directly on the trainer and
is forwarded explicitly inside `compute_kd_loss_func`. This eliminates:

- `mtd.convert()` in-place class swap and DynamicModule wrapping
- Forward hooks for capturing intermediate outputs
- `hide_teacher_model` / `hide_loss_modules` context managers for
checkpointing
- Deferred initialization branching (FSDP2 vs DDP/DeepSpeed)
- `save_model` and `QADTrainer._quantize_model` overrides

**Bug fix**: The previous `DistillationModel`/`mtd.convert()` approach
did not support CPU RAM-efficient loading for QAD. The teacher model had
to be fully loaded on GPU before wrapping, which doubled peak memory
during initialization. The new approach loads the teacher lazily on the
trainer, enabling standard HF device-map and low-cpu-mem-usage loading.

Only logit-level distillation is supported for the HF path. The core
`DistillationModel`/`mtd.convert()` API remains for Megatron and
advanced intermediate-layer distillation use cases.

## Test plan

- [x] `pytest tests/unit/torch/distill/` (29 passed)
- [x] `pytest tests/unit/torch/opt/plugins/test_hf_patching.py` (2
passed)
- [x] `pytest tests/unit/torch/opt/plugins/test_lr_config.py`
- [x] Pre-commit hooks pass
- [ ] GPU example tests: `pytest tests/examples/llm_qat/` (QAT, QAD,
LoRA QAT, QLoRA)
- [ ] GPU distill example: `pytest tests/examples/llm_distill/`

🤖 Generated with [Claude Code](https://claude.com/claude-code)

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

## Release Notes

* **New Features**
* Added Liger fused loss support in `ModelOptHFTrainer` for distributed
causal language models with JSD distillation loss support.
* Introduced `ModelOptTrainerArguments` with new training CLI flags:
per-parameter learning rates via YAML, parameter freezing, and manual
garbage collection.
* Simplified knowledge distillation trainer with logit-level
distillation support.

* **Documentation**
* Updated example configurations and documentation with new training
options and defaults.
  * Added learning rate configuration example guide.

* **Tests**
* Added test coverage for distillation training and per-parameter
optimizer configuration.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

Signed-off-by: realAsma <akuriparambi@nvidia.com>
This commit is contained in:
realAsma
2026-06-05 15:47:43 +00:00
committed by GitHub
parent 115cae2584
commit 433b549cd8
20 changed files with 1301 additions and 251 deletions
+3
View File
@@ -64,6 +64,9 @@ Changelog
**New Features**
- Add model-agnostic `Liger kernel <https://github.com/linkedin/Liger-Kernel>`_ fused loss support in ``ModelOptHFTrainer`` for any HuggingFace causal LM, with distributed param gathering for FSDP2, DeepSpeed ZeRO-3, and DDP. Extends HuggingFace's built-in Liger integration which is limited to `a fixed set of model architectures <https://github.com/linkedin/Liger-Kernel/blob/main/src/liger_kernel/transformers/monkey_patch.py>`_, FSDP only, and CrossEntropy loss. ModelOpt additionally supports Liger fused KD loss (JSD) for knowledge distillation.
- Add ``ModelOptTrainerArguments`` to ``ModelOptHFTrainer`` with ``--trainable_params``, ``--frozen_params``, ``--lr_config``, and ``--manual_gc`` flags. Add per-parameter learning rate support via YAML config. Saved checkpoints preserve the model's original dtype in ``config.json``.
- Simplify ``KDTrainer`` for HuggingFace knowledge distillation: remove ``mtd.convert()`` class-swap in favor of explicit teacher forwarding with logit-level distillation support.
- Support full Transformer Engine spec for Minitron pruning (``mcore_minitron``). Now we no longer need to use custom ModelOpt spec. Note that this does not affect the usage of the pruning workflow but makes pruning slightly faster and may result in slightly different pruned model because of different kernel and numerics.
- Add end-to-end tutorial for Minitron pruning + distillation + quantization + evaluation + vLLM deployment for Nemotron-Nano-9B-v2 → Pruned 7B along with data blend preparation steps (and ablation study). See `examples/pruning/minitron/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/pruning/minitron/>`_ for details.
- Add Puzzletron - a new algorithm for heterogeneous pruning of LLM and VLM models. See `examples/puzzletron/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/puzzletron>`_ for more details.
+2 -8
View File
@@ -27,7 +27,7 @@ from transformers import AutoTokenizer
from trl import SFTTrainer
import modelopt.torch.opt as mto
from modelopt.torch.distill.plugins.huggingface import KDTrainer, LMLogitsLoss
from modelopt.torch.distill.plugins.huggingface import KDTrainer
logger = get_logger(__name__, log_level="INFO")
@@ -115,12 +115,6 @@ def train():
model_args.teacher_name_or_path, dtype=torch.bfloat16 if training_args.bf16 else None
)
# Distillation configuration
kd_config = {
"teacher_model": teacher_model,
"criterion": LMLogitsLoss(),
}
# Fix problematic settings that logger.info excessive warnings
model.generation_config.temperature = None
model.generation_config.top_p = None
@@ -129,7 +123,7 @@ def train():
trainer = KDSFTTrainer(
model,
training_args,
distill_config=kd_config,
distill_args={"teacher_model": teacher_model},
train_dataset=dset_train,
eval_dataset=dset_eval,
formatting_func=lambda sample: _format_smoltalk_chat_template(sample, tokenizer),
+10 -3
View File
@@ -19,8 +19,10 @@
| Argument | Type | Default | Description |
|----------|------|---------|-------------|
| `--distill` | `bool` | `False` | Enable training with knowledge distillation. |
| `--teacher_model` | `str` | `None` | The name or path of the teacher model to use for distillation. |
| `--teacher_model` | `str` | `None` | The name or path of the teacher model. |
| `--criterion` | `str` | `"logits_loss"` | Distillation loss criterion. Currently only 'logits_loss' is supported. |
| `--temperature` | `float` | `1.0` | Softmax temperature for softening logits in KD loss. Used by both standard and Liger KD loss. |
| `--liger_jsd_beta` | `float` | `0.0` | JSD beta coefficient in [0, 1]. 0=forward KL, 1=reverse KL. Only used when --use_liger_kernel is enabled. |
## DataArguments
@@ -40,7 +42,8 @@
| Argument | Type | Default | Description |
|----------|------|---------|-------------|
| `--model_name_or_path` | `str` | `"Qwen/Qwen3-8B"` | HuggingFace model name or local path to the base model to quantize/train. |
| `--model_max_length` | `int` | `4096` | Maximum sequence length. Sequences will be right-padded (and possibly truncated). |
| `--model_max_length` | `int` | `8192` | Maximum sequence length. Sequences will be right-padded (and possibly truncated). |
| `--attn_implementation` | `str` | `None` | Attention implementation: 'flash_attention_2', 'flash_attention_3', 'sdpa', or 'eager'. |
## QuantizeArguments
@@ -59,5 +62,9 @@ Extends [HuggingFace TrainingArguments](https://huggingface.co/docs/transformers
| Argument | Type | Default | Description |
|----------|------|---------|-------------|
| `--cache_dir` | `str` | `None` | |
| `--trainable_params` | `list[str]` | `None` | Glob patterns (fnmatch) for parameters that should be trainable. All other parameters will be frozen. Mutually exclusive with frozen_params. |
| `--frozen_params` | `list[str]` | `None` | Glob patterns (fnmatch) for parameters that should be frozen. Mutually exclusive with trainable_params. |
| `--lr_config` | `str` | `None` | Path to a YAML file mapping fnmatch patterns to optimizer kwargs (e.g. lr, weight_decay). First matching pattern wins per parameter. See examples/llm_qat/configs/train/lr_config_example.yaml. |
| `--manual_gc` | `bool` | `False` | Run `gc.collect()` before each training/prediction step to work around GPU memory leaks during QAT/distillation. |
| `--liger_ce_label_smoothing` | `float` | `0.0` | Label smoothing for Liger fused CE loss. Only used when --use_liger_kernel is enabled. |
| `--lora` | `bool` | `False` | Whether to add LoRA (Low-Rank Adaptation) adapter before training. When using real quantization, the LoRA adapter must be set, as quantized weights will be frozen during training. |
+32 -16
View File
@@ -24,13 +24,21 @@ For background on how QAT enables low-precision accuracy recovery, see the [QAT/
### Prerequisites
Please refer to [llm_ptq/README.md](../llm_ptq/README.md#pre-requisites) for prerequisites.
Please refer to [llm_ptq/README.md](../llm_ptq/README.md#pre-requisites) for container
recommendations and base ModelOpt installation guidance. For this QAT/QAD example,
install the Hugging Face dependencies and the example-specific requirements:
```bash
pip install -U nvidia-modelopt[hf]
pip install -r examples/llm_qat/requirements.txt
```
The Qwen3-8B example below requires a minimum of **2 x 80GB GPUs**.
## Run End-to-End QAT/QAD Example
All arguments can be specified via YAML config, CLI flags, or both (CLI overrides YAML). See [ARGUMENTS.md](ARGUMENTS.md) for the full per-script argument reference, or run any script with `--help`.
All arguments can be set via YAML, CLI, or both (CLI overrides YAML). See
[ARGUMENTS.md](ARGUMENTS.md), `--help`, and [Configuration](#advanced-configuration).
### QAT
@@ -141,22 +149,23 @@ trainer.train()
trainer.save_model()
```
`QADTrainer` extends `QATTrainer` with distillation:
`QADTrainer` extends `QATTrainer` with distillation. Pass the teacher model and a `DistillArguments` instance:
```python
from modelopt.torch.distill.plugins.huggingface import LMLogitsLoss
from modelopt.torch.distill.plugins.huggingface import DistillArguments
from modelopt.torch.quantization.plugins.transformers_trainer import QADTrainer
distill_config = {
"teacher_model": teacher_model,
"criterion": LMLogitsLoss(),
}
distill_args = DistillArguments(
distill=True,
teacher_model="Qwen/Qwen3-8B",
criterion="logits_loss",
)
trainer = QADTrainer(
model=model, # pre-quantized model
processing_class=tokenizer,
args=training_args,
distill_config=distill_config,
distill_args=distill_args,
**data_module,
)
trainer.train()
@@ -188,13 +197,20 @@ See [custom calibration](https://nvidia.github.io/Model-Optimizer/guides/_pytorc
### Supported Quantization Formats
| Format | Precision | Recipe | Use Case |
|--------|-----------|--------|----------|
| **NVFP4** | W4A4 + FP8 KV | `general/ptq/nvfp4_default-kv_fp8` | Maximum compression for Blackwell GPUs |
| **FP8** | W8A8 + FP8 KV | `general/ptq/fp8_default-fp8_kv` | Balanced speed and accuracy |
| **INT4** weight-only | W4A16 | `general/ptq/int4_blockwise_weight_only` | Deployable on all Ampere or later GPUs |
Built-in recipes support full-model, partial-layer, and mixed-precision quantization. Common entry points:
> **NVFP4** uses 4-bit FP weights and activations (E2M1 with FP8 dynamic scales) plus FP8 KV cache. Partial variants are available for quantizing only specific layers (e.g., MLP-only, MoE experts-only) — see [`modelopt_recipes/general/ptq/`](../../modelopt_recipes/general/ptq/) for all options.
| Format | Precision | Example Recipe | Use Case |
|--------|-----------|----------------|----------|
| **NVFP4** | W4A4 + FP8 KV | `general/ptq/nvfp4_default-kv_fp8` | FP4 compute and compression on Blackwell GPUs |
| **FP8** | W8A8 + FP8 KV | `general/ptq/fp8_default-kv_fp8` | Near-BF16 accuracy on Hopper or later GPUs |
| **INT4** weight-only | W4A16 | `general/ptq/int4_blockwise_weight_only` | Deployable on all Ampere or later GPUs |
| **Partial / mixed** | Pattern-specific | `general/ptq/nvfp4_mlp_only-kv_fp8` | Quantize selected layers or combine precisions |
> Recipes can target different layers or GEMMs with different precisions, such as NVFP4
> for MLP/MoE GEMMs and FP8 for attention GEMMs or KV cache. See
> [`modelopt_recipes/general/ptq/`](../../modelopt_recipes/general/ptq/) and
> [`modelopt_recipes/configs/ptq/`](../../modelopt_recipes/configs/ptq/) for built-in
> options and reusable recipe units.
### Supported Backends
@@ -271,7 +287,7 @@ Common layer class names:
</details>
<details>
<details id="advanced-configuration">
<summary><b>Configuration</b></summary>
There are two types of configs:
+20 -4
View File
@@ -20,7 +20,11 @@ from dataclasses import field
import transformers
from modelopt.torch.distill.plugins.huggingface import DistillArguments
from modelopt.torch.opt.plugins.transformers import ModelOptArgParser, ModelOptHFArguments
from modelopt.torch.opt.plugins.transformers import (
ModelOptArgParser,
ModelOptHFArguments,
ModelOptTrainerArguments,
)
from modelopt.torch.quantization.plugins.transformers_trainer import (
QuantizationArguments as ModelOptQuantizationArguments,
)
@@ -34,13 +38,22 @@ class ModelArguments(ModelOptHFArguments):
},
)
model_max_length: int = field(
default=4096,
default=8192,
metadata={
"help": (
"Maximum sequence length. Sequences will be right-padded (and possibly truncated)."
)
},
)
attn_implementation: str | None = field(
default=None,
metadata={
"help": (
"Attention implementation: 'flash_attention_2', 'flash_attention_3', "
"'sdpa', or 'eager'."
)
},
)
class DataArguments(ModelOptHFArguments):
@@ -78,10 +91,13 @@ class DataArguments(ModelOptHFArguments):
)
class TrainingArguments(ModelOptHFArguments, transformers.TrainingArguments):
cache_dir: str | None = field(default=None)
class TrainingArguments(ModelOptTrainerArguments, transformers.TrainingArguments):
dataloader_drop_last: bool = field(default=True)
bf16: bool = field(default=True)
use_liger_kernel: bool = field(
default=True,
metadata={"help": "Use Liger kernel for fused loss computation. Reduces memory usage."},
)
lora: bool = field(
default=False,
metadata={
+2 -2
View File
@@ -15,10 +15,10 @@ learning_rate: 1e-5
per_device_train_batch_size: 2
per_device_eval_batch_size: 2
gradient_accumulation_steps: 2
model_max_length: 4096
model_max_length: 8192
warmup_ratio: 0.05
lr_scheduler_type: cosine
gradient_checkpointing: true
use_liger_kernel: true
seed: 42
# Checkpointing
@@ -0,0 +1,39 @@
# Per-parameter optimizer config example
#
# Maps fnmatch glob patterns to optimizer kwargs (lr, weight_decay, betas,
# eps, etc.). First matching pattern wins per parameter. Parameters not
# matching any pattern use the global values from the train config.
#
# Any keyword accepted by the optimizer constructor can be specified here.
# Common kwargs for AdamW:
# lr - learning rate
# weight_decay - L2 penalty (overrides the global --weight_decay)
# betas - Adam momentum coefficients [beta1, beta2]
# eps - term added to denominator for numerical stability
#
# Usage:
# --lr_config configs/train/lr_config_example.yaml
#
# Tip: use `model.named_parameters()` to find the exact parameter names
# for your model.
# Output head — lower LR, no weight decay
"*lm_head*":
lr: 1e-5
weight_decay: 0.0
# Attention layers — custom LR + more aggressive momentum
"*self_attn*":
lr: 5e-5
betas: [0.9, 0.95]
# MLP layers — custom LR + higher weight decay
"*mlp*":
lr: 5e-5
weight_decay: 0.05
# Embedding layers (often kept at a lower LR or frozen)
"*embed_tokens*":
lr: 1e-6
weight_decay: 0.0
eps: 1e-7
@@ -3,6 +3,7 @@
# Model
model_name_or_path: # e.g., Qwen/Qwen3-8B
output_dir: # e.g., qwen3-8b-qad-nvfp4
attn_implementation: flash_attention_2
# Distillation
distill: true
@@ -19,10 +20,11 @@ learning_rate: 1e-5
per_device_train_batch_size: 2
per_device_eval_batch_size: 2
gradient_accumulation_steps: 2
model_max_length: 4096
model_max_length: 8192
warmup_ratio: 0.05
lr_scheduler_type: cosine
gradient_checkpointing: true
use_liger_kernel: true
manual_gc: true
seed: 42
do_train: true
do_eval: true
@@ -3,6 +3,7 @@
# Model
model_name_or_path: # e.g., Qwen/Qwen3-8B
output_dir: # e.g., qwen3-8b-qat-nvfp4
attn_implementation: flash_attention_2
# Dataset
dataset_config: configs/dataset/blend.yaml
@@ -15,10 +16,11 @@ learning_rate: 1e-5
per_device_train_batch_size: 2
per_device_eval_batch_size: 2
gradient_accumulation_steps: 2
model_max_length: 4096
model_max_length: 8192
warmup_ratio: 0.05
lr_scheduler_type: cosine
gradient_checkpointing: true
use_liger_kernel: true
manual_gc: true
seed: 42
do_train: true
do_eval: true
@@ -18,10 +18,12 @@ learning_rate: 1e-3
per_device_train_batch_size: 2
per_device_eval_batch_size: 2
gradient_accumulation_steps: 2
model_max_length: 4096
model_max_length: 8192
warmup_ratio: 0.05
lr_scheduler_type: cosine
gradient_checkpointing: true
use_liger_kernel: true
manual_gc: true
seed: 42
do_train: true
do_eval: true
@@ -35,11 +35,12 @@ from transformers import AutoModelForCausalLM, PreTrainedModel, PreTrainedTokeni
from wrapt import register_post_import_hook
import modelopt.torch.opt as mto
from modelopt.torch.distill.plugins.huggingface import LMLogitsLoss
from modelopt.torch.quantization.plugins.transformers_trainer import QADTrainer, QATTrainer
logger = logging.get_logger(__name__)
_model_init_kwargs: dict[str, Any] = {}
@dataclass
class QuantizationArguments:
@@ -125,6 +126,8 @@ def patch_load_module(module):
config = load_config(model_args)
patch_config(config, tokenizer, model_args, init_kwargs, is_trainable)
global _model_init_kwargs
_model_init_kwargs = init_kwargs.copy()
assert not model_args.enable_liger_kernel, "Liger kernel is currently not supported."
assert not model_args.use_unsloth, "Unsloth is currently not supported."
@@ -216,12 +219,9 @@ def create_patch_module(quant_args=None, distill_args=None):
if distill_args and distill_args.distill:
teacher_model = transformers.AutoModelForCausalLM.from_pretrained(
distill_args.teacher_model,
**_model_init_kwargs,
)
distill_config = {
"teacher_model": teacher_model,
"criterion": LMLogitsLoss(),
}
modelopt_trainer_args["distill_config"] = distill_config
modelopt_trainer_args["distill_args"] = {"teacher_model": teacher_model}
super().__init__(*args, **modelopt_trainer_args, **kwargs)
# Replace the trainer class in the module
+4
View File
@@ -68,10 +68,14 @@ def quantize():
# Load model and tokenizer
print_rank_0(f"Loading model: {model_args.model_name_or_path}")
model_kwargs = {}
if model_args.attn_implementation:
model_kwargs["attn_implementation"] = model_args.attn_implementation
model = transformers.AutoModelForCausalLM.from_pretrained(
model_args.model_name_or_path,
dtype=torch.bfloat16,
device_map="auto",
**model_kwargs,
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
model_args.model_name_or_path, model_max_length=model_args.model_max_length
+2 -1
View File
@@ -1,3 +1,4 @@
flash-attn
flash-attn>=2.6.0
liger-kernel>=0.5.0; platform_system != 'Darwin' and platform_system != 'Windows'
py7zr
tensorboard
+35 -35
View File
@@ -42,8 +42,6 @@ import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
from warnings import warn
import torch
import transformers
from arguments import get_training_args
@@ -61,14 +59,6 @@ mto.enable_huggingface_checkpointing()
def train():
model_args, training_args, data_args, distill_args = get_training_args()
if distill_args.distill and getattr(training_args, "fsdp_config", None):
fsdp_cfg = training_args.fsdp_config
if fsdp_cfg.get("fsdp_cpu_ram_efficient_loading", True):
warn(
"Distillation with FSDP2 may require --fsdp_cpu_ram_efficient_loading False. "
"Set this if you encounter issues loading the teacher model."
)
print_rank_0(f"arguments: {model_args}, {training_args}, {data_args}, {distill_args}")
# Detecting last checkpoint.
@@ -77,10 +67,14 @@ def train():
last_checkpoint = get_last_checkpoint(training_args.output_dir)
print_rank_0(f"Last checkpoint detected: {last_checkpoint}")
model_kwargs = {}
if model_args.attn_implementation:
model_kwargs["attn_implementation"] = model_args.attn_implementation
model = transformers.AutoModelForCausalLM.from_pretrained(
model_args.model_name_or_path,
cache_dir=training_args.cache_dir,
dtype=torch.bfloat16,
**model_kwargs,
)
model.generation_config.do_sample = True
tokenizer = transformers.AutoTokenizer.from_pretrained(
@@ -106,33 +100,39 @@ def train():
if checkpoint is not None and training_args.lora:
raise RuntimeError("Does not support LoRA resuming training yet!")
# Torch >= 2.4 throws an error if `use_reentrant` is not set explicitly
if training_args.gradient_checkpointing and training_args.gradient_checkpointing_kwargs is None:
training_args.gradient_checkpointing_kwargs = {"use_reentrant": True}
distill_kwargs = {}
if distill_args.distill:
if distill_args.teacher_model is None:
raise ValueError("--teacher_model is required when --distill is enabled.")
teacher_model = transformers.AutoModelForCausalLM.from_pretrained(
distill_args.teacher_model,
cache_dir=training_args.cache_dir,
dtype=torch.bfloat16,
)
distill_kwargs = distill_args.to_distill_kwargs(teacher_model)
trainer_cls = QADTrainer if distill_args.distill else QATTrainer
if training_args.lora:
training_args.lora_config = get_lora_config()
trainer = trainer_cls(
model=model,
processing_class=tokenizer,
args=training_args,
**distill_kwargs,
**data_module,
)
distill_config = None
if distill_args.distill:
assert distill_args.teacher_model is not None, "Teacher model is required for distillation."
teacher = transformers.AutoModelForCausalLM.from_pretrained(
distill_args.teacher_model,
dtype=torch.bfloat16,
**model_kwargs,
)
distill_config = {
"teacher_model": teacher,
"temperature": distill_args.temperature,
"criterion": distill_args.criterion,
"liger_jsd_beta": distill_args.liger_jsd_beta,
}
if distill_config is None:
trainer = QATTrainer(
model=model,
processing_class=tokenizer,
args=training_args,
**data_module,
)
else:
trainer = QADTrainer(
model=model,
processing_class=tokenizer,
args=training_args,
distill_args=distill_config,
**data_module,
)
if training_args.do_train:
trainer.train(resume_from_checkpoint=checkpoint)
+271 -112
View File
@@ -13,20 +13,37 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""ModelOpt plugin to train HuggingFace models with knowledge distillation."""
"""ModelOpt plugin to train HuggingFace models with knowledge distillation.
Only logit-level distillation is supported. For intermediate-layer distillation
or Megatron models, use ``mtd.convert()`` directly.
"""
from contextlib import contextmanager, nullcontext
from dataclasses import field
from torch import Tensor
from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers.trainer_pt_utils import LabelSmoother
import torch
import torch.nn as nn
from transformers.trainer_pt_utils import find_batch_size
import modelopt.torch.distill as mtd
from modelopt.torch.distill.losses import LogitsDistillationLoss
from modelopt.torch.opt.plugins import ModelOptHFTrainer
from modelopt.torch.opt.plugins.transformers import ModelOptHFArguments
from modelopt.torch.opt.plugins.transformers import (
_LIGER_KERNEL_IMPORT_ERROR,
ModelOptHFArguments,
_forward_redirect,
is_liger_available,
)
from modelopt.torch.utils import print_rank_0
IGNORE_TOKEN_ID = LabelSmoother.ignore_index # equals -100
__all__ = [
"IGNORE_INDEX",
"DistillArgsWithTeacherModel",
"DistillArguments",
"KDTrainer",
]
IGNORE_INDEX = nn.CrossEntropyLoss().ignore_index
_SUPPORTED_CRITERIA = {"logits_loss"}
@@ -40,7 +57,7 @@ class DistillArguments(ModelOptHFArguments):
)
teacher_model: str | None = field(
default=None,
metadata={"help": "The name or path of the teacher model to use for distillation."},
metadata={"help": "The name or path of the teacher model."},
)
criterion: str = field(
default="logits_loss",
@@ -48,137 +65,279 @@ class DistillArguments(ModelOptHFArguments):
"help": "Distillation loss criterion. Currently only 'logits_loss' is supported."
},
)
def to_distill_kwargs(self, teacher_model) -> dict:
"""Convert distill args to kwargs for KDTrainer/QADTrainer.
Args:
teacher_model: The loaded teacher model instance.
Returns:
Dict with ``distill_config`` ready to pass to the trainer.
"""
if self.criterion not in _SUPPORTED_CRITERIA:
raise ValueError(
f"Unsupported criterion: {self.criterion!r}. Supported: {_SUPPORTED_CRITERIA}"
temperature: float = field(
default=1.0,
metadata={
"help": (
"Softmax temperature for softening logits in KD loss. "
"Used by both standard and Liger KD loss."
)
return {"distill_config": {"teacher_model": teacher_model, "criterion": LMLogitsLoss()}}
},
)
liger_jsd_beta: float = field(
default=0.0,
metadata={
"help": (
"JSD beta coefficient in [0, 1]. 0=forward KL, 1=reverse KL. "
"Only used when --use_liger_kernel is enabled."
)
},
)
class DistillArgsWithTeacherModel(DistillArguments):
"""Runtime distillation arguments with a pre-loaded teacher model."""
teacher_model: nn.Module | None = field(
default=None,
metadata={"help": "Pre-loaded teacher model."},
)
class KDTrainer(ModelOptHFTrainer):
"""Distillation trainer for HuggingFace models."""
"""Distillation trainer for HuggingFace models.
def __init__(self, *args, distill_config=None, **kwargs):
"""Initialize the trainer."""
Supports logit-level knowledge distillation only. The teacher model is stored
separately on the trainer and forwarded explicitly during loss computation.
No ``mtd.convert()`` or ``DistillationModel`` wrapping is used.
"""
def __init__(
self,
*args,
distill_args: DistillArgsWithTeacherModel | dict | None = None,
**kwargs,
):
"""Initialize the trainer.
Args:
distill_args: Runtime distillation config with a pre-loaded teacher model.
"""
super().__init__(*args, **kwargs)
if self.is_fsdp_enabled and not self.accelerator.is_fsdp2:
raise ValueError("FSDP1 is not supported for distillation. Use FSDP2 instead.")
assert distill_config is not None, "`distill_config` is required for distillation."
self.distill_config = distill_config
self._convert_to_distillation_model()
if distill_args is None:
raise ValueError("`distill_args` is required for distillation.")
if isinstance(distill_args, dict):
distill_args = DistillArgsWithTeacherModel(**distill_args)
def _convert_to_distillation_model(self):
"""Convert the model to a distillation model."""
mtd.convert(self.model, mode=[("kd_loss", self.distill_config)])
print_rank_0("Distillation model created.")
if distill_args.criterion not in _SUPPORTED_CRITERIA:
raise ValueError(
f"Unsupported criterion: {distill_args.criterion!r}. "
f"Supported: {_SUPPORTED_CRITERIA}"
)
def compute_loss(self, model, inputs, *args, **kwargs):
"""Compute loss for distillation.
teacher = distill_args.teacher_model
if teacher is None:
raise ValueError("`distill_args.teacher_model` is required.")
if not isinstance(teacher, nn.Module):
raise TypeError(
"`distill_args.teacher_model` must be a pre-loaded nn.Module. "
"Load the teacher in the training script before constructing KDTrainer."
)
Change the training loss to distillation loss and keep the original validation loss.
self._teacher_model = teacher
self._teacher_model.requires_grad_(False)
self._kd_criterion = LogitsDistillationLoss(
temperature=distill_args.temperature, reduction="none"
)
self._teacher_prepared = False
self._eval_kd_loss_totals = None
Args:
model: The model to compute loss for.
inputs: The inputs to the model.
if self.use_liger_kernel:
self._liger_temperature = distill_args.temperature
self._liger_jsd_beta = distill_args.liger_jsd_beta
def _setup_liger_fused_loss(self):
"""Set student fused-loss path and require Liger KD dependencies."""
if not is_liger_available():
raise ImportError(_LIGER_KERNEL_IMPORT_ERROR)
model = self.accelerator.unwrap_model(self.model)
if not hasattr(model, "lm_head"):
self.use_liger_kernel = False
return
self.compute_loss_func = self._liger_loss_func
def _ensure_teacher_prepared(self):
"""Prepare teacher model via accelerator (handles FSDP2, DeepSpeed, DDP)."""
if self._teacher_prepared:
return
self._teacher_prepared = True
self._teacher_model = self._prepare_model(self._teacher_model)
print_rank_0("Teacher model prepared for distillation.")
def _get_unwrapped_teacher(self):
"""Unwrap teacher model (removes FSDP/DDP/DeepSpeed wrapper)."""
return self.accelerator.unwrap_model(self._teacher_model)
@contextmanager
def _ds_gather(self, params):
"""Gather DS ZeRO-3 partitioned params; no-op if DeepSpeed disabled.
The teacher is loaded under an active ``zero.Init`` but not wrapped in a
DeepSpeedEngine, so its params have no per-module gather hooks and need an
explicit gather around any forward use.
"""
if not model.training:
_compute_loss_func = self.compute_loss_func
self.compute_loss_func = None
if self.is_deepspeed_enabled:
import deepspeed
loss = super().compute_loss(model, inputs, *args, **kwargs)
with deepspeed.zero.GatheredParameters(list(params), modifier_rank=None):
yield
else:
yield
if not model.training:
self.compute_loss_func = _compute_loss_func
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
"""Train on KD loss and evaluate on CE loss with KD as a metric."""
self._ensure_teacher_prepared()
kd_inputs = {k: v for k, v in inputs.items() if k != "labels"}
labels = inputs.get("labels")
is_training = model.training
if is_training:
student_context = self._liger_identity_lm_head if self.use_liger_kernel else nullcontext
with student_context():
outputs = model(**kd_inputs)
else:
loss, outputs = super().compute_loss(model, inputs, return_outputs=True, **kwargs)
kd_loss = self._compute_kd_loss(outputs, labels, kd_inputs, **kwargs)
if is_training:
loss = kd_loss
else:
batch_size = find_batch_size(inputs)
self._record_eval_kd_loss(kd_loss, batch_size)
return (loss, outputs) if return_outputs else loss
def _compute_kd_loss(self, outputs, labels, inputs, **kwargs):
"""Run teacher forward and compute KD loss.
The student forward has already run. When Liger is enabled, teacher
forward also needs the identity-lm_head context so
both student and teacher outputs are hidden states for fused KD.
"""
lm_head_context = (
self._teacher_liger_identity_lm_head if self.use_liger_kernel else nullcontext
)
with lm_head_context():
teacher_outputs = self._compute_teacher_outputs(inputs)
self._last_teacher_outputs = teacher_outputs
if self.use_liger_kernel:
return self._liger_kd_loss(outputs, labels, **kwargs)
return self._standard_kd_loss(outputs, labels, **kwargs)
def _compute_teacher_outputs(self, inputs):
with torch.no_grad(), self._ds_gather(self._teacher_model.parameters()):
self._teacher_model.eval()
return self._teacher_model(**inputs)
def _standard_kd_loss(self, outputs, labels, **kwargs):
"""KD loss with causal shift and ignore-index masking."""
# Match causal LM CE: logits at position t are scored against label t+1.
student_logits = outputs.logits[..., :-1, :].float()
teacher_logits = self._last_teacher_outputs.logits[..., :-1, :].float()
self._last_teacher_outputs = None
per_token_loss = self._kd_criterion(student_logits, teacher_logits)
if labels is None:
return per_token_loss.mean()
shift_labels = labels[..., 1:]
mask = shift_labels != IGNORE_INDEX
loss = (per_token_loss * mask).sum() / mask.sum().clamp(min=1)
return loss
def save_model(
self,
output_dir: str | None = None,
_internal_call: bool = False,
*args,
**kwargs,
):
"""Dumps model and ModelOpt states to disk.
@contextmanager
def _teacher_liger_identity_lm_head(self):
"""Patch teacher lm_head to identity for fused KD."""
teacher = self._get_unwrapped_teacher()
teacher_lm_head = self._get_lm_head(teacher)
teacher_orig = teacher_lm_head.forward
teacher_lm_head.forward = lambda x: x
try:
yield
finally:
teacher_lm_head.forward = teacher_orig
Note: Will save pretrained model in safetensors format if called manually, otherwise will
save in training checkpointformat (when called internally by transformers Trainer).
Args:
output_dir: The directory to save the model and ModelOpt states.
"""
if output_dir is None:
output_dir = self.args.output_dir
def _liger_kd_loss(self, outputs, labels, **kwargs):
"""Fused lm_head + JSD for KD."""
from liger_kernel.transformers import LigerFusedLinearJSD
model = self.accelerator.unwrap_model(self.model)
with model.hide_teacher_model(), model.hide_loss_modules(enable=not _internal_call):
if _internal_call:
return super().save_model(output_dir, _internal_call, *args, **kwargs)
teacher = self._get_unwrapped_teacher()
extra_kwargs = {}
if self.is_fsdp_enabled:
extra_kwargs["save_function"] = self.accelerator.save
extra_kwargs["state_dict"] = self.accelerator.get_state_dict(self.model)
self.accelerator.wait_for_everyone() # needed to prevent hang somehow
student_lm_head = self._get_lm_head(model)
teacher_lm_head = self._get_lm_head(teacher)
model.save_pretrained(
output_dir,
is_main_process=self.accelerator.is_main_process,
**extra_kwargs,
student_hs = outputs.logits.to(student_lm_head.weight.dtype) # RMSNorm may upcast to fp32
teacher_hs = self._last_teacher_outputs.logits.to(teacher_lm_head.weight.dtype)
self._last_teacher_outputs = None
# Causal LM shift
student_hs = student_hs[..., :-1, :].contiguous().view(-1, student_hs.size(-1))
teacher_hs = teacher_hs[..., :-1, :].contiguous().view(-1, teacher_hs.size(-1))
shift_labels = labels[..., 1:].contiguous().view(-1)
jsd = LigerFusedLinearJSD(
jsd_beta=self._liger_jsd_beta,
ignore_index=IGNORE_INDEX,
temperature=self._liger_temperature,
)
def _compute():
return self._teacher_liger_enabled(
lambda: jsd(
student_hs,
student_lm_head.weight,
teacher_hs,
teacher_lm_head.weight,
shift_labels,
),
teacher_lm_head,
)
self.processing_class.save_pretrained(output_dir)
def train(self, *args, **kwargs):
"""Train the model."""
return super()._sharded_liger_compute(_compute)
def _compute_kd_loss(outputs: Tensor, labels: Tensor | None, **kwargs):
def loss_reduction_fn(loss: Tensor):
if labels is None:
return loss.mean()
loss_mask = labels != IGNORE_TOKEN_ID
return (loss * loss_mask).sum() / loss_mask.sum().clamp(min=1)
def _teacher_liger_enabled(self, fn, teacher_lm_head):
if self.is_fsdp_enabled:
return _forward_redirect(self._teacher_model, fn)
if self.is_deepspeed_enabled:
# Teacher is not in the DS engine; gather its lm_head explicitly.
with self._ds_gather([teacher_lm_head.weight]):
return fn()
return fn()
return self.model.compute_kd_loss(loss_reduction_fn=loss_reduction_fn)
def evaluation_loop(
self,
dataloader,
description,
prediction_loss_only=None,
ignore_keys=None,
metric_key_prefix="eval",
):
"""Add KD loss as a secondary evaluation metric."""
self._eval_kd_loss_totals = None
output = super().evaluation_loop(
dataloader,
description,
prediction_loss_only=prediction_loss_only,
ignore_keys=ignore_keys,
metric_key_prefix=metric_key_prefix,
)
if self._eval_kd_loss_totals is not None:
output.metrics[f"{metric_key_prefix}_kd_loss"] = self._get_eval_kd_loss()
return output
self.compute_loss_func = _compute_kd_loss
return super().train(*args, **kwargs)
def _record_eval_kd_loss(self, loss, batch_size):
count = loss.new_tensor(float(batch_size or 1))
totals = torch.stack([loss.detach() * count, count])
self._eval_kd_loss_totals = (
totals
if self._eval_kd_loss_totals is None
else self._eval_kd_loss_totals + totals.to(self._eval_kd_loss_totals.device)
)
class LMLogitsLoss(mtd.LogitsDistillationLoss):
"""Logits loss for language-model knowledge distillation.
Defaults to ``reduction="none"`` to support per-token loss masking via ``loss_reduction_fn``
in :meth:`DistillationModel.compute_kd_loss`. This allows masking out padding and non-assistant
tokens before reducing the loss.
"""
def __init__(self, temperature: float = 1.0, reduction: str = "none"):
"""Constructor.
Args:
temperature: A value used to soften the logits before computing loss.
reduction: How to reduce the final pointwise loss. Defaults to ``"none"`` to
allow loss-masking via ``loss_reduction_fn`` in ``compute_kd_loss``.
"""
super().__init__(temperature=temperature, reduction=reduction)
def forward(self, out_student: CausalLMOutputWithPast, out_teacher: CausalLMOutputWithPast):
"""Forward pass for logits distillation loss.
Args:
out_student: The student model output.
out_teacher: The teacher model output.
"""
student_logits, teacher_logits = out_student.logits.float(), out_teacher.logits.float()
return super().forward(student_logits, teacher_logits)
def _get_eval_kd_loss(self):
totals = self.accelerator.gather_for_metrics(self._eval_kd_loss_totals)
totals = totals.reshape(-1, 2).sum(dim=0)
return (totals[0] / totals[1].clamp(min=1)).item()
+427 -6
View File
@@ -16,19 +16,25 @@
"""ModelOpt plugin for enabling automatic save/restore of ModelOpt state for HuggingFace models."""
import dataclasses
import fnmatch
import gc
import json
import os
import sys
import types
from contextlib import contextmanager
import warnings
from contextlib import contextmanager, suppress
from pathlib import Path
from typing import Any
import torch
import torch.nn as nn
import transformers
from packaging.version import Version
from transformers import HfArgumentParser, PreTrainedModel, Trainer, TrainerCallback
from transformers import modeling_utils as tf_modeling_utils
from modelopt.torch.utils import report_memory
from modelopt.torch.utils import print_rank_0, report_memory
from ..conversion import ModeloptStateManager, load_modelopt_state
from .huggingface import (
@@ -39,7 +45,23 @@ from .huggingface import (
register_for_patching,
)
__all__ = ["ModelOptArgParser", "ModelOptHFArguments", "ModelOptHFTrainer"]
IGNORE_INDEX = nn.CrossEntropyLoss().ignore_index
_LIGER_KERNEL_IMPORT_ERROR = "`use_liger_kernel=True` requires the optional `liger-kernel` package."
__all__ = [
"ModelOptArgParser",
"ModelOptHFArguments",
"ModelOptHFTrainer",
"ModelOptTrainerArguments",
]
def is_liger_available():
try:
__import__("liger_kernel")
except ImportError:
return False
return True
@contextmanager
@@ -177,6 +199,62 @@ class ModelOptHFArguments:
dataclasses.dataclass(cls)
class ModelOptTrainerArguments(ModelOptHFArguments):
"""Arguments for ModelOptHFTrainer controlling param freezing, LR config, and save dtype.
This class can be used with HuggingFace's ``HfArgumentParser`` for CLI parsing.
"""
trainable_params: list[str] | None = dataclasses.field(
default=None,
metadata={
"nargs": "+",
"help": (
"Glob patterns (fnmatch) for parameters that should be trainable. "
"All other parameters will be frozen. Mutually exclusive with frozen_params."
),
},
)
frozen_params: list[str] | None = dataclasses.field(
default=None,
metadata={
"nargs": "+",
"help": (
"Glob patterns (fnmatch) for parameters that should be frozen. "
"Mutually exclusive with trainable_params."
),
},
)
lr_config: str | None = dataclasses.field(
default=None,
metadata={
"help": (
"Path to a YAML file mapping fnmatch patterns to optimizer kwargs "
"(e.g. lr, weight_decay). First matching pattern wins per parameter. "
"See examples/llm_qat/configs/train/lr_config_example.yaml."
),
},
)
manual_gc: bool = dataclasses.field(
default=False,
metadata={
"help": (
"Run `gc.collect()` before each training/prediction step to work around "
"GPU memory leaks during QAT/distillation."
),
},
)
liger_ce_label_smoothing: float = dataclasses.field(
default=0.0,
metadata={
"help": (
"Label smoothing for Liger fused CE loss. "
"Only used when --use_liger_kernel is enabled."
),
},
)
class ModelOptArgParser(HfArgumentParser):
"""HfArgumentParser with ``--config`` YAML support and ``--generate_docs`` for ARGUMENTS.md."""
@@ -363,14 +441,357 @@ class _MemoryReportCallback(TrainerCallback):
_report_memory("Memory usage at evaluation")
def _forward_redirect(module, fn):
"""Run ``fn`` inside ``module``'s forward to trigger distributed param gathering.
Works for both FSDP2 (unshards DTensor params) and DeepSpeed ZeRO-3
(gathers partitioned params via per-module forward hooks).
"""
original_forward = module.forward
def wrapped_forward(*a, **kw):
module.forward = original_forward
return fn()
module.forward = wrapped_forward
try:
dummy = torch.empty(1, device=next(module.parameters()).device)
return module(dummy)
except Exception:
module.forward = original_forward
raise
class ModelOptHFTrainer(Trainer):
"""A drop-in replacement of HuggingFace's Trainer for ModelOpt.
This class adds extra utilities for ModelOpt checkpointing and memory reporting.
This class adds extra utilities for ModelOpt checkpointing, memory reporting,
parameter freezing, per-layer learning rates, Liger fused loss, and original-dtype-preserving save.
**Liger kernel support:** When ``--use_liger_kernel`` is set, this trainer provides
model-agnostic fused loss computation that extends HuggingFace's built-in Liger
integration in three ways:
1. **Model-agnostic**: Works with any causal LM that has an ``lm_head``, unlike
HF's Liger which only supports `a fixed set of model architectures
<https://github.com/linkedin/Liger-Kernel/blob/main/src/liger_kernel/transformers/monkey_patch.py>`_.
2. **DeepSpeed ZeRO-3 support**: HF's Liger integration only works with FSDP.
ModelOpt adds distributed param gathering for DeepSpeed ZeRO-3 and DDP as well.
3. **KD loss support**: ``KDTrainer`` extends fused loss to knowledge distillation
via ``LigerFusedLinearJSD`` for fused lm_head + Jensen-Shannon divergence.
"""
def __init__(self, *args, **kwargs):
"""Initialize."""
def __init__(
self,
*args,
trainer_args: ModelOptTrainerArguments | None = None,
lr_config: dict[str, dict[str, Any]] | None = None,
**kwargs,
):
"""Initialize.
Args:
trainer_args: Optional arguments for param freeze and lr config.
lr_config: Optional dict for per-pattern optimizer param groups
(overrides trainer_args.lr_config).
"""
enable_huggingface_checkpointing()
super().__init__(*args, **kwargs)
_raw_dtype = getattr(getattr(self.model, "config", None), "dtype", None) or getattr(
getattr(self.model, "config", None), "torch_dtype", None
)
self._original_dtype = None if _raw_dtype is None else str(_raw_dtype).rsplit(".", 1)[-1]
if trainer_args is None and isinstance(self.args, ModelOptTrainerArguments):
trainer_args = self.args
self.trainer_args = trainer_args or ModelOptTrainerArguments()
self._lr_config = self._resolve_lr_config(lr_config, self.trainer_args)
self._apply_gradient_checkpointing_defaults()
self.add_callback(_MemoryReportCallback())
self.use_liger_kernel = getattr(self.args, "use_liger_kernel", False)
if self.use_liger_kernel:
if self.is_fsdp_enabled and not self.accelerator.is_fsdp2:
raise ValueError("Liger fused loss is not supported with FSDP1. Use FSDP2 instead.")
self._setup_liger_fused_loss()
self._configure_trainable_params()
def _prepare_model(self, model):
"""Prepare a model via accelerator (materializes meta-device params, applies sharding).
Uses a dummy optimizer because ``accelerator.prepare`` requires one for FSDP2.
Works generically for FSDP2, DDP, and DeepSpeed backends. For fully-frozen models
under DS ZeRO-3, falls back to inference-mode prep since ZeRO-3 asserts on empty
trainable_param_groups; in that case the caller is responsible for gathering
``zero.Init``-partitioned params around forward passes.
"""
if self.is_deepspeed_enabled and not any(p.requires_grad for p in model.parameters()):
return self.accelerator.prepare_model(model, evaluation_mode=True)
dummy_optimizer = torch.optim.SGD([next(model.parameters())], lr=0.0)
model, _ = self.accelerator.prepare(model, dummy_optimizer)
return model
def training_step(self, *args, **kwargs):
"""Run gc.collect() before the training step if manual_gc is enabled."""
if self.trainer_args.manual_gc:
gc.collect()
return super().training_step(*args, **kwargs)
def prediction_step(self, *args, **kwargs):
"""Run gc.collect() before the prediction step if manual_gc is enabled."""
if self.trainer_args.manual_gc:
gc.collect()
return super().prediction_step(*args, **kwargs)
def _load_best_model(self, *args, **kwargs):
"""Run gc.collect() before loading the best model if manual_gc is enabled."""
if self.trainer_args.manual_gc:
gc.collect()
return super()._load_best_model(*args, **kwargs)
def _apply_gradient_checkpointing_defaults(self):
"""Ensure non-reentrant gradient checkpointing when no explicit kwargs are set."""
args = self.args
if not getattr(args, "gradient_checkpointing", False):
return
if args.gradient_checkpointing_kwargs is None:
args.gradient_checkpointing_kwargs = {"use_reentrant": False}
else:
if args.gradient_checkpointing_kwargs.get("use_reentrant", False):
warnings.warn(
"ModelOpt overriding `use_reentrant=True` to `use_reentrant=False` "
"for gradient checkpointing compatibility.",
UserWarning,
stacklevel=2,
)
args.gradient_checkpointing_kwargs["use_reentrant"] = False
def _configure_trainable_params(self):
"""Freeze/unfreeze parameters based on trainer_args.trainable_params or frozen_params."""
trainable = self.trainer_args.trainable_params
frozen = self.trainer_args.frozen_params
if not trainable and not frozen:
return
if trainable and frozen:
raise ValueError("trainable_params and frozen_params are mutually exclusive.")
def _matches(name, patterns):
return any(fnmatch.fnmatch(name, p) for p in patterns)
model = self.model
if trainable:
for name, param in model.named_parameters():
param.requires_grad_(_matches(name, trainable))
else:
for name, param in model.named_parameters():
if _matches(name, frozen):
param.requires_grad_(False)
trainable_count = sum(p.requires_grad for p in model.parameters())
total_count = sum(1 for _ in model.parameters())
print_rank_0(
f"Trainable params: {trainable_count}/{total_count} "
f"({100 * trainable_count / max(total_count, 1):.1f}%)"
)
@staticmethod
def _resolve_lr_config(
lr_config: dict[str, dict[str, Any]] | None,
trainer_args: ModelOptTrainerArguments,
) -> dict[str, dict[str, Any]] | None:
if lr_config is not None:
return lr_config
path = getattr(trainer_args, "lr_config", None)
if path is not None:
return ModelOptHFTrainer.load_lr_config(path)
return None
@staticmethod
def load_lr_config(path: str) -> dict[str, dict[str, Any]]:
"""Load an lr_config YAML file mapping fnmatch patterns to optimizer kwargs.
Example YAML::
"*lm_head*":
lr: 1e-5
"*mlp*":
lr: 5e-5
Returns:
Ordered dict of ``{pattern: {kwarg: value, ...}}``.
"""
import yaml
with open(path) as f:
cfg = yaml.safe_load(f)
if not isinstance(cfg, dict):
raise ValueError(f"lr_config must be a YAML mapping, got {type(cfg).__name__}")
for pattern, kwargs in cfg.items():
if not isinstance(pattern, str) or not isinstance(kwargs, dict):
raise ValueError(
f"lr_config entry must be str -> dict, got {pattern!r} -> {kwargs!r}"
)
for key, val in kwargs.items():
if isinstance(val, str):
with suppress(ValueError):
kwargs[key] = float(val)
return cfg
def _match_lr_config_pattern(self, name: str) -> str | None:
"""Return the first lr_config pattern matching ``name``, or None."""
for pattern in self._lr_config: # type: ignore[union-attr]
if fnmatch.fnmatch(name, pattern):
return pattern
return None
def create_optimizer(self):
"""Build per-pattern param groups from lr_config, then delegate to HF Trainer."""
if self._lr_config is None:
return super().create_optimizer()
if self.optimizer is not None:
return self.optimizer
opt_model = self.model
if self.optimizer_cls_and_kwargs is not None:
optimizer_cls, optimizer_kwargs = self.optimizer_cls_and_kwargs
else:
optimizer_cls, optimizer_kwargs = self.get_optimizer_cls_and_kwargs(
self.args, opt_model
)
decay_parameters = self.get_decay_parameter_names(opt_model)
groups: dict[tuple[str | None, bool], list] = {}
for name, param in opt_model.named_parameters():
if not param.requires_grad:
continue
pattern = self._match_lr_config_pattern(name)
is_decay = name in decay_parameters
groups.setdefault((pattern, is_decay), []).append(param)
param_groups = []
for (pattern, is_decay), params in groups.items():
group: dict[str, Any] = {
"params": params,
"weight_decay": self.args.weight_decay if is_decay else 0.0,
}
if pattern is not None:
group.update(self._lr_config[pattern])
param_groups.append(group)
optimizer_kwargs["params"] = param_groups
self.optimizer_cls_and_kwargs = (optimizer_cls, optimizer_kwargs)
result = super().create_optimizer()
self._log_lr_config_summary()
return result
def _log_lr_config_summary(self):
if self.optimizer is None:
return
lines = ["lr_config optimizer param groups:"]
for i, group in enumerate(self.optimizer.param_groups):
lr = group.get("lr", "default")
wd = group.get("weight_decay", "default")
n_params = len(group["params"])
lines.append(f" group {i}: {n_params} params, lr={lr}, weight_decay={wd}")
print_rank_0("\n".join(lines))
def _get_lm_head(self, model):
"""Resolve lm_head from model at call time (no cached pointer to FSDP-managed params)."""
return model.lm_head
def _setup_liger_fused_loss(self):
"""Set compute_loss_func for fused CE."""
if not is_liger_available():
raise ImportError(_LIGER_KERNEL_IMPORT_ERROR)
model = self.accelerator.unwrap_model(self.model)
if not hasattr(model, "lm_head"):
self.use_liger_kernel = False
return
self.compute_loss_func = self._liger_loss_func
@contextmanager
def _liger_identity_lm_head(self):
"""Temporarily patch lm_head to identity for fused loss computation."""
model = self.accelerator.unwrap_model(self.model)
lm_head = self._get_lm_head(model)
original_forward = lm_head.forward
lm_head.forward = lambda x: x
try:
yield
finally:
lm_head.forward = original_forward
def _sharded_liger_compute(self, fn):
"""Route fn through sharded DP to ensure lm_head params are gathered. No-op for DDP."""
if self.is_fsdp_enabled:
return _forward_redirect(self.model, fn)
if self.is_deepspeed_enabled:
lm_head = self._get_lm_head(self.accelerator.unwrap_model(self.model))
return _forward_redirect(lm_head, fn)
return fn()
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
"""Compute loss, patching lm_head to identity when using liger fused loss."""
if self.use_liger_kernel:
with self._liger_identity_lm_head():
return super().compute_loss(model, inputs, return_outputs=return_outputs, **kwargs)
return super().compute_loss(model, inputs, return_outputs=return_outputs, **kwargs)
def _liger_loss_func(self, outputs, labels, num_items_in_batch=None, **kwargs):
"""Fused lm_head + CE loss via liger kernel."""
from liger_kernel.transformers.model.loss_utils import LigerForCausalLMLoss
model = self.accelerator.unwrap_model(self.model)
lm_head = self._get_lm_head(model)
hidden_states = outputs.logits.to(lm_head.weight.dtype) # RMSNorm may upcast to fp32
def _compute():
return LigerForCausalLMLoss(
hidden_states=hidden_states,
lm_head_weight=lm_head.weight,
labels=labels,
hidden_size=hidden_states.size(-1),
num_items_in_batch=num_items_in_batch,
ignore_index=IGNORE_INDEX,
label_smoothing=self.trainer_args.liger_ce_label_smoothing,
)
return self._sharded_liger_compute(_compute)
def save_model(self, *args, **kwargs):
"""Save the model and rewrite config.json dtype to preserve the original model dtype."""
outputs = super().save_model(*args, **kwargs)
if (not self.is_in_train) and self.args.should_save:
out_dir = args[0] if args else self.args.output_dir
self._update_config_json_dtype(out_dir, self._original_dtype)
return outputs
def _update_config_json_dtype(self, output_dir: str, dtype_str: str | None) -> None:
"""Rewrite <output_dir>/config.json 'dtype' (preferred) or 'torch_dtype' to dtype_str."""
if dtype_str is None:
return
cfg_path = os.path.join(output_dir, "config.json")
if not os.path.isfile(cfg_path):
print_rank_0(f"[warn] config.json not found under {output_dir}; skip dtype rewrite.")
return
try:
with open(cfg_path, encoding="utf-8") as f:
data = json.load(f)
key_to_update = (
"dtype" if "dtype" in data else ("torch_dtype" if "torch_dtype" in data else None)
)
if key_to_update is None:
print_rank_0(
"[warn] Neither 'dtype' nor 'torch_dtype' present in config.json; "
"skip dtype rewrite."
)
return
if data.get(key_to_update) != dtype_str:
data[key_to_update] = dtype_str
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
print_rank_0(f'Updated config.json: {key_to_update} -> "{dtype_str}"')
except Exception as e:
print_rank_0(f"[warn] Failed to update dtype in config.json: {e}")
@@ -17,7 +17,6 @@
import contextlib
import gc
import json
import os
import types
import warnings
@@ -232,10 +231,6 @@ class QATTrainer(ModelOptHFTrainer):
elif is_quantized(self.model):
self._save_modelopt_state_with_weights()
self._original_dtype = getattr(
getattr(self.model, "config", None), "dtype", None
) or getattr(getattr(self.model, "config", None), "torch_dtype", None)
def _save_modelopt_state_with_weights(self):
"""Save the modelopt weights for fsdp2 models."""
if torch.distributed.is_initialized():
@@ -311,9 +306,7 @@ class QATTrainer(ModelOptHFTrainer):
"""Evaluate the model."""
if self.args.do_eval and not self.args.do_train and self.accelerator.is_fsdp2:
# [Not related to ModelOpt] HF does not support eval only for FSDP2.
# This is a hack to make it work
dummy_optimizer = torch.optim.SGD([next(self.model.parameters())], lr=0.0)
self.model, _ = self.accelerator.prepare(self.model, dummy_optimizer)
self.model = self._prepare_model(self.model)
return super().evaluate(*args, **kwargs)
def train(self, *args, **kwargs):
@@ -344,11 +337,6 @@ class QATTrainer(ModelOptHFTrainer):
self.accelerator.state.fsdp_plugin.set_state_dict_type(original_type)
else:
outputs = super().save_model(*args, **kwargs)
if (not self.is_in_train) and self.args.should_save:
out_dir = args[0]
# FSDP may upcast parameter dtype to float32 during mixed-precision training,
# we convert it back to original dtype by updating `torch-dtype` in `config.json`
self._update_config_json_dtype(out_dir, str(self._original_dtype).split(".")[1])
return outputs
def _load_best_model(self, *args, **kwargs):
@@ -373,32 +361,6 @@ class QATTrainer(ModelOptHFTrainer):
else:
super()._load_best_model(*args, **kwargs)
def _update_config_json_dtype(self, output_dir: str, dtype_str: str | None) -> None:
"""Rewrite <output_dir>/config.json 'dtype' (preferred) or 'torch_dtype' to dtype_str."""
cfg_path = os.path.join(output_dir, "config.json")
if not os.path.isfile(cfg_path):
print_rank_0(f"[warn] config.json not found under {output_dir}; skip dtype rewrite.")
return
try:
with open(cfg_path, encoding="utf-8") as f:
data = json.load(f)
# Prefer 'dtype', else fall back to 'torch_dtype'
key_to_update = (
"dtype" if "dtype" in data else ("torch_dtype" if "torch_dtype" in data else None)
)
if key_to_update is None:
print_rank_0(
"[warn] Neither 'dtype' nor 'torch_dtype' present in config.json; skip dtype rewrite."
)
return
if data.get(key_to_update) != dtype_str:
data[key_to_update] = dtype_str
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
print_rank_0(f'Updated config.json: {key_to_update} -> "{dtype_str}"')
except Exception as e:
print_rank_0(f"[warn] Failed to update dtype in config.json: {e}")
def _patch_accelerate_for_fsdp2_fix(self):
"""Patch accelerate FSDP2 prepare for TensorQuantizer buffers."""
_patch_fsdp2_post_backward()
@@ -459,9 +421,3 @@ class QADTrainer(QATTrainer, KDTrainer):
and
:class:`KDTrainer <modelopt.torch.distill.plugins.huggingface.KDTrainer>`.
"""
def _quantize_model(self):
"""Quantize the model."""
model = self.accelerator.unwrap_model(self.model)
with model.hide_teacher_model(), model.only_student_forward():
return super()._quantize_model()
+59 -8
View File
@@ -27,6 +27,18 @@ BACKEND_CONFIGS = {
# Backends that need gradient checkpointing
GRADIENT_CHECKPOINTING_BACKENDS = {"ddp", "deepspeed"}
# Fast training overrides (short runs with frequent eval)
FAST_TRAIN_ARGS = [
"--model_max_length",
"128",
"--num_train_epochs",
"1.0",
"--save_steps",
"5",
"--eval_steps",
"5",
]
# fmt: off
def _fast_data_args(cache_dir: str) -> list[str]:
@@ -40,10 +52,11 @@ def _fast_data_args(cache_dir: str) -> list[str]:
]
def _run_quantize(extra_cmd_args: list[str], cache_dir: str = ""):
def _run_quantize(config: str, extra_cmd_args: list[str], cache_dir: str = ""):
run_example_command(
[
"python", "quantize.py",
"--config", config,
*_fast_data_args(cache_dir),
*extra_cmd_args,
],
@@ -51,7 +64,7 @@ def _run_quantize(extra_cmd_args: list[str], cache_dir: str = ""):
)
def _run_train(extra_cmd_args: list[str], backend: str = "fsdp2", cache_dir: str = ""):
def _run_train(config: str, extra_cmd_args: list[str], backend: str = "fsdp2", cache_dir: str = ""):
config_file = BACKEND_CONFIGS[backend]
gradient_args = (
["--gradient_checkpointing", "True"]
@@ -63,13 +76,9 @@ def _run_train(extra_cmd_args: list[str], backend: str = "fsdp2", cache_dir: str
"accelerate", "launch",
"--config-file", config_file,
"train.py",
"--config", config,
*_fast_data_args(cache_dir),
"--num_train_epochs", "0.3",
"--learning_rate", "1e-5",
"--per_device_train_batch_size", "2",
"--per_device_eval_batch_size", "2",
"--save_steps", "5",
"--eval_steps", "5",
*FAST_TRAIN_ARGS,
*gradient_args,
*extra_cmd_args,
],
@@ -104,6 +113,7 @@ def test_qwen3_qat_nvfp4(tiny_qwen3_path, tmp_path, backend):
# Step 1: Quantize
_run_quantize(
"configs/train/qat_nvfp4.yaml",
[
"--model_name_or_path", tiny_qwen3_path,
"--recipe", "general/ptq/nvfp4_default-kv_fp8",
@@ -115,6 +125,7 @@ def test_qwen3_qat_nvfp4(tiny_qwen3_path, tmp_path, backend):
# Step 2: QAT
_run_train(
"configs/train/qat_nvfp4.yaml",
[
"--model_name_or_path", str(ptq_output_dir),
"--do_train", "True",
@@ -130,6 +141,7 @@ def test_qwen3_lora_qat_nvfp4(tiny_qwen3_path, tmp_path):
# Step 1: Quantize
_run_quantize(
"configs/train/qat_nvfp4.yaml",
[
"--model_name_or_path", tiny_qwen3_path,
"--recipe", "general/ptq/nvfp4_default-kv_fp8",
@@ -141,6 +153,7 @@ def test_qwen3_lora_qat_nvfp4(tiny_qwen3_path, tmp_path):
# Step 2: LoRA QAT
_run_train(
"configs/train/qat_nvfp4.yaml",
[
"--model_name_or_path", str(ptq_output_dir),
"--do_train", "True",
@@ -152,12 +165,49 @@ def test_qwen3_lora_qat_nvfp4(tiny_qwen3_path, tmp_path):
)
@pytest.mark.parametrize("backend", [
"fsdp2",
"deepspeed",
])
def test_qwen3_qad_nvfp4(tiny_qwen3_path, tmp_path, backend):
ptq_output_dir = tmp_path / "ptq"
qad_output_dir = tmp_path / "qad"
cache_dir = str(tmp_path / "dataset_cache")
# Step 1: Quantize student
_run_quantize(
"configs/train/qad_nvfp4.yaml",
[
"--model_name_or_path", tiny_qwen3_path,
"--recipe", "general/ptq/nvfp4_default-kv_fp8",
"--calib_size", "64",
"--output_dir", str(ptq_output_dir),
],
cache_dir=cache_dir,
)
# Step 2: QAD (quantization-aware distillation)
_run_train(
"configs/train/qad_nvfp4.yaml",
[
"--model_name_or_path", str(ptq_output_dir),
"--do_train", "True",
"--output_dir", str(qad_output_dir),
"--distill", "True",
"--teacher_model", tiny_qwen3_path,
],
backend=backend,
cache_dir=cache_dir,
)
def test_qwen3_qlora_nvfp4(tiny_qwen3_path, tmp_path):
ptq_output_dir = tmp_path / "ptq"
cache_dir = str(tmp_path / "dataset_cache")
# Step 1: Quantize with compression for QLoRA
_run_quantize(
"configs/train/qlora_nvfp4.yaml",
[
"--model_name_or_path", tiny_qwen3_path,
"--recipe", "general/ptq/nvfp4_default-kv_fp8",
@@ -170,6 +220,7 @@ def test_qwen3_qlora_nvfp4(tiny_qwen3_path, tmp_path):
# Step 2: QLoRA training
_run_train(
"configs/train/qlora_nvfp4.yaml",
[
"--model_name_or_path", str(ptq_output_dir),
"--do_train", "True",
@@ -0,0 +1,163 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
from types import SimpleNamespace
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset
transformers = pytest.importorskip("transformers")
TrainingArguments = transformers.TrainingArguments
default_data_collator = transformers.default_data_collator
from transformers.modeling_outputs import CausalLMOutputWithPast
from modelopt.torch.distill.losses import LogitsDistillationLoss
from modelopt.torch.distill.plugins.huggingface import IGNORE_INDEX, KDTrainer
class _TinyCausalLM(nn.Module):
def __init__(self, name=None, events=None):
super().__init__()
self.name = name
self.events = events
self.config = SimpleNamespace(use_cache=False)
self.embed = nn.Embedding(8, 6)
self.lm_head = nn.Linear(6, 8, bias=False)
def forward(self, input_ids, labels=None):
if self.events is not None:
self.events.append((self.name, labels is not None))
logits = self.lm_head(self.embed(input_ids))
loss = None
if labels is not None:
loss = F.cross_entropy(
logits[..., :-1, :].contiguous().view(-1, logits.size(-1)),
labels[..., 1:].contiguous().view(-1),
ignore_index=IGNORE_INDEX,
)
return CausalLMOutputWithPast(loss=loss, logits=logits)
class _ToyDataset(Dataset):
def __init__(self):
self.examples = [
{
"input_ids": torch.tensor([1, 2, 3, 4]),
"labels": torch.tensor([1, 2, 3, 4]),
},
{
"input_ids": torch.tensor([2, 3, 4, 5]),
"labels": torch.tensor([2, IGNORE_INDEX, 4, 5]),
},
]
def __len__(self):
return len(self.examples)
def __getitem__(self, idx):
return self.examples[idx]
def _make_models(events=None):
torch.manual_seed(0)
student = _TinyCausalLM("student", events)
teacher = _TinyCausalLM("teacher", events)
with torch.no_grad():
teacher.lm_head.weight.add_(0.25)
return student, teacher
def _make_batch():
return default_data_collator([_ToyDataset()[0], _ToyDataset()[1]])
def _make_trainer(tmp_path, student, teacher, use_liger_kernel=False):
training_args = TrainingArguments(
output_dir=str(tmp_path),
per_device_eval_batch_size=2,
report_to=[],
use_cpu=True,
)
training_args.use_liger_kernel = use_liger_kernel
return KDTrainer(
model=student,
args=training_args,
eval_dataset=_ToyDataset(),
data_collator=default_data_collator,
distill_args={"teacher_model": teacher},
)
def _manual_kd_loss(student, teacher, batch):
with torch.no_grad():
student_outputs = student(input_ids=batch["input_ids"])
teacher_outputs = teacher(input_ids=batch["input_ids"])
criterion = LogitsDistillationLoss(reduction="none")
per_token_loss = criterion(
student_outputs.logits[..., :-1, :].contiguous().float(),
teacher_outputs.logits[..., :-1, :].contiguous().float(),
)
mask = batch["labels"][..., 1:].contiguous() != IGNORE_INDEX
return (per_token_loss * mask).sum() / mask.sum().clamp(min=1)
def test_training_loss_is_kd_and_skips_ce(tmp_path):
events = []
student, teacher = _make_models(events)
batch = _make_batch()
expected_kd_loss = _manual_kd_loss(student, teacher, batch)
trainer = _make_trainer(tmp_path, student, teacher)
events.clear()
trainer.model.train()
loss = trainer.compute_loss(trainer.model, batch.copy())
assert events == [("student", False), ("teacher", False)]
assert loss.item() == pytest.approx(expected_kd_loss.item())
def test_eval_loss_is_ce_and_kd_is_secondary_metric(tmp_path):
student, teacher = _make_models()
batch = _make_batch()
expected_ce_loss = student(**batch).loss.detach()
trainer = _make_trainer(tmp_path, student, teacher)
metrics = trainer.evaluate()
assert metrics["eval_loss"] == pytest.approx(expected_ce_loss.item())
assert "eval_kd_loss" in metrics
assert metrics["eval_kd_loss"] != pytest.approx(metrics["eval_loss"])
def test_standard_kd_loss_without_labels_uses_mean(tmp_path):
student, teacher = _make_models()
batch = _make_batch()
trainer = _make_trainer(tmp_path, student, teacher)
with torch.no_grad():
outputs = student(input_ids=batch["input_ids"])
teacher_outputs = teacher(input_ids=batch["input_ids"])
trainer._last_teacher_outputs = teacher_outputs
loss = trainer._standard_kd_loss(outputs, labels=None)
expected = LogitsDistillationLoss(reduction="none")(
outputs.logits[..., :-1, :].contiguous().float(),
teacher_outputs.logits[..., :-1, :].contiguous().float(),
).mean()
assert loss.item() == pytest.approx(expected.item())
@@ -0,0 +1,214 @@
# SPDX-FileCopyrightText: Copyright (c) 2024 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.
"""Unit tests for ModelOptHFTrainer lr_config (per-parameter optimizer kwargs)."""
import pytest
import torch
import yaml
from torch import nn
transformers = pytest.importorskip("transformers")
TrainingArguments = transformers.TrainingArguments
from modelopt.torch.opt.plugins.transformers import ModelOptHFTrainer, ModelOptTrainerArguments
class TinyModel(nn.Module):
"""Minimal model with named submodules to exercise fnmatch patterns."""
def __init__(self):
super().__init__()
self.embed_tokens = nn.Embedding(32, 16)
self.self_attn = nn.Linear(16, 16)
self.mlp = nn.Linear(16, 16)
self.lm_head = nn.Linear(16, 32, bias=False)
def forward(self, input_ids, labels=None, **kwargs):
x = self.embed_tokens(input_ids)
x = self.self_attn(x)
x = self.mlp(x)
logits = self.lm_head(x)
loss = None
if labels is not None:
loss = nn.functional.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1))
return type("Out", (), {"loss": loss, "logits": logits})()
@pytest.fixture
def dummy_dataset():
"""Tiny dataset for trainer initialization."""
return [
{"input_ids": torch.randint(0, 32, (8,)), "labels": torch.randint(0, 32, (8,))}
for _ in range(4)
]
def _write_lr_config(tmp_path, cfg: dict) -> str:
path = tmp_path / "lr_config.yaml"
path.write_text(yaml.dump(cfg))
return str(path)
def _make_trainer(
tmp_path,
dummy_dataset,
lr_config_dict=None,
lr_config_path=None,
):
"""Build a ModelOptHFTrainer with the given lr_config."""
training_args = TrainingArguments(
output_dir=str(tmp_path / "output"),
do_train=True,
num_train_epochs=1,
per_device_train_batch_size=2,
learning_rate=1e-3,
weight_decay=0.01,
report_to="none",
use_cpu=True,
)
trainer_args = ModelOptTrainerArguments()
if lr_config_path is not None:
trainer_args.lr_config = lr_config_path
trainer = ModelOptHFTrainer(
model=TinyModel(),
args=training_args,
trainer_args=trainer_args,
lr_config=lr_config_dict,
train_dataset=dummy_dataset,
)
return trainer
class TestLoadLrConfig:
def test_load_basic(self, tmp_path):
cfg = {"*lm_head*": {"lr": 1e-5}, "*mlp*": {"lr": 5e-5}}
path = _write_lr_config(tmp_path, cfg)
loaded = ModelOptHFTrainer.load_lr_config(path)
assert loaded == cfg
def test_load_with_weight_decay_and_betas(self, tmp_path):
cfg = {
"*self_attn*": {"lr": 5e-5, "betas": [0.9, 0.95]},
"*mlp*": {"lr": 5e-5, "weight_decay": 0.05},
"*embed_tokens*": {"lr": 1e-6, "eps": 1e-7},
}
path = _write_lr_config(tmp_path, cfg)
loaded = ModelOptHFTrainer.load_lr_config(path)
assert loaded["*self_attn*"]["betas"] == [0.9, 0.95]
assert loaded["*mlp*"]["weight_decay"] == 0.05
assert loaded["*embed_tokens*"]["eps"] == 1e-7
def test_load_invalid_not_dict(self, tmp_path):
path = tmp_path / "lr_config.yaml"
path.write_text("- item1\n- item2\n")
with pytest.raises(ValueError, match="YAML mapping"):
ModelOptHFTrainer.load_lr_config(str(path))
def test_load_invalid_entry(self, tmp_path):
path = tmp_path / "lr_config.yaml"
path.write_text('"*lm_head*": 0.001\n')
with pytest.raises(ValueError, match="str -> dict"):
ModelOptHFTrainer.load_lr_config(str(path))
class TestCreateOptimizerWithLrConfig:
def test_lr_applied_per_group(self, tmp_path, dummy_dataset):
lr_config = {
"*lm_head*": {"lr": 1e-5},
"*self_attn*": {"lr": 2e-5},
}
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_dict=lr_config)
trainer.create_optimizer()
groups = trainer.optimizer.param_groups
lrs = {g["lr"] for g in groups}
assert 1e-5 in lrs
assert 2e-5 in lrs
def test_weight_decay_override(self, tmp_path, dummy_dataset):
lr_config = {
"*lm_head*": {"lr": 1e-5, "weight_decay": 0.0},
"*mlp*": {"lr": 5e-5, "weight_decay": 0.05},
}
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_dict=lr_config)
trainer.create_optimizer()
for group in trainer.optimizer.param_groups:
if group["lr"] == 1e-5:
assert group["weight_decay"] == 0.0
elif group["lr"] == 5e-5:
assert group["weight_decay"] == 0.05
def test_betas_override(self, tmp_path, dummy_dataset):
lr_config = {
"*self_attn*": {"lr": 5e-5, "betas": [0.9, 0.95]},
}
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_dict=lr_config)
trainer.create_optimizer()
found = False
for group in trainer.optimizer.param_groups:
if group["lr"] == 5e-5:
assert group["betas"] == (0.9, 0.95) or group["betas"] == [0.9, 0.95]
found = True
assert found, "No group found with lr=5e-5"
def test_eps_override(self, tmp_path, dummy_dataset):
lr_config = {
"*embed_tokens*": {"lr": 1e-6, "eps": 1e-7},
}
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_dict=lr_config)
trainer.create_optimizer()
found = False
for group in trainer.optimizer.param_groups:
if group["lr"] == 1e-6:
assert group["eps"] == 1e-7
found = True
assert found, "No group found with lr=1e-6"
def test_lr_config_from_yaml_path(self, tmp_path, dummy_dataset):
cfg = {
"*lm_head*": {"lr": 1e-5, "weight_decay": 0.0},
"*self_attn*": {"lr": 2e-5, "betas": [0.9, 0.95]},
}
path = _write_lr_config(tmp_path, cfg)
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_path=path)
trainer.create_optimizer()
lrs = {g["lr"] for g in trainer.optimizer.param_groups}
assert 1e-5 in lrs
assert 2e-5 in lrs
def test_unmatched_params_use_global_lr(self, tmp_path, dummy_dataset):
lr_config = {
"*lm_head*": {"lr": 1e-5},
}
trainer = _make_trainer(tmp_path, dummy_dataset, lr_config_dict=lr_config)
trainer.create_optimizer()
global_lr = 1e-3 # set in _make_trainer
for group in trainer.optimizer.param_groups:
if group["lr"] != 1e-5:
assert group["lr"] == global_lr
def test_no_lr_config_uses_default(self, tmp_path, dummy_dataset):
trainer = _make_trainer(tmp_path, dummy_dataset)
trainer.create_optimizer()
for group in trainer.optimizer.param_groups:
assert group["lr"] == 1e-3 # global default