首页
/ LLaMA-Factory v1 DataEngine 深度解析:数据集加载、索引构建与样本格式标准化

LLaMA-Factory v1 DataEngine 深度解析:数据集加载、索引构建与样本格式标准化

2026-09-04 09:35:08作者:董灵辛Dennis

本文基于 LLaMA-Factory 仓库中的 DataEngine 设计文档,完整讲解 v1 数据引擎的接口定义、初始化三步流水线(元信息解析 → 数据集加载 → 数据索引构建)、converter 插件化格式标准化机制,以及实例化后的索引/切片访问方式;并结合 data_engine.pydata_plugins 等源码补充实现细节与验证方式,帮助你在多数据集混训、多轮对话前缀展开和自定义数据接入场景中正确使用和扩展该引擎。

1. DataEngine 是什么

DataEngine 是 LLaMA-Factory v1 数据处理的核心类,继承自 PyTorch 的 Dataset,负责各种数据插件的接入——数据格式转换、数据加载等功能均以插件形式实现并挂接到 DataEngine 上。其设计目标在 源码模块 docstring 中也有明确表述:

  • 数据引擎等价于一个 torch Datasetdata_engine[i] 即可取样本);
  • 数据引擎与模型无关(agnostic to the model),即它只负责产出标准化样本,不感知下游训练任务与模型结构。

实例化时传入数据配置,DataEngine 内部维护四个核心状态(见 类属性定义):

属性 类型 含义
datasets dict[str, HFDataset] 数据集名称 → 已加载数据对象(datasets.Dataset / IterableDataset)的映射
dataset_infos dict[str, DatasetInfo] 数据集名称 → 元信息(路径、来源、converter、size、weight 等)的映射
data_index list[tuple] 全局数据索引,将多个数据集拉平成一个统一索引序列
streaming bool 是否处于流式模式

这种“元信息(dataset_infos)+ 数据体(datasets)+ 统一索引(data_index)”三分离的结构,是多数据集混训时能独立控制采样频率、按权重调整规模的基础。

2. 接口定义:DataArguments 与 DataEngine

设计文档给出的接口定义如下。DataEngine 接受一个唯一入参:DataArguments 实例,所有元数据集信息均通过该参数传入。

@dataclass
class DataArguments:
    """ `DataEngine`初始化入参

    args:
        dataset (str): 数据集路径,远程数据集 repo id / dataset_info.yaml 路径,或本地数据集路径/dataset_info.yaml路径
        cutoff_len (int): 数据集截止长度,即数据集最大样本采样数量
    """
    ...

两个核心参数的语义:

  • dataset:数据集路径,支持本地或远程。当传入本地数据集文件路径时,要求该数据集已是标准格式;否则需要传入 dataset_info.yaml 来配置该数据集的 converter 等元信息,以告知 DataEngine 应当如何处理该数据。
  • cutoff_len:数据集的截止长度,即该数据集的最大样本数量。

需要说明的是,当前仓库的 DataArguments 实现 中,字段为 train_dataseteval_dataset(均为 str | None):

@dataclass
class DataArguments:
    train_dataset: str | None = field(default=None, metadata={"help": "Path to the training dataset."})
    eval_dataset: str | None = field(default=None, metadata={"help": "Path to the evaluation dataset."})

也就是说,文档以“接口设计”视角描述参数(dataset/cutoff_len),而仓库当前实现中,DataEngine 构造入口实际接收的是数据集路径字符串(dataset_path),测试用例模块 __main__ 入口均以 DataEngine(data_args.train_dataset) 的形式调用。阅读下文时应同时理解两个视角。

DataEngine 类的完整职责(继承自设计文档的接口定义):

class DataEngine(Dataset):
    """数据引擎(DataEngine)

    `DataEngine` 负责数据集的加载与统一管理,支持:
        - 从本地路径或 Hugging Face Hub 加载数据
        - 通过插件机制加载自定义数据
        - 构建统一的数据索引
        - 支持流式(streaming)与非流式数据访问
    """

    def __init__(self, data_args: DataArguments) -> None: ...
    def get_dataset_info(self) -> None: ...
    def load_dataset(self) -> None: ...
    def build_data_index(self) -> None: ...
    def _convert_data_sample(self, raw_sample: dict[str, Any], dataset_name: str) -> Sample: ...
    def __len__(self) -> int: ...
    def __getitem__(self, index: Union[int, Any]) -> Union[Sample, list[Sample]]: ...
    def __iter__(self) -> Iterable: ...
    async def __aiter__(self) -> AsyncIterable: ...

各方法的接口契约(摘自设计文档):

方法 契约要点
__init__(data_args) 初始化时自动依次执行:1) get_dataset_info 读取并解析数据集元信息;2) load_dataset 按配置加载数据集;3) build_data_index 构建统一索引列表
get_dataset_info() 根据 self.args.dataset 确定数据源,支持四种:本地 YAML 配置路径、HF Hub 上的 YAML 配置路径、本地数据集文件路径、HF Hub 数据集 repo id
load_dataset() 每个数据集条目支持:hf_hub_urldatasets.load_dataset 加载;本地数据文件经 DataLoaderPlugin 插件加载;streaming 是否启用流式。执行后更新 self.datasets,且若任一数据集为流式则置 self.streaming = True
build_data_index() 为所有数据集创建全局索引 (dataset_name, sample_index);流式模式下生成固定长度(例如 1000)的占位索引;DataIndexPlugin 可根据 size/weight 调整索引分布
_convert_data_sample() 根据 dataset_infoconverter 字段调用对应转换插件,把原始样本标准化为统一结构;converter 为空则假定已是标准格式。由 __getitem__ 调用
__len__() 返回总样本数;流式数据集返回 -1
__getitem__(index) index 支持 int 或 list[int],返回单个样本或样本列表
__iter__() 用于非流式数据集的顺序或随机访问
__aiter__() 用于流式数据集或异步数据加载场景,在异步环境中按流读取样本

2.1 DatasetInfo:每个数据集条目的元信息字段

元信息在 types.py 中被定义为 DatasetInfo TypedDict,各字段及默认值如下(这是对文档中 converter/size/weight/streaming 字段的完整落点):

字段 类型 默认值 含义
path str 必填 数据集路径(本地路径或 HF Hub repo id)
source Literal["hf_hub", "ms_hub", "local"] "hf_hub" 数据集来源,决定走 HF Hub 直载还是本地 loader 插件
split str "train" 读取的数据集 split
converter str 数据转换插件名,如 alpaca/sharegpt/pair
size int 全部样本 该数据集期望的样本数量(截断/重采样目标)
weight float 1.0 数据集权重,按比例缩放索引规模
streaming bool False 是否流式加载

仓库自带的 data/v1_sft_demo.yaml 是一个最小可参考的 dataset_info 文件:

identity:
  path: data/identity.json
  source: local
  converter: alpaca
alpaca_en_demo:
  path: data/alpaca_en_demo.json
  source: local
  converter: alpaca
  size: 500

其中 identityconverter: alpaca 声明原始格式,alpaca_en_demo 额外用 size: 500 控制进入全局索引的样本规模。

3. 初始化流水线:三步自动执行

DataEngine 实例化时自动执行三步:解析元信息 → 加载数据 → 构建索引。构造函数 中对应:

self._get_dataset_info()
self._load_dataset()
self._build_data_index()

3.1 _get_dataset_info:加载数据元信息

根据 dataset 参数加载数据集配置,获取数据位置、数据格式、插件配置等所有数据元信息。实现 按四个分支判定数据源,与文档声明的四种选项一一对应:

if self.path.endswith(".yaml") and os.path.isfile(self.path):  # 1. 本地 YAML 配置文件
    self.dataset_infos = OmegaConf.load(self.path)
elif self.path.endswith(".yaml"):  # 2. HF Hub 上的 YAML,如 llamafactory/v1-sft-demo/dataset_info.yaml
    repo_id, filename = os.path.split(self.path)
    filepath = hf_hub_download(repo_id=repo_id, filename=filename, repo_type="dataset")
    self.dataset_infos = OmegaConf.load(filepath)
elif os.path.exists(self.path):  # 3. 本地数据集文件
    self.dataset_infos = {"default": {"path": self.path, "source": "local"}}
else:  # 4. HF Hub 数据集 repo id,如 llamafactory/v1-sft-demo
    self.dataset_infos = {"default": {"path": self.path}}

注意分支 3/4:直接给文件路径或 repo id 时,DataEngine 会自动合成一条名为 default 的元信息(source: local 或默认 hf_hub)。这印证了文档的说法——直接传本地文件路径时要求数据为标准格式(合成条目不带 converter 字段,样本会被原样透传,见第 4 节)。

3.2 load_dataset:按元信息加载所有数据集

文档给出的流程是:遍历所有数据源,按不同数据源加载数据:

for key, value in self.dataset_infos.items():
    split = value.get("split", "train")
    streaming = value.get("streaming", False)

    if "hf_hub_url" in value:
        # 从 HF Hub 加载
        dataset = load_dataset(value["hf_hub_url"], split=split, streaming=streaming)
    else:
        # 使用 DataLoaderPlugin 加载本地文件
        dataset = DataLoaderPlugin(args=self.args).auto_load_data(value)

    self.datasets[key] = dataset

当前源码实现 的核心判定改为读取 source 字段(默认 "hf_hub"),并比文档描述多了一条流式一致性校验:

is_streaming = [dataset_info.get("streaming", False) for dataset_info in self.dataset_infos.values()]
self.streaming = any(is_streaming)
if all(is_streaming) != any(is_streaming):
    raise ValueError("All datasets must be streaming or non-streaming.")

for dataset_name, dataset_info in self.dataset_infos.items():
    split = dataset_info.get("split", "train")
    if dataset_info.get("source", "hf_hub") == "hf_hub":
        from datasets import load_dataset
        self.datasets[dataset_name] = load_dataset(dataset_info["path"], split=split, streaming=self.streaming)
    else:  # data loader plugin
        from ..plugins.data_plugins.loader import DataLoaderPlugin
        self.datasets[dataset_name] = DataLoaderPlugin(dataset_info["source"]).load(dataset_info)

两条值得注意的实现事实:

  1. 流式模式要求“全流式或全非流式”,混用会直接抛 ValueError。文档中“任一数据集为流式则置 streaming=True”的行为只在全量流式时成立。
  2. 本地加载走 DataLoaderPlugin。注册为 "local" 的 loader 支持文件或目录两种输入,并按扩展名路由到 datasets.load_dataset 的对应 builder(arrow/csv/json/parquet/textjsonl 归入 jsontxt 归入 text):
@DataLoaderPlugin("local").register()
def load_data_from_file(filepath: str, split: str, streaming: bool) -> HFDataset:
    if os.path.isdir(filepath):
        filetype = _get_builder_name(os.listdir(filepath)[0])
        dataset = load_dataset(filetype, data_dir=filepath, split=split)
    elif os.path.isfile(filepath):
        filetype = _get_builder_name(filepath)
        dataset = load_dataset(filetype, data_files=filepath, split=split)
    else:
        raise ValueError(f"Can not load dataset from {filepath}.")

    if streaming:  # faster when data is streamed from local files
        dataset = dataset.to_iterable_dataset()
    return dataset

即 v1 数据引擎当前支持的本地文件格式为:arrowcsvjson/jsonlparquettxt。从源码结构看,DatasetInfo.source 的类型标注中还包含 ms_hub,但目前 loader 插件注册表中只注册了 local 实现,传入 ms_hub 会在插件解析阶段报“未注册”错误。

3.3 build_data_index:构建统一索引

为每个数据集创建索引列表 [(dataset_name, sample_index), ...]DataIndexPlugin(当前实现为 adjust_data_index 函数)在此处被调用,可控制各数据集的采样频率、采样方式。文档给出的流程代码:

for dataset_name, dataset in self.datasets.items():
    # 创建基础索引
    data_index = [(dataset_name, idx) for idx in range(len(dataset))]

    # 根据 size 和 weight 调整索引
    size = self.dataset_infos[dataset_name].get("size")
    weight = self.dataset_infos[dataset_name].get("weight")
    if size or weight:
        data_index = DataIndexPlugin().adjust_data_index(data_index, size, weight)

    self.data_index.extend(data_index)

当前源码 在此基础上有两点增强:

(1)流式占位索引。流式数据集无法预先计数,生成固定 1000 个占位索引:

if self.streaming:  # cannot pre-count turns -> keep whole, unsplit
    data_index = [(dataset_name, -1, None) for _ in range(1000)]

(2)多轮 SFT 对话的前缀展开(prefix expansion)。非流式路径下,索引条目是三元的 (dataset_name, sample_index, cut),其中 cutmessages[:cut] 的前缀长度:

for sample_index in range(len(dataset)):
    sample = self._convert_data_sample(dataset[sample_index], dataset_name)
    for cut in self._prefix_cuts(sample):
        data_index.append((dataset_name, sample_index, cut))

_prefix_cuts 对每个有监督的 assistant 轮loss_weight > 1e-6)切出一个前缀:

@staticmethod
def _prefix_cuts(sample: Sample) -> list[int | None]:
    """``u1 a1 u2 a2`` -> ``[2, 4]`` (samples ``messages[:2]`` and ``messages[:4]``,
    each trained on its last assistant turn)."""
    messages = sample.get("messages")
    if not messages:
        return [None]
    cuts = [i + 1 for i, m in enumerate(messages) if m["role"] == "assistant" and m.get("loss_weight", 1.0) > 1e-6]
    return cuts or [None]

以一段 u1 a1 u2 a2 的多轮对话为例,它会被拆成 messages[:2]messages[:4] 两条训练样本,每条只训练其最后一个 assistant 轮(前面各轮的 loss_weight 由渲染阶段负责)。这使 len(data_engine) 反映的是真实训练样本数而非原始行数;对 DPO、流式或无监督轮样本则整体保留(cut=None)。

(3)size / weight 调整语义adjust_data_index实际实现 基于 random.choices有放回采样):

if size is not None:
    data_index = random.choices(data_index, k=size)          # 截断/重采样到 size 条
if weight is not None:
    data_index = random.choices(data_index, int(len(data_index) * weight))  # 按权重比例缩放

也就是说 size 把该数据集在混合索引中的规模固定为目标值,weight 按原始规模的 weight 倍缩放,两者可叠加生效(先 size 后 weight)。

4. _convert_data_sample:数据格式标准化与 converter 插件

该方法把原始样本转换为统一格式,由 __getitem__ 触发。DataConverterPlugin 在此处被调用,具体调用的插件由 get_dataset_info 得到的 converter 元信息指定;converter 为空则假定数据集已是标准格式,直接透传。文档给出的代码:

def _convert_data_sample(self, raw_sample: dict, dataset_name: str) -> Sample:
    converter = self.dataset_infos[dataset_name].get("converter")
    if converter is not None:
        # 使用指定的转换器
        from ..plugins.data_plugins.converter import get_converter
        return {"_dataset_name": dataset_name, **get_converter(converter)(raw_sample)}
    else:
        # 已经是标准格式
        return {"_dataset_name": dataset_name, **raw_sample}

当前源码 通过 DataConverterPlugin(converter)(raw_sample) 按插件名路由,效果一致。

4.1 内置的三种 converter

converter.py 中注册了三个转换器,均把原始样本归一到 messages 结构(每条消息含 rolecontent 内容块列表、loss_weight):

插件名 原始格式 转换行为
alpaca instruction/input/output/system(可选 images/videos/audios 列) system→system 消息(loss_weight=0.0);instruction+input→user 消息(loss_weight=0.0);output→assistant 消息(loss_weight=1.0
sharegpt conversationsfrom: human/gpt/system/function_call/observation,可选 tools 角色映射:human→user、gpt→assistant(loss_weight=1.0)、observation→tool;function_call 解析为 JSON 后落成 assistant 消息中的 tool_call 内容块;tools 字段规范化为 JSON 字符串
pair chosen/rejected 两组 OpenAI 风格消息 产出 DPO 样本 chosen_messages/rejected_messages,每侧独立消费媒体占位符;assistant 轮 loss_weight=1.0,其余 0.0

alpaca 转换器 为例:

@DataConverterPlugin("alpaca").register()
def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample:
    messages = []
    media_iters = _build_media_iters(raw_sample)
    if "system" in raw_sample:
        messages.append({"role": "system", "content": [{"type": "text", "value": raw_sample["system"]}], "loss_weight": 0.0})
    if "instruction" in raw_sample or "input" in raw_sample:
        messages.append({
            "role": "user",
            "content": _to_content_blocks(raw_sample.get("instruction", "") + raw_sample.get("input", ""), media_iters),
            "loss_weight": 0.0,
        })
    if "output" in raw_sample:
        messages.append({"role": "assistant", "content": [{"type": "text", "value": raw_sample["output"]}], "loss_weight": 1.0})
    _assert_media_consumed(media_iters)
    return {"messages": messages}

多模态占位符处理是转换器的重要一环:_to_content_blocks实现)按内联媒体标签把文本切分成 text/image_url/video_url/audio_url 内容块,每个占位标签按文档顺序消费该模态列(images/videos/audios)中的下一个路径;标签数多于媒体文件、或少于媒体文件(存在未引用的媒体)都会显式抛错,避免静默错位。

4.2 标准化样本结构

转换结果被归一为 types.py 中的 TypedDict:content 块类型支持 text/reasoning/tool_call/image_url/video_url/audio_urlSFTSamplemessages 与可选 tools(JSON 字符串)、extra_info_dataset_nameDPOSamplechosen_messages/rejected_messages。仓库自带的标准格式样例 data/v1_sft_demo.jsonl 第一行即展示了这种结构:

{"messages": [
  {"role": "user", "content": [{"type": "text", "value": "hi"}], "loss_weight": 0.0},
  {"role": "assistant", "content": [{"type": "text", "value": "Hello! I am {{name}}, an AI assistant developed by {{author}}. How can I assist you today?"}], "loss_weight": 1.0}
]}

loss_weight 在引擎层已有实际用途:前缀展开时用它判断“有监督的 assistant 轮”(见 3.3 节)。

5. 使用:初始化与数据访问

5.1 初始化

按设计文档的接口视角,初始化只需传入构建好的 DataArguments

from llamafactory.v1.config.data_args import DataArguments
from llamafactory.v1.core.data_engine import DataEngine

# 1. 创建数据参数
data_args = DataArguments(
    dataset="~/data/v1_sft_demo.jsonl",
    cutoff_len=2048
)

# 2. 初始化 Data Engine
data_engine = DataEngine(data_args=data_args)

# 3. 访问数据
sample = data_engine[0]  # 获取第一个样本

对应到仓库当前实现的调用方式(与 测试用例 和模块 __main__ 一致):

from llamafactory.v1.config.data_args import DataArguments
from llamafactory.v1.core.data_engine import DataEngine

data_args = DataArguments(train_dataset="data/v1_sft_demo.yaml")  # yaml 或标准格式文件路径
data_engine = DataEngine(data_args.train_dataset)

print(data_engine[0])  # 获取第一个标准化样本

也可以在命令行直接冒烟验证(模块内置入口,见 data_engine.py 尾部):

python -m llamafactory.v1.core.data_engine --train_dataset data/v1_sft_demo.yaml
python -m llamafactory.v1.core.data_engine --train_dataset data/v1_dpo_demo.yaml

5.2 数据访问:等价于 Python 列表

实例化后的 DataEngine 支持整数索引、切片、列表索引,用法等价于 Python 列表:

sample = data_engine[0]       # 获取第一个样本
sample = data_engine[0:10]    # 获取前 10 个样本
sample = data_engine[[0, 5, 10]]  # 获取指定索引的样本

底层实现 的路径是:

  • int 索引:直接取 self.data_index[index](即某个 (dataset_name, sample_index, cut) 三元组),经 _get 取原始行 → converter 标准化 → 按 cut 截断 messages 前缀后返回;
  • slice / list 索引:交给 select_data_sample 在全局索引上选择(slice 按 range(*index.indices(...)) 展开,list 逐个取),逐个转换后返回列表。

5.3 流式模式的限制

从源码可以确认以下适用边界:

  • __len__ 对流式数据集返回 -1实现),与文档契约一致;
  • 流式数据集不支持索引访问__getitem__self.streaming 为真时抛 ValueError("Streaming dataset does not support index access.")
  • 当前源码中 __iter__ 尚抛 NotImplementedError,其注释指出流式迭代需要对齐 HF IterableDataset 的 worker id 分片与 shuffle 逻辑。可以推断,文档接口中的 __iter__/__aiter__ 是面向流式场景的既定设计,而当前仓库版本尚未放开该路径,实际使用流式数据前应先确认迭代入口的可用性。

6. 验证与扩展:插件机制与测试

6.1 插件注册机制

数据加载与转换插件都基于 BasePlugin:每个插件族子类通过 __init_subclass__ 拥有独立注册表,register() 把实现绑定到插件名,__call__ 按名字解析路由。现有注册示例:

@DataLoaderPlugin("local").register()
def load_data_from_file(filepath: str, split: str, streaming: bool) -> HFDataset: ...

@DataConverterPlugin("alpaca").register()
def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample: ...

接入新的数据源或格式时,遵循同一模式即可:为 DataLoaderPlugin 注册新的 source 名(在 dataset_info 中以 source 字段引用),或为 DataConverterPlugin 注册新的 converter 名(以 converter 字段引用),无需改动 DataEngine 主体——这正是文档所强调的“其他功能(如数据格式转换、数据加载等)均通过插件的形式实现并接入 DataEngine”的落点。

6.2 测试佐证

tests_v1/core/test_data_engine.py 验证了索引访问与透传语义:以 llamafactory/v1-sft-demo 为标准格式 HF Hub 数据集构造引擎,随机抽取索引断言:

data_args = DataArguments(train_dataset="llamafactory/v1-sft-demo")
data_engine = DataEngine(data_args.train_dataset)
original_data = load_dataset("llamafactory/v1-sft-demo", split="train")
indexes = random.choices(range(len(data_engine)), k=num_samples)
for index in indexes:
    assert data_engine[index] == {"_dataset_name": "default", **original_data[index]}

该断言恰好印证了两件事:标准格式数据经 DataEngine 后仅附加 _dataset_name 元字段、内容逐字节透传;多数据集被拉平为单一全局索引,data_engine[i] 与列表语义一致。

7. 小结与适用前提

  • DataEngine 用“元信息 + 数据体 + 统一索引”的三段式结构,把多数据集加载、格式转换、采样控制解耦:source 决定加载通道(HF Hub 或本地 loader 插件),converter 决定格式归一(alpaca/sharegpt/pair,空则透传),size/weight 决定该数据集在混合索引中的规模,多轮 SFT 还会按有监督 assistant 轮做前缀展开,使 len() 等于真实训练样本数。
  • 使用前提:当前实现位于 src/llamafactory/v1 实验性模块下;文档与源码在个别细节上存在视角差——文档给出的是接口设计(DataArguments.dataset/cutoff_len、二元索引、DataIndexPlugin 命名、hf_hub_url 字段),而仓库当前实现以数据集路径字符串为构造入参、索引为含 cut 的三元组、来源判定基于 source 字段、索引调整由 adjust_data_index 函数承担。以本文 3/5 节结合源码给出的调用方式(如 DataEngine(data_args.train_dataset)python -m llamafactory.v1.core.data_engine)为准可直接运行验证。
  • 流式模式为全有或全无,且当前不支持索引访问、迭代入口尚未实现,生产使用建议以非流式 map 数据集为主。
登录后查看全文
热门项目推荐
相关项目推荐