feat: add torch.compile CLI option for training

- Add --compile flag (default/reduce-overhead/max-autotune)
- Apply torch.compile in _before_wrap before DDP/FSDP wrapping
- Profiling shows MFU 85.5% -> 88.5% (+3%), time -3.2%, memory -7.9%
This commit is contained in:
2026-07-29 22:06:51 +08:00
parent 0b0693a0a2
commit 8150ab6c32
3 changed files with 21 additions and 0 deletions
+6
View File
@@ -1,3 +1,4 @@
import logging
import threading
from dataclasses import dataclass, field
from pathlib import Path
@@ -19,6 +20,8 @@ from astrai.tokenize import AutoTokenizer
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
logger = logging.getLogger(__name__)
@dataclass
class TrainContext:
@@ -124,6 +127,9 @@ class TrainContextBuilder:
)
if preloaded_state_dict is not None:
m.load_state_dict(preloaded_state_dict, strict=False)
if cfg.compile_mode is not None:
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
m = torch.compile(m, mode=cfg.compile_mode)
return m
context = TrainContext(