feat: 新增 GradientCheckpointingCallback

- TrainConfig.gradient_checkpointing_modules 指定模块类型
- apply 递归遍历,兼容 DDP,不硬编码模型结构
- modules=None 时静默跳过,零开销
This commit is contained in:
2026-05-17 18:21:05 +08:00
parent 7621f05d3f
commit 2c2697390d
4 changed files with 166 additions and 2 deletions
+4
View File
@@ -39,6 +39,10 @@ class TrainConfig(BaseConfig):
max_grad_norm: float = field(
default=1.0, metadata={"help": "Maximum gradient norm."}
)
gradient_checkpointing_modules: list = field(
default_factory=list,
metadata={"help": "Module types to enable activation checkpointing for."},
)
# checkpoint setting
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})