# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file. import math import warnings from dataclasses import dataclass @dataclass class TrainArgs: """Training-related arguments""" save_interval: int | None = 1000 """Number of optimizer steps between saving checkpoints""" log_interval: int = 1 """Number of iterations between logging calls""" global_batch_size: int = 64 """Number of samples between optimizer steps across data-parallel ranks""" micro_batch_size: int = 4 """Number of samples per data-parallel rank""" lr_warmup_steps: int | None = 100 """Number of iterations with learning rate warmup active""" lr_warmup_fraction: float | None = None """The fraction of an epoch to use for learning rate warmup""" epochs: int | None = None """Number of epochs to train on""" # TODO: `pretrain` is the only script using `max_tokens` explicitly. replace it with epoch_size*epochs? max_tokens: int | None = None """Total number of tokens to train on""" max_steps: int | None = None """Limits the number of optimizer steps to run""" max_time: float | None = None """Limits the number of seconds to train for""" max_seq_length: int | None = None """Limits the length of samples""" tie_embeddings: bool | None = None """Whether to tie the embedding weights with the language modeling head weights""" # Optimization args max_norm: float | None = None min_lr: float = 6e-5 lr_schedule: str = "cosine" """Learning rate schedule. Use `cosine`, `onecycle`, or `wsd` for warmup-stable-decay.""" lr_decay_start_fraction: float = 0.7 """For `wsd`, fraction of training after which cosine decay starts.""" compile_model: bool = True """Compile the model with torch.compile after Fabric setup.""" compile_mode: str = "default" """torch.compile mode used when compile_model is enabled.""" mtp_loss_weight: float = 0.0 """Weight of the training-only teacher-forced MTP auxiliary loss.""" nta_margin_loss_weight: float = 0.0 """Optional next-token margin-ranking auxiliary loss weight.""" nta_margin: float = 0.0 """Required correct-logit margin over the strongest wrong token for the NTA auxiliary loss.""" nta_margin_error_only: bool = False """Apply NTA margin ranking only where the target is not the current top-1 prediction.""" channel_memory_lr_mult: float = 1.0 """Learning-rate multiplier for TileRoutedChannelMemoryDSwiGLUMLP gain parameters.""" grouped_mlp_lr_mult: float = 1.0 """Learning-rate multiplier for full-active grouped MLP projection tensors.""" grouped_mlp_weight_decay_mult: float = 1.0 """Weight-decay multiplier for full-active grouped MLP projection tensors.""" def __post_init__(self) -> None: if self.lr_warmup_fraction and self.lr_warmup_steps: raise ValueError( "Can't provide both `--train.lr_warmup_fraction` and `--train.lr_warmup_steps`. Choose one." ) if self.lr_warmup_fraction and not (0 <= self.lr_warmup_fraction <= 1): raise ValueError("`--train.lr_warmup_fraction` must be between 0 and 1.") if self.lr_warmup_steps and self.max_steps and (self.lr_warmup_steps >= self.max_steps): warnings.warn( "`--train.lr_warmup_steps` should be less than `--train.max_steps`." f" Got {self.lr_warmup_steps} lr_warmup_steps and {self.max_steps} max_steps.", UserWarning, ) if self.lr_schedule not in {"cosine", "onecycle", "wsd"}: raise ValueError("`--train.lr_schedule` must be either 'cosine', 'onecycle', or 'wsd'.") if not (0.0 <= self.lr_decay_start_fraction <= 1.0): raise ValueError("`--train.lr_decay_start_fraction` must be between 0 and 1.") if self.compile_mode not in {"default", "reduce-overhead", "max-autotune"}: raise ValueError("`--train.compile_mode` must be 'default', 'reduce-overhead', or 'max-autotune'.") if self.mtp_loss_weight < 0.0: raise ValueError("`--train.mtp_loss_weight` must be non-negative.") if self.nta_margin_loss_weight < 0.0: raise ValueError("`--train.nta_margin_loss_weight` must be non-negative.") if self.nta_margin < 0.0: raise ValueError("`--train.nta_margin` must be non-negative.") if self.channel_memory_lr_mult <= 0.0: raise ValueError("`--train.channel_memory_lr_mult` must be positive.") if self.grouped_mlp_lr_mult <= 0.0: raise ValueError("`--train.grouped_mlp_lr_mult` must be positive.") if self.grouped_mlp_weight_decay_mult < 0.0: raise ValueError("`--train.grouped_mlp_weight_decay_mult` must be non-negative.") def gradient_accumulation_iters(self, devices: int, num_nodes: int = 1) -> int: """Number of iterations between gradient synchronizations""" gradient_accumulation_iters = self.batch_size(devices, num_nodes) // self.micro_batch_size assert gradient_accumulation_iters > 0 return gradient_accumulation_iters def batch_size(self, devices: int, num_nodes: int = 1) -> int: """Number of samples between optimizer steps per data-parallel rank""" batch_size = self.global_batch_size // (devices * num_nodes) assert batch_size > 0 return batch_size def warmup_iters(self, devices: int, num_nodes: int, max_iters: int, train_dataloader) -> int: """Number of iterations to warm up the learning rate.""" if self.lr_warmup_fraction: return min(max_iters, math.ceil(self.lr_warmup_fraction * len(train_dataloader))) if self.lr_warmup_steps: return min(max_iters, self.lr_warmup_steps * self.gradient_accumulation_iters(devices, num_nodes)) return 0 @dataclass class EvalArgs: """Evaluation-related arguments""" interval: int = 600 """Number of optimizer steps between evaluation calls""" max_new_tokens: int | None = None """Number of tokens to generate""" max_iters: int = 100 """Number of iterations""" initial_validation: bool = False """Whether to evaluate on the validation set at the beginning of the training""" final_validation: bool = True """Whether to evaluate on the validation set at the end of the training""" evaluate_example: str | int = "first" """How to pick an example instruction to evaluate periodically during training. Can be "first", "random", or an integer index to pick a specific example.""" @dataclass class LogArgs: """Logging-related arguments. Different loggers use different fields.""" # === WandB Fields === project: str | None = None """WandB project name""" run: str | None = None """WandB run name (defaults to generated name)""" group: str | None = None """WandB group name""" # === LitLogger Fields (Lightning.ai) === teamspace: str | None = None """Teamspace name where charts and artifacts will appear""" metadata: dict | None = None """Extra metadata to associate with the experiment as tags""" log_model: bool = False """If True, automatically log model checkpoints as artifacts""" save_logs: bool = True """If True, capture and upload terminal logs""" checkpoint_name: str | None = None """Override the base name for logged checkpoints"""