espnet2.speechlm.trainer.titan_trainer.TitanTrainer
espnet2.speechlm.trainer.titan_trainer.TitanTrainer
class espnet2.speechlm.trainer.titan_trainer.TitanTrainer(train_data_factory, valid_data_factories: Dict, model: Module, resume_path: Path | None, output_dir: Path, trainer_args: Dict[str, Any], parallel_dims=None)
Bases: object
TorchTitan-based trainer with FSDP2 support for SpeechLM.
This trainer provides distributed training using PyTorch’s native FSDP2 (Fully Sharded Data Parallel) instead of DeepSpeed ZeRO. It maintains interface compatibility with DeepSpeedTrainer for easy switching.
IMPORTANT: wandb is MANDATORY and must be initialized before creating this trainer. The trainer will raise an error if wandb.run is None. Wandb should always be initialized in offline mode for local-only logging. All training metrics, losses, and stats are logged to local wandb files.
Key Features: : - FSDP2/HSDP for memory-efficient data parallelism
- Activation checkpointing for memory optimization
- PyTorch Distributed Checkpoint (DCP) for reshardable checkpoints
- Compatible with existing DataIteratorFactory interface
Initialize TorchTitan trainer.
- Parameters:
train_data_factory – Training data iterator factory
valid_data_factories – Dictionary of validation data factories
model – Model to train (HuggingFace model)
resume_path – Path to checkpoint for resuming training
output_dir – Directory for saving outputs
trainer_args –
Training configuration dictionary containing:
- max_step: Maximum number of training steps
- log_interval: Steps between logging
- save_interval: Steps between checkpoints
- freeze_param: List of parameter prefixes to freeze
- titan_config: TorchTitan configuration dict with:
- dp_shard: FSDP degree (-1 = auto)
- dp_replicate: HSDP replicate degree (default: 1)
- mixed_precision_param: Parameter dtype (default: “bfloat16”)
- mixed_precision_reduce: Reduce dtype (default: “float32”)
- gradient_clipping: Max gradient norm (default: 1.0)
- optimizer: Optimizer config dict
- lr_scheduler: LR scheduler config dict
- activation_checkpoint: “none”, “selective”, or “full”
parallel_dims – Pre-built ParallelDims from train.py (avoids double init_parallel_dims). If None, builds internally.
count_normalized_keys = ['loss', 'ce_loss', 'z_loss', 'z_loss_s0', 'z_loss_mm', 'load_balance_loss', 'acc_layer0']
run() → None
Main training loop.
train() → None
Execute one training epoch (save_interval optimizer steps).
With gradient accumulation, each optimizer step consumes gradient_accumulation_steps micro-batches. The iterator is sized so that save_interval optimizer steps are performed.
valid() → None
Run validation on all validation datasets.
