espnet2.speechlm.trainer.titan_trainer_pp.TitanPPTrainer
espnet2.speechlm.trainer.titan_trainer_pp.TitanPPTrainer
class espnet2.speechlm.trainer.titan_trainer_pp.TitanPPTrainer(*args, **kwargs)
Bases: TitanTrainer
Pipeline-parallel trainer.
Inherits from TitanTrainer and overrides train() and valid() to use the PP schedule for forward/backward instead of explicit gradient accumulation.
self.model is always an nn.ModuleList of model chunks. For single-stage schedules (1F1B) it contains one element; for multi-stage schedules (Interleaved1F1B) it contains vpp_degree elements. The last element is always the last virtual stage (which computes the loss).
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.
train() → None
PP training: schedule.step() handles microbatching and backward.
valid() → None
Run validation on all validation datasets.
