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 Data Collator

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


Knowledge Sources
Domains Data Processing, Training
Last Updated 2026-02-06 19:00 GMT

Overview

Concrete data collators for batching multimodal and training-paradigm-specific inputs provided by LLaMA Factory.

Description

This module provides four data collator classes and one utility function that bridge tokenized dataset examples and model-ready batched tensors. The base class MultiModalDataCollatorForSeq2Seq extends HuggingFace's DataCollatorForSeq2Seq to handle images, videos, and audios, including generating fake multimodal inputs to prevent distributed training hangs when a batch contains no media. Three specialized subclasses handle different training paradigms:

  • SFTDataCollatorWith4DAttentionMask -- Adds 4D attention mask support for packed sequence training with block-diagonal attention
  • PairwiseDataCollatorWithPadding -- Reformats chosen/rejected pairs for DPO and reward modeling
  • KTODataCollatorWithPadding -- Handles KTO training with separate target and KL-reference inputs

The standalone prepare_4d_attention_mask function expands a 2D attention mask with packing indices into a 4D lower-triangular mask that prevents cross-sequence attention in packed batches.

Usage

These collators are instantiated by the training workflow based on the training stage and passed to the HuggingFace Trainer as the data_collator argument. They are called automatically during each training step to collate individual dataset examples into batched tensors.

Code Reference

Source Location

Signature

def prepare_4d_attention_mask(
    attention_mask_with_indices: "torch.Tensor",
    dtype: "torch.dtype",
) -> "torch.Tensor": ...

@dataclass
class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
    template: Optional["Template"] = None
    processor: Optional["ProcessorMixin"] = None
    def __call__(self, features: list[dict[str, Any]]) -> dict[str, "torch.Tensor"]: ...

@dataclass
class SFTDataCollatorWith4DAttentionMask(MultiModalDataCollatorForSeq2Seq):
    block_diag_attn: bool = False
    attn_implementation: Literal["eager", "sdpa", "flash_attention_2"] = "eager"
    compute_dtype: "torch.dtype" = torch.float32
    def __call__(self, features: list[dict[str, Any]]) -> dict[str, "torch.Tensor"]: ...

@dataclass
class PairwiseDataCollatorWithPadding(MultiModalDataCollatorForSeq2Seq):
    def __call__(self, features: list[dict[str, Any]]) -> dict[str, "torch.Tensor"]: ...

@dataclass
class KTODataCollatorWithPadding(MultiModalDataCollatorForSeq2Seq):
    def __call__(self, features: list[dict[str, Any]]) -> dict[str, "torch.Tensor"]: ...

Import

from llamafactory.data.collator import (
    prepare_4d_attention_mask,
    MultiModalDataCollatorForSeq2Seq,
    SFTDataCollatorWith4DAttentionMask,
    PairwiseDataCollatorWithPadding,
    KTODataCollatorWithPadding,
)

I/O Contract

Inputs

Name Type Required Description
features list[dict[str, Any]] Yes List of tokenized examples, each containing input_ids, attention_mask, labels, and optionally images, videos, audios
template Template Yes Chat template with multimodal plugin for processing media inputs
processor ProcessorMixin No HuggingFace processor for multimodal models
model PreTrainedModel No Model reference used for mRoPE position ID computation

Outputs

Name Type Description
batch dict[str, torch.Tensor] Batched tensors including input_ids, attention_mask, labels, and multimodal inputs (pixel_values, image_grid_thw, etc.)

Usage Examples

from llamafactory.data.collator import SFTDataCollatorWith4DAttentionMask

collator = SFTDataCollatorWith4DAttentionMask(
    tokenizer=tokenizer,
    model=model,
    template=template,
    processor=processor,
    block_diag_attn=True,
    attn_implementation="sdpa",
    compute_dtype=torch.bfloat16,
)

# Called by the Trainer during training
batch = collator(features)
# batch contains: input_ids, attention_mask (4D), labels, pixel_values, ...

Related Pages

Page Connections

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