| |
| 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""" |
| |
| 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""" |
|
|
| |
| 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.""" |
|
|
| |
| 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""" |
|
|
| |
| 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""" |
|
|