Metris's picture
Upload 78 files
236083b verified
Raw
History Blame Contribute Delete
51.4 kB
# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file.
import math
import os
import pprint
import time
import warnings
from dataclasses import asdict
from datetime import timedelta
from functools import partial
from pathlib import Path
from typing import Literal
import lightning as L
import torch
import torch.nn as nn
from lightning.fabric.strategies import FSDPStrategy
from lightning.fabric.utilities.throughput import ThroughputMonitor, measure_flops
from torch.utils.data import DataLoader
from torchmetrics.aggregation import RunningMean
from litgpt import Tokenizer
from litgpt.args import EvalArgs, LogArgs, TrainArgs
from litgpt.config import name_to_config
from litgpt.constants import _TORCH_EQUAL_2_7, _TORCH_EQUAL_2_8
from litgpt.data import DataModule, TinyLlama
from litgpt.model import (
GPT,
AdaptiveDSwiGLUMLP,
Block,
BlockSparseAdaptiveDSwiGLUMLP,
CausalSelfAttention,
Config,
DSwiGLUMLP,
LLaMAMLP,
TileRoutedChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedChannelMemoryDSwiGLUMLP,
TileRoutedChannelStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelStateMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLP,
TileRoutedGroupStateMemoryDSwiGLUMLP,
TileRoutedKeyedChannelMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLPStaticA2,
TileRoutedDSwiGLUMLPStaticGPTS,
)
from litgpt.optim import AdaptiveRowTAdamW, HybridMuonAdamW, OfficialMuonWithAuxAdam, RMNPGrouped
from litgpt.parser_config import save_hyperparameters
from litgpt.types import LoggerChoice
from litgpt.utils import (
CycleIterator,
capture_hparams,
check_nvlink_connectivity,
choose_logger,
chunked_cross_entropy,
copy_config_files,
extend_checkpoint_dir,
find_resume_path,
get_default_supported_precision,
init_out_dir,
instantiate_torch_optimizer,
num_parameters,
parse_devices,
reset_parameters,
save_config,
)
torch.set_float32_matmul_precision("high")
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.backends.cuda.enable_flash_sdp(True)
torch.backends.cuda.enable_mem_efficient_sdp(True)
def setup(
model_name: str,
model_config: Config | None = None,
out_dir: Path = Path("out/pretrain"),
precision: Literal["bf16-true", "bf16-mixed", "32-true", None] = None,
initial_checkpoint_dir: Path | None = None,
resume: bool | Literal["auto"] | Path = False,
data: DataModule | None = None,
train: TrainArgs = TrainArgs(
save_interval=1000,
log_interval=1,
global_batch_size=512,
micro_batch_size=4,
max_tokens=int(3e12), # 3 trillion
max_norm=1.0,
min_lr=4e-5,
lr_warmup_steps=2000,
tie_embeddings=False,
),
eval: EvalArgs = EvalArgs(interval=1000, max_iters=100),
log: LogArgs = LogArgs(),
optimizer: str | dict = "AdamW",
devices: int | str = "auto",
num_nodes: int = 1,
tokenizer_dir: Path | None = None,
logger_name: LoggerChoice = "tensorboard",
seed: int = 42,
):
"""Pretrain a model.
Arguments:
model_name: The name of the model to pretrain. Choose from names in ``litgpt.config``. Use "list" to list the supported models.
model_config: A ``litgpt.Config`` object to define the model architecture. Mutually exclusive with
``model_config``. Overrides the `model_name` if specified.
out_dir: Directory in which to save checkpoints and logs. If running in a Lightning Studio Job, look for it in
/teamspace/jobs/<job-name>/share.
precision: The precision to use for finetuning. Determines a compatible precision setting by default.
initial_checkpoint_dir: Optional path to a checkpoint directory to initialize the model from.
Useful for continued pretraining. Mutually exclusive with ``resume``.
resume: Path to a checkpoint directory to resume from in case training was interrupted, or ``True`` to resume
from the latest checkpoint in ``out_dir``. An error will be raised if no checkpoint is found. Passing
``'auto'`` will resume from the latest checkpoint but not error if no checkpoint exists.
data: Data-related arguments. If not provided, the default is ``litgpt.data.TinyLlama``.
train: Training-related arguments. See ``litgpt.args.TrainArgs`` for details.
eval: Evaluation-related arguments. See ``litgpt.args.EvalArgs`` for details.
optimizer: An optimizer name (such as "AdamW") or config.
devices: How many devices/GPUs to use. Uses all GPUs by default.
num_nodes: How many nodes the code is being run on.
tokenizer_dir: Optional path to the tokenizer dir that was used for preprocessing the dataset. Only some data
module require this.
logger_name: The name of the logger to send metrics to.
seed: The random seed to use for reproducibility.
"""
if model_name == "list":
available_models = "\n".join(sorted(name_to_config))
print(f"Available values:\n{available_models}")
quit()
if initial_checkpoint_dir is not None:
initial_checkpoint_dir = extend_checkpoint_dir(initial_checkpoint_dir)
if tokenizer_dir is not None:
tokenizer_dir = extend_checkpoint_dir(tokenizer_dir)
if model_config is None:
# Support both model_name options: meta-llama/Meta-Llama-3-8B & Meta-Llama-3-8B
try:
model_config = Config.from_name(model_name)
except ValueError:
print(f"Model name {model_name} is not supported.\n")
available_models = "\n".join(sorted(name_to_config))
print(f"Available values:\n{available_models}")
quit()
hparams = capture_hparams()
data = TinyLlama() if data is None else data
config = Config.from_name(model_name) if model_config is None else model_config
precision = precision or get_default_supported_precision(training=True)
devices = parse_devices(devices)
out_dir = init_out_dir(out_dir)
# in case the dataset requires the Tokenizer
tokenizer = Tokenizer(tokenizer_dir) if tokenizer_dir is not None else None
logger = choose_logger(
logger_name,
out_dir,
name=f"pretrain-{config.name}",
resume=bool(resume),
log_interval=train.log_interval,
log_args=asdict(log),
)
if devices * num_nodes > 1:
strategy = FSDPStrategy(auto_wrap_policy={Block}, state_dict_type="full", sharding_strategy="HYBRID_SHARD")
else:
strategy = "auto"
fabric = L.Fabric(devices=devices, num_nodes=num_nodes, strategy=strategy, precision=precision, loggers=[logger])
if torch.cuda.is_available() and devices > 1:
check_nvlink_connectivity(fabric)
fabric.launch()
fabric.print(pprint.pformat(hparams))
if logger_name in ("tensorboard", "wandb", "mlflow"):
fabric.logger.log_hyperparams(hparams)
main(
fabric=fabric,
devices=devices,
num_nodes=num_nodes,
seed=seed,
initial_checkpoint_dir=initial_checkpoint_dir,
resume=resume,
config=config,
data=data,
out_dir=out_dir,
tokenizer_dir=tokenizer_dir,
tokenizer=tokenizer,
train=train,
eval=eval,
optimizer=optimizer,
)
def main(
fabric: L.Fabric,
devices: int,
seed: int,
initial_checkpoint_dir: Path | None,
resume: bool | Literal["auto"] | Path,
config: Config,
data: DataModule,
out_dir: Path,
tokenizer_dir: Path | None,
tokenizer: Tokenizer | None,
train: TrainArgs,
eval: EvalArgs,
optimizer: str | dict,
num_nodes: int = 1,
) -> None:
validate_args(train, eval, initial_checkpoint_dir, resume)
if fabric.global_rank == 0:
out_dir.mkdir(parents=True, exist_ok=True)
fabric.seed_everything(seed) # same seed for every process to init model (FSDP)
t0 = time.perf_counter()
with fabric.init_module(empty_init=True):
model = GPT(config)
initialize_weights(fabric, model, n_layer=config.n_layer, n_embd=config.n_embd)
if train.tie_embeddings and not config.mamba3_hierarchical_vocab:
model.transformer.wte.weight = model.lm_head.weight
if train.max_seq_length:
model.max_seq_length = train.max_seq_length
if config.sliding_window_size is not None:
model.prepare_sliding_window_masks(model.max_seq_length, fabric.device)
fabric.print(f"Time to instantiate model: {time.perf_counter() - t0:.02f} seconds.")
fabric.print(f"Total parameters: {num_parameters(model):,}")
if bool(train.mtp_loss_weight) != bool(config.mtp_num_layers):
raise ValueError(
"Paper-faithful MTP requires both model.mtp_num_layers=1 and "
"train.mtp_loss_weight>0, or both disabled."
)
if train.compile_model:
if config.kda_enabled:
# FLA stores a distinct integer ``layer_idx`` on each KDA module.
# Dynamo otherwise specializes the same forward on every layer,
# hits its recompile limit after eight variants, and leaves later
# layers on a slower fallback. The attribute is only a cache key;
# allowing it to remain symbolic preserves the exact computation.
torch._dynamo.config.allow_unspec_int_on_nn_module = True
compile_kwargs = {}
if train.compile_mode != "default":
compile_kwargs["mode"] = train.compile_mode
fabric.print(f"Compiling model with torch.compile({compile_kwargs or 'default'})")
model = torch.compile(model, **compile_kwargs)
model = fabric.setup(model)
if train.mtp_loss_weight:
model.mark_forward_method('mtp_forward')
extra_kwargs = {"fused": fabric.device.type == "cuda"}
if isinstance(optimizer, dict) and "fused" in optimizer.get("init_args", {}):
# Respect an explicit optimizer setting. This is required for Gefen's
# documented pure-PyTorch fallback when a matching nvcc is unavailable.
extra_kwargs = {}
optimizer_class_path = optimizer.get("class_path", "") if isinstance(optimizer, dict) else ""
if optimizer_class_path == "litgpt.optim.OfficialMuonWithAuxAdam":
if train.channel_memory_lr_mult != 1.0:
raise ValueError("`train.channel_memory_lr_mult` is not supported with OfficialMuonWithAuxAdam.")
optimizer = instantiate_official_muon_optimizer(optimizer, model)
elif optimizer_class_path == "litgpt.optim.AdaptiveRowTAdamW":
if train.channel_memory_lr_mult != 1.0:
raise ValueError("`train.channel_memory_lr_mult` is not supported with AdaptiveRowTAdamW.")
optimizer = instantiate_adaptive_row_tadamw(optimizer, model)
elif optimizer_class_path == "litgpt.optim.RMNPGrouped":
if train.channel_memory_lr_mult != 1.0:
raise ValueError("`train.channel_memory_lr_mult` is not supported with RMNPGrouped.")
optimizer = instantiate_rmnp_optimizer(optimizer, model)
elif config.mlp_matrix_optimizer == "himuon_tile32":
if train.channel_memory_lr_mult != 1.0:
raise ValueError("`train.channel_memory_lr_mult` is only supported for torch optimizers.")
optimizer = instantiate_hybrid_mlp_optimizer(optimizer, model, **extra_kwargs)
else:
if optimizer_class_path == "gefen.Gefen":
import gefen.gefen as gefen_impl
backend = os.environ.get("GEFEN_FIND_PERIOD_BACKEND")
if backend:
if backend not in {"cpu", "gpu", "cuda_kernel"}:
raise ValueError(f"Unsupported GEFEN_FIND_PERIOD_BACKEND={backend!r}")
gefen_impl.FIND_PERIOD_BACKEND = backend
fabric.print(f"Gefen period backend: {backend}")
if os.environ.get("GEFEN_FUSE_HISTOGRAM_FOR_EXACT") == "0":
gefen_impl.FUSE_HISTOGRAM_FOR_EXACT = False
fabric.print("Gefen exact histogram fusion: disabled")
# Gefen accepts named parameters and uses their names in its learned
# partition diagnostics. Preserve that information for grouped tensors.
model_parameters = model.named_parameters() if optimizer_class_path == "gefen.Gefen" else model.parameters()
if config.rwkv7_enabled:
if optimizer_class_path == "gefen.Gefen":
raise ValueError("RWKV-7 parameter rules are not implemented for Gefen.")
if not isinstance(optimizer, dict):
raise ValueError("RWKV-7 optimizer groups require optimizer.init_args.")
init_args = optimizer.get("init_args", {})
base_lr = init_args.get("lr")
weight_decay = init_args.get("weight_decay", 0.0)
if base_lr is None:
raise ValueError("RWKV-7 optimizer groups require optimizer.init_args.lr.")
decay_parameters = []
no_decay_parameters = []
decay_lr2_parameters = []
decay_suffixes = (
".receptance.weight",
".key.weight",
".value.weight",
".output.weight",
".local_screen.query_weight",
".local_screen.key_weight",
".local_screen.value_weight",
".local_screen.gate_weight",
".local_screen.output_weight",
)
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
if name.endswith(".w0"):
decay_lr2_parameters.append(parameter)
elif (
name in {"transformer.wte.weight", "lm_head.weight"}
or name.endswith(decay_suffixes)
):
decay_parameters.append(parameter)
else:
no_decay_parameters.append(parameter)
model_parameters = [
{"params": no_decay_parameters, "weight_decay": 0.0},
{"params": decay_parameters, "weight_decay": float(weight_decay)},
{
"params": decay_lr2_parameters,
"weight_decay": 0.0,
"lr": 2.0 * float(base_lr),
},
]
fabric.print(
"RWKV-7 optimizer groups: "
f"{len(decay_parameters)} large tensors at weight_decay={float(weight_decay):g}; "
f"{len(no_decay_parameters)} tensors without decay; "
f"{len(decay_lr2_parameters)} time-decay tensors at 2x LR"
)
elif config.mamba3_enabled:
if optimizer_class_path == "gefen.Gefen":
raise ValueError("Mamba-3 parameter rules are not implemented for Gefen.")
if not isinstance(optimizer, dict):
raise ValueError("Mamba-3 optimizer groups require optimizer.init_args.")
init_args = optimizer.get("init_args", {})
weight_decay = float(init_args.get("weight_decay", 0.0))
decay_parameters = []
no_decay_parameters = []
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
no_weight_decay = bool(getattr(parameter, "_no_weight_decay", False))
is_norm_or_bias = (
parameter.ndim < 2
or ".norm." in name
or name.endswith(".bias")
or "basis_coefficients" in name
)
(no_decay_parameters if no_weight_decay or is_norm_or_bias else decay_parameters).append(parameter)
model_parameters = [
{"params": no_decay_parameters, "weight_decay": 0.0},
{"params": decay_parameters, "weight_decay": weight_decay},
]
fabric.print(
"Mamba-3 optimizer groups: "
f"{len(decay_parameters)} matrices at weight_decay={weight_decay:g}; "
f"{len(no_decay_parameters)} state, norm, and bias tensors without decay"
)
use_memory_group = train.channel_memory_lr_mult != 1.0
use_grouped_mlp_group = (
train.grouped_mlp_lr_mult != 1.0
or train.grouped_mlp_weight_decay_mult != 1.0
)
controller_parameter_names = (
"cvr_energy_coeff",
"cvr_agreement_coeff",
"caet_temperature_raw",
"adaptive_commit_router",
)
use_controller_group = config.sandwich_adaptive_commit > 0 or config.mlp_class_name in (
"TileRoutedContextualValueRotationEulerGLUMLP",
"TileRoutedContextAdaptiveEulerTemperatureGLUMLP",
)
if (
use_memory_group or use_grouped_mlp_group or use_controller_group
) and not (config.rwkv7_enabled or config.mamba3_enabled):
if not isinstance(optimizer, dict):
raise ValueError("Optimizer parameter groups require optimizer.init_args.")
init_args = optimizer.get("init_args", {})
base_lr = init_args.get("lr")
base_weight_decay = init_args.get("weight_decay")
if base_lr is None:
raise ValueError("Optimizer parameter groups require optimizer.init_args.lr.")
if use_grouped_mlp_group and base_weight_decay is None:
raise ValueError("Grouped MLP decay requires optimizer.init_args.weight_decay.")
base_parameters = []
memory_parameters = []
grouped_mlp_parameters = []
controller_parameters = []
memory_parameter_names = (
"channel_memory_gain_raw",
"group_state_memory_gain_raw",
"keyed_memory_key",
"keyed_memory_value_raw",
"grouped_mlp_mix_alpha",
"grouped_mlp_mixer",
"shared_alpha",
"shared_path",
)
grouped_mlp_parameter_names = (
".mlp.up_weight",
".mlp.value_weight",
".mlp.down_weight",
)
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
if use_controller_group and name.endswith(controller_parameter_names):
controller_parameters.append(parameter)
elif use_grouped_mlp_group and name.endswith(grouped_mlp_parameter_names):
grouped_mlp_parameters.append(parameter)
elif use_memory_group and any(
memory_name in name for memory_name in memory_parameter_names
):
memory_parameters.append(parameter)
else:
base_parameters.append(parameter)
if use_memory_group and not memory_parameters:
raise ValueError(
"`train.channel_memory_lr_mult` was set, but no trainable channel memory gains were found."
)
if use_grouped_mlp_group and not grouped_mlp_parameters:
raise ValueError(
"Grouped MLP optimizer settings were set, but no grouped MLP projection tensors were found."
)
if use_controller_group and not controller_parameters:
raise ValueError("Adaptive model was selected, but no controller parameters were found.")
model_parameters = [{"params": base_parameters}]
if use_controller_group:
fabric.print(
"Adaptive controller group: "
f"{len(controller_parameters)} tensors at lr={float(base_lr):g}, weight_decay=0"
)
model_parameters.append(
{"params": controller_parameters, "lr": float(base_lr), "weight_decay": 0.0}
)
if use_memory_group:
memory_lr = float(base_lr) * float(train.channel_memory_lr_mult)
fabric.print(
"Channel memory LR multiplier: "
f"{train.channel_memory_lr_mult:g}x ({len(memory_parameters)} tensors at lr={memory_lr:g})"
)
model_parameters.append({"params": memory_parameters, "lr": memory_lr})
if use_grouped_mlp_group:
grouped_mlp_lr = float(base_lr) * float(train.grouped_mlp_lr_mult)
grouped_mlp_weight_decay = float(base_weight_decay) * float(train.grouped_mlp_weight_decay_mult)
fabric.print(
"Grouped MLP AdamW group: "
f"{len(grouped_mlp_parameters)} tensors at lr={grouped_mlp_lr:g}, "
f"weight_decay={grouped_mlp_weight_decay:g}"
)
model_parameters.append(
{
"params": grouped_mlp_parameters,
"lr": grouped_mlp_lr,
"weight_decay": grouped_mlp_weight_decay,
}
)
optimizer = instantiate_torch_optimizer(optimizer, model_parameters, **extra_kwargs)
optimizer = fabric.setup_optimizers(optimizer)
train_dataloader, val_dataloader = get_dataloaders(fabric, data, tokenizer, train, model.max_seq_length)
train_dataloader, val_dataloader = fabric.setup_dataloaders(train_dataloader, val_dataloader)
if initial_checkpoint_dir:
fabric.load_raw(initial_checkpoint_dir / "lit_model.pth", model)
state = {
"model": model,
"optimizer": optimizer,
"train_dataloader": train_dataloader,
"iter_num": 0,
"step_count": 0,
}
resume = find_resume_path(resume, out_dir)
if resume:
fabric.print(f"Resuming training from {resume}")
fabric.load(resume, state)
train_time = time.perf_counter()
# work around PyTorch issue https://github.com/pytorch/pytorch/issues/152162
# which does not like the lazy initialization to be called in dynamo.
# TODO: Happens with PyTorch 2.7+
if (
(_TORCH_EQUAL_2_7 or _TORCH_EQUAL_2_8)
and (model._forward_module.__class__.__name__ == "OptimizedModule")
and (model._forward_module._orig_mod.__class__.__name__ == "FullyShardedDataParallel")
):
from torch.distributed.fsdp._runtime_utils import _root_pre_forward
_root_pre_forward(model._forward_module._orig_mod, model._forward_module._orig_mod, [], {})
fit(
fabric=fabric,
devices=devices,
num_nodes=num_nodes,
state=state,
train_dataloader=train_dataloader,
val_dataloader=val_dataloader,
out_dir=out_dir,
tokenizer_dir=tokenizer_dir,
train=train,
eval=eval,
)
# Save final checkpoint
save_checkpoint(fabric, state, tokenizer_dir, out_dir / "final" / "lit_model.pth")
total_tokens = state["iter_num"] * train.micro_batch_size * model.max_seq_length * fabric.world_size
# Print formatted output
separator = "-" * 40
fabric.print(separator)
fabric.print("| Performance")
fabric.print(f"| - Total tokens : {total_tokens:,}")
training_elapsed = time.perf_counter() - train_time
fabric.print(f"| - Training Time : {training_elapsed:.2f} s")
fabric.print(f"| - Tok/sec : {total_tokens / training_elapsed:.2f} tok/s")
fabric.print("| " + "-" * 40)
if fabric.device.type == "cuda":
memory_used = torch.cuda.max_memory_allocated() / 1e9
fabric.print("| Memory Usage")
fabric.print(f"| - Memory Used : {memory_used:.2f} GB")
fabric.print(separator)
def instantiate_hybrid_mlp_optimizer(optimizer_config, model: nn.Module, **kwargs) -> HybridMuonAdamW:
if isinstance(optimizer_config, str):
init_args = {"lr": 1e-3}
else:
init_args = dict(optimizer_config.get("init_args", {}))
init_args.update({key: value for key, value in kwargs.items() if key in {"lr", "weight_decay", "betas", "eps"}})
init_args.pop("fused", None)
muon_params = []
adamw_params = []
for name, param in model.named_parameters():
is_mlp_matrix = ".mlp." in name and param.ndim >= 2
if is_mlp_matrix:
muon_params.append(param)
else:
adamw_params.append(param)
param_groups = []
if adamw_params:
param_groups.append({"params": adamw_params, "use_muon": False})
if muon_params:
param_groups.append({"params": muon_params, "use_muon": True, "tile_size": 32})
return HybridMuonAdamW(param_groups, **init_args)
def instantiate_rmnp_optimizer(optimizer_config, model: nn.Module) -> RMNPGrouped:
"""Map LitGPT parameter names onto the authors' RMNP/Adam branches."""
if not isinstance(optimizer_config, dict):
raise ValueError("RMNPGrouped requires optimizer.init_args.")
init_args = dict(optimizer_config.get("init_args", {}))
init_args.pop("fused", None)
rmnp_params = []
adam_params = []
for name, param in model.named_parameters():
# The authors assign every hidden matrix to RMNP and keep embeddings,
# lm_head, vectors, and scalars on Adam. LitGPT calls its embedding
# `transformer.wte`, so exclude that exact semantic equivalent too.
is_embedding_or_head = "embed" in name or ".wte." in name or name.endswith("lm_head.weight")
if param.ndim >= 2 and not is_embedding_or_head:
rmnp_params.append(param)
else:
adam_params.append(param)
param_groups = []
if rmnp_params:
param_groups.append({"params": rmnp_params, "is_rmnp": True})
if adam_params:
param_groups.append({"params": adam_params, "is_rmnp": False})
if not rmnp_params:
raise ValueError("RMNPGrouped found no hidden matrix parameters.")
return RMNPGrouped(param_groups, **init_args)
def instantiate_adaptive_row_tadamw(optimizer_config, model: nn.Module) -> AdaptiveRowTAdamW:
"""Apply adaptive row correction only to hidden matrix parameters."""
if not isinstance(optimizer_config, dict):
raise ValueError("AdaptiveRowTAdamW requires optimizer.init_args.")
init_args = dict(optimizer_config.get("init_args", {}))
init_args.pop("fused", None)
matrix_params = []
base_params = []
for name, param in model.named_parameters():
is_embedding_or_head = "embed" in name or ".wte." in name or name.endswith("lm_head.weight")
if param.ndim >= 2 and not is_embedding_or_head:
matrix_params.append(param)
else:
base_params.append(param)
param_groups = []
if matrix_params:
param_groups.append({"params": matrix_params, "use_row_blend": True})
if base_params:
param_groups.append({"params": base_params, "use_row_blend": False})
return AdaptiveRowTAdamW(param_groups, **init_args)
def instantiate_official_muon_optimizer(optimizer_config, model: nn.Module) -> OfficialMuonWithAuxAdam:
"""Map hidden matrices to Muon and embeddings/head/vectors to auxiliary Adam."""
if not isinstance(optimizer_config, dict):
raise ValueError("OfficialMuonWithAuxAdam requires optimizer.init_args.")
init_args = dict(optimizer_config.get("init_args", {}))
init_args.pop("fused", None)
muon_params = []
per_head_muon_params: list[tuple[nn.Parameter, int]] = []
adam_params = []
for name, param in model.named_parameters():
is_embedding_or_head = "embed" in name or ".wte." in name or name.endswith("lm_head.weight")
if param.ndim >= 2 and not is_embedding_or_head:
# Kimi K3 section 2.5 orthogonalizes each Q/K/V head block
# independently instead of coupling all heads in one matrix.
# FLA's KDA projections lay heads contiguously on output axis 0.
is_kda_qkv = (
".attn.attn." in name
and name.endswith(("q_proj.weight", "k_proj.weight", "v_proj.weight"))
and param.ndim == 2
)
if is_kda_qkv:
num_heads = int(model.config.kda_num_heads)
per_head_muon_params.append((param, num_heads))
else:
muon_params.append(param)
else:
adam_params.append(param)
param_groups = []
if muon_params:
param_groups.append({"params": muon_params, "use_muon": True})
for param, num_heads in per_head_muon_params:
param_groups.append(
{
"params": [param],
"use_muon": True,
"muon_head_splits": num_heads,
}
)
if adam_params:
param_groups.append({"params": adam_params, "use_muon": False})
return OfficialMuonWithAuxAdam(param_groups, **init_args)
def fit(
fabric: L.Fabric,
devices: int,
state: dict,
train_dataloader: DataLoader,
val_dataloader: DataLoader,
out_dir: Path,
tokenizer_dir: Path | None,
train: TrainArgs,
eval: EvalArgs,
num_nodes: int = 1,
) -> None:
model = state["model"]
optimizer = state["optimizer"]
if eval.initial_validation:
val_loss = validate(fabric, model, val_dataloader, max_iters=eval.max_iters)
val_loss = f"{val_loss:.3f}"
else:
fabric.print("Verifying settings ...")
validate(fabric, model, val_dataloader, max_iters=2, verbose=False) # sanity check
val_loss = "n/a"
throughput = ThroughputMonitor(fabric, window_size=5)
if (
model.config.mlp_class_name == "BlockSparseAdaptiveDSwiGLUMLP"
or model.config.multiscreen_enabled
or model.config.mamba3_enabled
or model.config.rwkv7_enabled
or model.config.kda_enabled
):
measured_flops = 0
fabric.print("Measured TFLOPs: skipped for custom recurrent/screening kernel")
else:
with torch.device("meta"):
meta_model = GPT(model.config)
x = torch.randint(0, 1, (train.micro_batch_size, meta_model.max_seq_length))
model_fwd = lambda: meta_model(x) # noqa: F821
model_loss = lambda y: chunked_cross_entropy(y, x, chunk_size=0) # noqa: F821
measured_flops = measure_flops(meta_model, model_fwd, model_loss)
fabric.print(f"Measured TFLOPs: {measured_flops * fabric.world_size / 1e12:.2f}")
del meta_model, x
max_tokens_per_device = train.max_tokens // fabric.world_size
tokens_per_iter = train.micro_batch_size * model.max_seq_length
max_iters = max_tokens_per_device // tokens_per_iter
log_iter_interval = train.log_interval * train.gradient_accumulation_iters(devices, num_nodes)
initial_iter = state["iter_num"]
train_iterator = CycleIterator(train_dataloader)
running_loss = RunningMean(window=train.gradient_accumulation_iters(devices, num_nodes), sync_on_compute=False).to(
fabric.device
)
fabric.barrier()
total_t0 = time.perf_counter()
warmup_iters = train.warmup_iters(devices, num_nodes, max_iters, train_dataloader)
sliding_window_enabled: bool | None = None
for train_data in train_iterator:
if state["iter_num"] >= max_iters:
break
if model.config.sliding_window_curriculum == "qg_global_qg":
progress = (state["iter_num"] + 1) / max_iters
next_sliding_window_enabled = (
progress <= model.config.sliding_window_curriculum_early_fraction
or progress >= model.config.sliding_window_curriculum_late_fraction
)
if next_sliding_window_enabled != sliding_window_enabled:
changed = model.set_sliding_window_enabled(next_sliding_window_enabled)
phase = "QG-windowed" if next_sliding_window_enabled else "all-global"
fabric.print(
f"Sliding-window curriculum -> {phase} at progress={progress:.6f} "
f"(attention modules changed={changed})"
)
sliding_window_enabled = next_sliding_window_enabled
fade_start = model.config.depth_memory_fade_start_fraction
if model.config.depth_memory_mode == "adaptive" and fade_start < 1.0:
progress = (state["iter_num"] + 1) / max_iters
if progress <= fade_start:
depth_memory_scale = 1.0
else:
fade_progress = (progress - fade_start) / (1.0 - fade_start)
depth_memory_scale = 0.5 * (1.0 + math.cos(math.pi * fade_progress))
model.depth_memory_scale.fill_(depth_memory_scale)
# determine and set the learning rate for this iteration
lr = get_lr(
optimizer.defaults["lr"],
state["iter_num"],
warmup_iters,
max_iters,
train.min_lr,
train.lr_schedule,
train.lr_decay_start_fraction,
)
for param_group in optimizer.param_groups:
group_base_lr = param_group.get("schedule_base_lr")
if group_base_lr is None:
param_group["lr"] = lr
else:
reference_lr = optimizer.defaults["lr"]
group_min_lr = train.min_lr * group_base_lr / reference_lr
param_group["lr"] = get_lr(
group_base_lr,
state["iter_num"],
warmup_iters,
max_iters,
group_min_lr,
train.lr_schedule,
train.lr_decay_start_fraction,
)
state["iter_num"] += 1
iter_t0 = time.perf_counter()
input_ids = train_data[:, 0 : model.max_seq_length].contiguous().long()
targets = train_data[:, 1 : (model.max_seq_length + 1)].contiguous().long()
is_accumulating = state["iter_num"] % train.gradient_accumulation_iters(devices, num_nodes) != 0
with fabric.no_backward_sync(model, enabled=is_accumulating):
use_multiscreen_chunked_loss = (
model.config.multiscreen_chunked_lm_loss
and train.mtp_loss_weight == 0.0
and train.nta_margin_loss_weight == 0.0
)
if use_multiscreen_chunked_loss:
loss = model(input_ids, targets=targets)
logits = None
elif (
train.mtp_loss_weight > 0.0
and model.config.multiscreen_chunked_lm_loss
):
loss, hidden = model(
input_ids,
return_hidden=True,
targets=targets,
)
logits = None
elif train.mtp_loss_weight > 0.0:
logits, hidden = model(input_ids, return_hidden=True)
else:
logits = model(input_ids)
if (
not use_multiscreen_chunked_loss
and not (
train.mtp_loss_weight > 0.0
and model.config.multiscreen_chunked_lm_loss
)
):
loss = chunked_cross_entropy(logits, targets)
if train.mtp_loss_weight > 0.0:
mtp_targets = train_data[:, 2 : (model.max_seq_length + 1)].contiguous().long()
mtp_inputs = train_data[:, 1 : model.max_seq_length].contiguous().long()
if model.config.multiscreen_chunked_lm_loss:
mtp_loss = model.mtp_forward(
hidden[:, : mtp_targets.shape[1], :],
mtp_inputs,
targets=mtp_targets,
)
else:
mtp_logits = model.mtp_forward(
hidden[:, : mtp_targets.shape[1], :],
mtp_inputs,
)
mtp_loss = chunked_cross_entropy(mtp_logits, mtp_targets)
loss = loss + train.mtp_loss_weight * mtp_loss
if train.nta_margin_loss_weight > 0.0:
loss = loss + train.nta_margin_loss_weight * next_token_margin_loss(
logits,
targets,
train.nta_margin,
error_only=train.nta_margin_error_only,
)
fabric.backward(loss / train.gradient_accumulation_iters(devices, num_nodes))
running_loss.update(loss.detach())
if not is_accumulating:
fabric.clip_gradients(model, optimizer, max_norm=train.max_norm)
optimizer.step()
optimizer.zero_grad()
state["step_count"] += 1
if state["iter_num"] % log_iter_interval == 0:
loss = running_loss.compute().item() # expensive device-to-host synchronization
t1 = time.perf_counter()
throughput.update(
time=(t1 - total_t0),
flops=(measured_flops * log_iter_interval),
batches=state["iter_num"],
samples=(state["iter_num"] * train.micro_batch_size),
lengths=(state["iter_num"] * train.micro_batch_size * model.max_seq_length),
)
metrics = {
"loss": loss,
"iter": state["iter_num"],
"step": state["step_count"],
"epoch": train_iterator.epoch,
"iter_time": t1 - iter_t0,
"remaining_time": (
(t1 - total_t0) / (state["iter_num"] - initial_iter) * (max_iters - state["iter_num"])
),
"tokens": state["iter_num"] * train.micro_batch_size * model.max_seq_length,
"total_tokens": (state["iter_num"] * train.micro_batch_size * model.max_seq_length * fabric.world_size),
"learning_rate": lr,
}
if isinstance(val_loss, float):
val_loss = f"{val_loss:.3f}"
fabric.print(
f"Epoch {metrics['epoch'] + 1} | iter {metrics['iter']} step {metrics['step']} |"
f" loss train: {metrics['loss']:.3f},"
f" val: {val_loss} |"
f" iter time: {metrics['iter_time'] * 1000:.2f} ms"
f"{' (step)' if not is_accumulating else ''}"
f" remaining time: {timedelta(seconds=int(metrics['remaining_time']))!s}"
)
throughput_metrics = throughput.compute()
metrics.update(throughput_metrics)
fabric.log_dict(metrics, step=state["iter_num"] - 1)
if val_dataloader is not None and not is_accumulating and state["step_count"] % eval.interval == 0:
t0 = time.perf_counter()
val_loss = validate(fabric, model, val_dataloader, max_iters=eval.max_iters)
val_loss = val_loss.item()
td = time.perf_counter() - t0
fabric.print(f"iter {state['iter_num']}: val loss {val_loss:.4f}, val time: {td * 1000:.2f} ms")
metrics = {"val_loss": val_loss, "val_ppl": math.exp(val_loss)}
fabric.log_dict(metrics, step=state["iter_num"] - 1)
fabric.barrier()
if train.save_interval is not None and not is_accumulating and state["step_count"] % train.save_interval == 0:
save_checkpoint(fabric, state, tokenizer_dir, out_dir / f"step-{state['step_count']:08d}" / "lit_model.pth")
# Final validation
if eval.final_validation:
val_loss = validate(fabric, model, val_dataloader, max_iters=eval.max_iters)
metrics = {"val_loss": val_loss, "val_ppl": math.exp(val_loss)}
fabric.log_dict(metrics, step=state["iter_num"])
fabric.print(f"Final evaluation | val loss: {val_loss.item():.3f} | val ppl: {math.exp(val_loss):.3f}")
@torch.no_grad()
def validate(
fabric: L.Fabric, model: nn.Module, val_dataloader: DataLoader, max_iters: int, verbose: bool = True
) -> torch.Tensor:
fabric.barrier()
if verbose:
fabric.print("Validating ...")
model.eval()
losses = []
for k, batch in enumerate(val_dataloader):
if k >= max_iters:
break
input_ids = batch[:, 0 : model.max_seq_length].contiguous().long()
targets = batch[:, 1 : (model.max_seq_length + 1)].contiguous().long()
if model.config.mamba3_hierarchical_vocab:
# The two 64-way cross-entropies are exactly the NLL of the
# Cartesian-product 4096-way distribution. Avoid materializing a
# B x T x 4096 validation tensor at throughput microbatches.
loss = model(input_ids, targets=targets)
else:
logits = model(input_ids)
loss = chunked_cross_entropy(logits, targets)
losses.append(loss)
val_loss = torch.stack(losses).mean()
model.train()
fabric.barrier()
return val_loss
def get_dataloaders(
fabric: L.Fabric, data: DataModule, tokenizer: Tokenizer, train: TrainArgs, block_size: int
) -> tuple[DataLoader, DataLoader]:
data.connect(tokenizer=tokenizer, batch_size=train.micro_batch_size, max_seq_length=block_size)
with fabric.rank_zero_first():
data.prepare_data()
data.setup()
train_dataloader = data.train_dataloader()
val_dataloader = data.val_dataloader()
return train_dataloader, val_dataloader
# learning rate decay scheduler
def next_token_margin_loss(
logits: torch.Tensor,
targets: torch.Tensor,
margin: float = 0.0,
*,
error_only: bool = False,
) -> torch.Tensor:
"""Hinge loss that rewards the target logit outranking the strongest wrong logit.
``error_only`` concentrates the auxiliary gradient on positions currently
counted as NTA errors, avoiding extra margin pressure on already-correct
predictions while cross-entropy remains active everywhere.
"""
top_values, top_indices = logits.float().topk(2, dim=-1)
target_logits = logits.float().gather(-1, targets.unsqueeze(-1)).squeeze(-1)
top_is_target = top_indices[..., 0].eq(targets)
strongest_wrong = torch.where(top_is_target, top_values[..., 1], top_values[..., 0])
penalties = torch.relu(float(margin) + strongest_wrong - target_logits)
if not error_only:
return penalties.mean()
error_mask = (~top_is_target).to(dtype=penalties.dtype)
return (penalties * error_mask).sum() / error_mask.sum().clamp_min(1.0)
def get_lr(
learning_rate: float,
it: int,
warmup_iters: int,
max_iters: int,
min_lr: float,
schedule: str = "cosine",
decay_start_fraction: float = 0.7,
) -> float:
if schedule == "wsd":
if warmup_iters > 0 and it < warmup_iters:
return learning_rate * it / warmup_iters
decay_start = max(warmup_iters, int(max_iters * decay_start_fraction))
if it < decay_start:
return learning_rate
if it > max_iters:
return min_lr
decay_span = max(1, max_iters - decay_start)
decay_ratio = (it - decay_start) / decay_span
decay_ratio = min(max(decay_ratio, 0.0), 1.0)
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
return min_lr + coeff * (learning_rate - min_lr)
if schedule == "onecycle":
if warmup_iters > 0 and it < warmup_iters:
return min_lr + (learning_rate - min_lr) * it / warmup_iters
if it > max_iters:
return min_lr
decay_span = max(1, max_iters - warmup_iters)
decay_ratio = (it - warmup_iters) / decay_span
decay_ratio = min(max(decay_ratio, 0.0), 1.0)
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
return min_lr + coeff * (learning_rate - min_lr)
# 1) linear warmup for warmup_iters steps
if it < warmup_iters:
return learning_rate * it / warmup_iters
# 2) if it > max_iters, return min learning rate
if it > max_iters:
return min_lr
# 3) in between, use cosine decay down to min learning rate
decay_ratio = (it - warmup_iters) / (max_iters - warmup_iters)
assert 0 <= decay_ratio <= 1
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio)) # coeff ranges 0..1
return min_lr + coeff * (learning_rate - min_lr)
def initialize_weights(fabric: L.Fabric, model: GPT, n_layer: int, n_embd: int) -> None:
"""GPT-NeoX weight initialization (https://arxiv.org/abs/2204.06745)."""
# Adapted from https://github.com/jzhang38/TinyLlama
def init_weights(module, std):
nn.init.normal_(module.weight, mean=0.0, std=std)
if getattr(module, "bias", None) is not None:
nn.init.zeros_(module.bias)
def init_block_sparse_weights(module, std, down_std):
nn.init.normal_(module.up_weight, mean=0.0, std=std)
nn.init.normal_(module.value_weight, mean=0.0, std=std)
nn.init.normal_(module.down_weight, mean=0.0, std=down_std)
if hasattr(module, "router"):
init_weights(module.router, std)
if hasattr(module, "threshold"):
nn.init.zeros_(module.threshold.weight)
init = min(max(module.config.drelu_threshold_init, 1e-4), 1 - 1e-4)
nn.init.constant_(module.threshold.bias, math.log(init / (1 - init)))
if hasattr(module, "direction_input"):
# LoRA-style zero-output start: the base Euler+Sandwich function is
# exact at step zero, while random input/router factors break route
# symmetry and allow the conditional residual to grow by gradient.
nn.init.normal_(module.direction_input, mean=0.0, std=0.02)
nn.init.zeros_(module.direction_up)
nn.init.zeros_(module.direction_value)
init_weights(module.direction_router, 0.02)
module.direction_strength_raw.data.fill_(-2.0)
if model.config.initializer_range is not None:
std = model.config.initializer_range
for mod in model.modules():
if isinstance(mod, (nn.Embedding, nn.Linear)):
mod.reset_parameters = partial(init_weights, mod, std=std)
elif isinstance(
mod,
(
BlockSparseAdaptiveDSwiGLUMLP,
TileRoutedChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedChannelMemoryDSwiGLUMLP,
TileRoutedChannelStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelStateMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLP,
TileRoutedGroupStateMemoryDSwiGLUMLP,
TileRoutedKeyedChannelMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLPStaticA2,
TileRoutedDSwiGLUMLPStaticGPTS,
),
):
mod.reset_parameters = partial(init_block_sparse_weights, mod, std=std, down_std=std)
if not isinstance(fabric.strategy, FSDPStrategy):
reset_parameters(model)
return
for mod in model.modules():
if isinstance(mod, (nn.Embedding, nn.Linear)):
mod.reset_parameters = partial(init_weights, mod, std=math.sqrt(2.0 / 5 / n_embd))
if model.config.mamba3_enabled:
def init_mamba_input(module):
nn.Linear.reset_parameters(module)
def init_mamba_output(module):
nn.Linear.reset_parameters(module)
with torch.no_grad():
module.weight.div_(math.sqrt(n_layer))
# Preserve the official Mamba initialization instead of replacing its
# selective-state projections with GPT-NeoX normal initialization.
for block in model.transformer.h:
block.mixer.in_proj.reset_parameters = partial(init_mamba_input, block.mixer.in_proj)
block.mixer.out_proj.reset_parameters = partial(init_mamba_output, block.mixer.out_proj)
# need a separate loop because the residual projection linears below are `nn.Linear` too
for mod in model.modules():
if isinstance(mod, BlockSparseAdaptiveDSwiGLUMLP):
mod.reset_parameters = partial(
init_block_sparse_weights,
mod,
std=math.sqrt(2.0 / 5 / n_embd),
down_std=(1 / math.sqrt(n_embd) / n_layer),
)
if isinstance(
mod,
(
TileRoutedChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedChannelMemoryDSwiGLUMLP,
TileRoutedChannelStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelGroupStateMemoryDSwiGLUMLP,
TileRoutedCenteredChannelStateMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLP,
TileRoutedGroupStateMemoryDSwiGLUMLP,
TileRoutedKeyedChannelMemoryDSwiGLUMLP,
TileRoutedDSwiGLUMLPStaticA2,
TileRoutedDSwiGLUMLPStaticGPTS,
),
):
mod.reset_parameters = partial(
init_block_sparse_weights,
mod,
std=math.sqrt(2.0 / 5 / n_embd),
down_std=(1 / math.sqrt(n_embd) / n_layer),
)
if isinstance(
mod,
(LLaMAMLP, DSwiGLUMLP, AdaptiveDSwiGLUMLP, CausalSelfAttention),
):
mod.proj.reset_parameters = partial(init_weights, mod.proj, std=(1 / math.sqrt(n_embd) / n_layer))
if not isinstance(fabric.strategy, FSDPStrategy):
reset_parameters(model)
def save_checkpoint(fabric, state, tokenizer_dir, checkpoint_file):
model = state["model"]
checkpoint_file.parent.mkdir(parents=True, exist_ok=True)
fabric.print(f"Saving checkpoint to {str(checkpoint_file)!r}")
fabric.save(checkpoint_file, state)
if fabric.global_rank == 0:
save_hyperparameters(setup, checkpoint_file.parent)
if tokenizer_dir is not None:
copy_config_files(tokenizer_dir, checkpoint_file.parent)
save_config(model.config, checkpoint_file.parent)
def validate_args(train: TrainArgs, eval: EvalArgs, initial_checkpoint_dir, resume) -> None:
issues = []
unsupported = [(train, ["epochs"]), (eval, ["max_new_tokens"])]
for args, names in unsupported:
for name in names:
if getattr(args, name) is not None:
issues.append(f"{__file__} doesn't support the {name!r} argument. This is set in {args}")
if train.max_steps is not None:
warnings.warn(
"`train.max_steps` is intended for profiling or debug runs only. "
"For full pretraining runs, prefer `train.max_tokens` or `train.max_time`.",
UserWarning,
)
required = [(train, ["max_tokens", "max_norm"])]
for args, names in required:
for name in names:
if getattr(args, name) is None:
issues.append(f"{__file__} requires the {name!r} argument. This is set in {args}")
if initial_checkpoint_dir and resume:
issues.append("Can't provide both `--resume` and `--initial_checkpoint_dir`. Choose one.")
if issues:
raise ValueError("\n".join(issues))