mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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.
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
@@ -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,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={
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user