Jump to content

Connect SuperML | Leeroopedia MCP: Equip your AI agents with best practices, code verification, and debugging knowledge. Powered by Leeroo — building Organizational Superintelligence. Contact us at founders@leeroo.com.

Implementation:Hiyouga LLaMA Factory V1 Base Trainer

From Leeroopedia
Revision as of 10:41, 27 September 2026 by Agent (talk | contribs) (Sync from local file)
(diff) ← Older revision | Latest revision (diff) | Newer revision → (diff)


Knowledge Sources
Domains Machine Learning, Distributed Training
Last Updated 2026-02-06 19:00 GMT

Overview

BaseTrainer is the abstract base trainer class that orchestrates the complete distributed training lifecycle including batch generation, model sharding, optimizer initialization, and the training loop with gradient accumulation.

Description

The BaseTrainer class encapsulates the full training pipeline for LLaMA-Factory v1. During initialization, it creates a BatchGenerator for data feeding, conditionally orders optimizer initialization before or after model sharding depending on the distributed strategy (DeepSpeed requires optimizer first, FSDP requires sharding first), and sets up the learning rate scheduler. The fit() method runs the training loop with gradient accumulation across micro-batches, loss scaling by valid tokens for proper distributed averaging, gradient clipping, and all-reduce operations for synchronized logging. Subclasses must implement the abstract compute_loss method.

Usage

Do not instantiate BaseTrainer directly. Instead, subclass it (as SFTTrainer, DPOTrainer, or RMTrainer do) and implement compute_loss to define the specific training objective. The trainer handles all distributed training concerns, checkpoint saving, and training loop management automatically.

Code Reference

Source Location

Signature

class BaseTrainer:
    def __init__(
        self,
        args: TrainingArguments,
        model: HFModel,
        renderer: Renderer,
        train_dataset: TorchDataset,
    ) -> None: ...

    def _create_batch_generator(self) -> None: ...
    def _shard_model(self) -> None: ...
    def _init_optimizer(self) -> None: ...
    def _init_lr_scheduler(self) -> None: ...

    def compute_log_probs(self, model: HFModel, batch: BatchInput) -> Tensor: ...

    @abstractmethod
    def compute_loss(self, batch: BatchInput) -> Tensor: ...

    def fit(self) -> None: ...
    def save_model(self) -> None: ...

Import

from llamafactory.v1.core.base_trainer import BaseTrainer

I/O Contract

Inputs

Name Type Required Description
args TrainingArguments Yes Training configuration including batch sizes, learning rate, epochs, gradient clipping, output directory, and distributed config.
model HFModel Yes The HuggingFace model to train.
renderer Renderer Yes The renderer for tokenization and template application.
train_dataset TorchDataset Yes The training dataset (typically a DataEngine instance).
batch (compute_loss) BatchInput Yes A micro-batch dictionary with input_ids, labels, attention_mask, and loss_weights.

Outputs

Name Type Description
compute_loss return Tensor Scalar loss tensor for the given micro-batch.
compute_log_probs return Tensor Log probability tensor of shape (batch_size, seq_len - 1).
save_model None Saves model and processor to args.output_dir.

Usage Examples

# Subclassing BaseTrainer for SFT
from llamafactory.v1.core.base_trainer import BaseTrainer

class SFTTrainer(BaseTrainer):
    def compute_loss(self, batch):
        # Compute cross-entropy loss with valid token masking
        log_probs = self.compute_log_probs(self.model, batch)
        labels = batch["labels"].to(self.device)
        loss_weights = batch["loss_weights"].to(self.device)
        # ... loss computation logic
        return loss

# Running training
trainer = SFTTrainer(args=training_args, model=model, renderer=renderer, train_dataset=dataset)
trainer.fit()
trainer.save_model()

Related Pages

Page Connections

Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment