MindSpore Transformers LLM 数据预处理:支持多源数据集混合
大模型预训练、SFT 监督微调场景,经常需要同时使用多个来源数据集,例如通用对话数据集、指令数据集、领域知识库数据。多源数据集混合核心难点:不同数据源格式不一致、样本权重配比、样本洗牌、格式统一转换、边界过滤、分布式训练数据分片。本文基于 MindSpore‑Transformers 实现多源数据集混合预处理,支持权重采样、格式适配、缓存加速。
环境:MindSpore 2.4,mindspore‑transformers,昇腾 NPU;数据集格式为 jsonl。
多源混合预处理设计思路
多数据集加载:分别读取不同路径下的数据集文件;
格式适配器:不同来源数据做字段映射,统一转为 prompt‑response 标准格式;
加权混合采样:设置各数据集采样权重,控制每类数据在训练中的占比;
全局 shuffle:混合后整体打乱,避免数据集顺序带来训练偏差;
分词与组批:统一 tokenizer 处理,过滤超长、非法样本;
分布式适配:支持多卡训练数据集分片,避免多卡重复读取样本。
数据集样例
数据集 A:通用指令数据集instruct_data.jsonl
{"instruction":"解释什么是昇思MindSpore","output":"昇思MindSpore是全栈AI框架"}
数据集 B:领域对话数据集domain_chat.jsonl
{"question":"什么是Ascend‑C","answer":"Ascend‑C是昇腾算子开发语言"}
1. 数据集适配器,统一不同数据源字段
不同数据集 key 不一致,通过适配器函数映射为统一prompt、target字段。
# multi_source_dataset.py
import json
import os
import random
from typing import List,Dict
import mindspore as ms
from mindspore.dataset import GeneratorDataset
from mindspore_transformers import AutoTokenizer
def load_jsonl(path:str)->List[Dict]:
"""加载jsonl数据集"""
samples = []
with open(path,"r",encoding="utf‑8") as f:
for line in f:
line = line.strip()
if not line:
continue
samples.append(json.loads(line))
return samples
# 适配器函数,不同数据源做字段转换
def adapter_instruct(sample:Dict):
"""指令数据集适配器"""
prompt = f"用户:{sample['instruction']}\n助手:"
target = sample["output"]
return {"prompt":prompt,"target":target}
def adapter_domain_chat(sample:Dict):
"""领域对话数据集适配器"""
prompt = f"用户:{sample['question']}\n助手:"
target = sample["answer"]
return {"prompt":prompt,"target":target}
DATASET_ADAPTER_MAP = {
"instruct": adapter_instruct,
"domain": adapter_domain_chat
}
class MultiSourceMixDataset:
"""
多源数据集混合
dataset_list: [{"path":"xxx.jsonl","type":"instruct","weight":1.0},...]
"""
def __init__(self,dataset_list:List[Dict], global_shuffle=True):
self.dataset_list = dataset_list
self.global_shuffle = global_shuffle
self.all_samples = []
self._load_and_mix()
def _load_and_mix(self):
for ds_cfg in self.dataset_list:
path = ds_cfg["path"]
ds_type = ds_cfg["type"]
weight = ds_cfg["weight"]
raw_data = load_jsonl(path)
adapter_func = DATASET_ADAPTER_MAP[ds_type]
adapted = [adapter_func(s) for s in raw_data]
# 根据权重做重复采样,实现数据集比例控制
sample_num = int(len(adapted)*weight)
sampled = random.choices(adapted,k=sample_num)
self.all_samples.extend(sampled)
# 全局打乱全部样本
if self.global_shuffle:
random.shuffle(self.all_samples)
def __len__(self):
return len(self.all_samples)
def __getitem__(self,idx):
item = self.all_samples[idx]
return item["prompt"], item["target"]
2. 封装 MindSpore GeneratorDataset,接入训练流水线
结合 tokenizer,拼接 prompt 与 target,完成 SFT 格式 token 化,过滤超长样本。
def sft_tokenize_fn(prompt,target,tokenizer,max_seq_len=512):
"""SFT样本分词拼接"""
full_text = prompt + target + tokenizer.eos_token
token_out = tokenizer(
full_text,
max_length=max_seq_len,
truncation=True,
padding="max_length",
return_tensors="ms"
)
input_ids = token_out["input_ids"][0]
attention_mask = token_out["attention_mask"][0]
# SFT label与input_ids一致
labels = input_ids.copy()
return input_ids,attention_mask,labels
def build_mix_train_dataset(tokenizer,max_seq_len=512):
dataset_config = [
{
"path":"./instruct_data.jsonl",
"type":"instruct",
"weight":0.8
},
{
"path":"./domain_chat.jsonl",
"type":"domain",
"weight":1.2
}
]
raw_ds = MultiSourceMixDataset(dataset_config,global_shuffle=True)
gen_ds = GeneratorDataset(
source=raw_ds,
column_names=["prompt","target"],
shuffle=False
)
# map分词处理
def map_func(prompt,target):
return sft_tokenize_fn(prompt,target,tokenizer,max_seq_len)
train_ds = gen_ds.map(operations=map_func)
return train_ds
if __name__ == "__main__":
ms.set_context(mode=ms.GRAPH_MODE,device_target="Ascend")
tokenizer = AutoTokenizer.from_pretrained("llama2‑7b‑zh")
train_dataset = build_mix_train_dataset(tokenizer,max_seq_len=512)
train_dataset = train_dataset.batch(4)
for batch in train_dataset.create_tuple_iterator():
input_ids,attn_mask,labels = batch
print(f"batch input_ids shape:{input_ids.shape}")
break
3. 分布式训练适配代码
多卡训练场景,数据集需要分片,避免每张卡加载全部样本,造成数据重复。
def get_dist_dataset(train_ds:GeneratorDataset,batch_size:int):
"""分布式数据集分片"""
from mindspore.communication import get_rank,get_group_size
rank_id = get_rank()
rank_size = get_group_size()
# 按rank分片
train_ds = train_ds.shard(num_shards=rank_size,shard_id=rank_id)
train_ds = train_ds.batch(batch_size,drop_remainder=True)
return train_ds
4. 数据集缓存与过滤扩展
增加样本过滤逻辑,过滤空文本、过长文本,开启数据集缓存加速预处理。
def filter_sample_func(prompt,target):
"""过滤无效样本"""
if len(prompt.strip()) ==0 or len(target.strip())==0:
return False
if len(prompt+target) > 1800:
return False
return True
# 在MultiSourceMixDataset加载阶段增加过滤
# filtered = [x for x in adapted if filter_sample_func(x["prompt"],x["target"])]
工程调优要点
权重采样:weight 不等于数据集比例,weight 大于 1 会扩充样本,小于 1 做降采样;领域小数据集可以调高 weight 提升占比。
全局 shuffle:必须在多源合并之后做全局 shuffle,否则训练会出现按数据集顺序训练,收敛效果变差。
分布式 shard:GeneratorDataset 必须调用 shard,保证多卡之间数据互不重复。
格式适配器扩展:新增数据集只需要新增 adapter 函数,不需要修改主流程,便于接入更多第三方开源数据集。
性能优化:预处理耗时大时,可以将预处理结果保存为 mindrecord 格式,训练直接读取 mindrecord,避免重复预处理。
样本质量优先,混合后建议做数据统计,打印各类数据集样本数量,确认配比符合预期。
总结
多源数据集混合预处理是 LLM 微调的基础工程模块。本文实现完整的多源数据集加载方案:通过适配器模式统一不同数据集字段,支持权重采样控制各数据源占比,全局 shuffle 打乱样本,对接 MindSpore GeneratorDataset,完成 SFT 分词处理,同时提供分布式分片适配。该框架可以灵活接入指令集、对话、领域知识库等多种格式数据集。实际项目开发中,建议增加样本过滤、mindrecord 缓存、数据集统计打印功能,保障混合后的数据质量。高质量的多源混合数据集,是大模型微调效果的关键前提。
鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。
更多推荐



所有评论(0)