diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index dcf9ff6..82d514d 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -45,6 +45,12 @@ class TrainConfig(BaseConfig): default_factory=list, metadata={"help": "Module types to enable activation checkpointing for."}, ) + compile_mode: Optional[str] = field( + default=None, + metadata={ + "help": "torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None to disable." + }, + ) # checkpoint setting start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."}) diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index 3771ee5..8b7c51a 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -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( diff --git a/scripts/tools/train.py b/scripts/tools/train.py index abbb758..7d9a2d6 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -208,6 +208,13 @@ _START_METHODS = ["spawn", "fork", "forkserver"] default=False, help="Enable activation checkpointing.", ) +@click.option( + "--compile", + "compile_mode", + type=click.Choice(["default", "reduce-overhead", "max-autotune"]), + default=None, + help="torch.compile mode. Omit to disable.", +) @click.option( "--ckpt_interval", type=int, default=5000, help="Steps between checkpoints." ) @@ -495,6 +502,7 @@ def train( ) grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else [] + compile_mode = kwargs.pop("compile_mode", None) collate_fn = None if train_type == "dpo": @@ -532,6 +540,7 @@ def train( val_step=val_step, metrics=metrics, gradient_checkpointing_modules=grad_ckpt_modules, + compile_mode=compile_mode, executor_kwargs=executor_kwargs, extra_kwargs=strategy_kwargs, neftune_alpha=neftune_alpha,