Implementation:Hiyouga LLaMA Factory V1 Types
| Knowledge Sources | |
|---|---|
| Domains | Machine Learning, Type Systems |
| Last Updated | 2026-02-06 19:00 GMT |
Overview
This module defines the central type system for the LLaMA-Factory v1 architecture, providing TypedDict structures for data, messages, model I/O, and TYPE_CHECKING-guarded type aliases for PyTorch and HuggingFace objects.
Description
The module defines eleven TypedDict and NamedTuple classes that structure data flow throughout the v1 pipeline. DatasetInfo describes dataset metadata (path, source, split, converter, size, weight, streaming). DistributedConfig parameterizes distributed training (mp_replicate_size, mp_shard_size, dp_size, cp_size, timeout). Content and Message define the chat message format with role, typed content items (text, reasoning, tool_call, image_url), and optional loss weights. SFTSample and DPOSample structure training samples for supervised fine-tuning and direct preference optimization. ToolCall describes function call structures. ModelInput and BatchInput define the tokenized input format (input_ids, attention_mask, labels, loss_weights, position_ids, token_type_ids) at per-sample and batched tensor levels. BatchInfo carries micro-batch metadata, and ModelOutput wraps model logits as a NamedTuple. Under TYPE_CHECKING, type aliases (Tensor, HFModel, Processor, Optimizer, etc.) resolve to actual PyTorch/HuggingFace types for IDE support, but are set to None at runtime to avoid import overhead.
Usage
Import these types for type annotations throughout the v1 codebase. Use the TypedDict classes to construct and validate structured data in training pipelines, rendering plugins, and inference engines. The TYPE_CHECKING-guarded aliases should be used in type hints only, not at runtime.
Code Reference
Source Location
- Repository: Hiyouga_LLaMA_Factory
- File: src/llamafactory/v1/utils/types.py
- Lines: 1-180
Signature
# TYPE_CHECKING-guarded aliases
Tensor = torch.Tensor
TensorLike = Union[int, float, list[int], list[float], np.ndarray, Tensor]
TorchDataset = Union[torch.utils.data.Dataset, torch.utils.data.IterableDataset]
HFDataset = Union[datasets.Dataset, datasets.IterableDataset]
DataCollator = transformers.DataCollator
DataLoader = torch.utils.data.DataLoader
HFConfig = transformers.PretrainedConfig
HFModel = transformers.PreTrainedModel
DistModel = Union[torch.nn.parallel.DistributedDataParallel, FullyShardedDataParallel]
Processor = Union[transformers.PreTrainedTokenizer, transformers.ProcessorMixin]
Optimizer = torch.optim.Optimizer
Scheduler = torch.optim.lr_scheduler.LRScheduler
ProcessGroup = ProcessGroup
# TypedDict classes
class DatasetInfo(TypedDict, total=False): ...
class DistributedConfig(TypedDict, total=False): ...
class Content(TypedDict): ...
class Message(TypedDict): ...
class SFTSample(TypedDict): ...
class DPOSample(TypedDict): ...
Sample = Union[SFTSample, DPOSample]
class ToolCall(TypedDict): ...
class ModelInput(TypedDict, total=False): ...
class BatchInput(TypedDict, total=False): ...
class BatchInfo(TypedDict): ...
class ModelOutput(NamedTuple): ...
Import
from llamafactory.v1.utils.types import (
Message, Content, ModelInput, BatchInput, BatchInfo, ModelOutput,
SFTSample, DPOSample, Sample, ToolCall, DatasetInfo, DistributedConfig,
HFModel, Processor, Tensor, TorchDataset, HFDataset,
)
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| Content.type | Literal["text", "reasoning", "tool_call", "image_url"] | Yes | The type of the content item |
| Content.value | str | Yes | The value/payload of the content item |
| Message.role | Literal["system", "user", "assistant", "tool"] | Yes | The role of the message sender |
| Message.content | list[Content] | Yes | List of content items in the message |
| Message.loss_weight | float | No | Loss weight for training (default 1.0) |
| ModelInput.input_ids | list[int] | Yes | Tokenized input IDs |
| ModelInput.attention_mask | list[int] | Yes | Attention mask (1 for real tokens, 0 for padding) |
| ModelInput.labels | list[int] | Yes | Target labels (IGNORE_INDEX for non-loss tokens) |
| ModelInput.loss_weights | list[float] | Yes | Per-token loss weights |
Outputs
| Name | Type | Description |
|---|---|---|
| ModelOutput.logits | Tensor | Model output logits tensor |
| Sample | Union[SFTSample, DPOSample] | A training sample for either SFT or DPO |
| BatchInfo.data_iter | Iterator[list[ModelInput]] | Iterator yielding micro-batches of model inputs |
Usage Examples
from llamafactory.v1.utils.types import Message, Content, ModelInput, SFTSample
# Construct a chat message
message = Message(
role="user",
content=[Content(type="text", value="What is the capital of France?")],
loss_weight=0.0,
)
# Construct an SFT training sample
sample = SFTSample(
messages=[
Message(role="user", content=[Content(type="text", value="Hello")], loss_weight=0.0),
Message(role="assistant", content=[Content(type="text", value="Hi!")], loss_weight=1.0),
]
)
# Construct a model input
model_input = ModelInput(
input_ids=[1, 2, 3, 4],
attention_mask=[1, 1, 1, 1],
labels=[-100, -100, 3, 4],
loss_weights=[0.0, 0.0, 1.0, 1.0],
)
Related Pages
- Implementation:Hiyouga_LLaMA_Factory_V1_Rendering_Plugin - Uses Message, ModelInput, Content, and ToolCall types
- Implementation:Hiyouga_LLaMA_Factory_V1_CLI_Sampler - Uses Message, Sample, HFModel, and TorchDataset types
- Implementation:Hiyouga_LLaMA_Factory_V1_Kernel_Base - Uses HFModel type alias