Implementation:Hiyouga LLaMA Factory Data Collator
| 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
- Repository: Hiyouga_LLaMA_Factory
- File: src/llamafactory/data/collator.py
- Lines: 1-331
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
- Hiyouga_LLaMA_Factory_Chat_Template - Template class used for multimodal plugin access
- Hiyouga_LLaMA_Factory_Multimodal_Plugin - Plugin system invoked to generate mm_inputs
- Hiyouga_LLaMA_Factory_Data_Loader - Dataset loading pipeline that produces the features consumed by collators