大模型预训练、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 缓存、数据集统计打印功能,保障混合后的数据质量。高质量的多源混合数据集,是大模型微调效果的关键前提。

Logo

作为“人工智能6S店”的官方数字引擎,为AI开发者与企业提供一个覆盖软硬件全栈、一站式门户。

更多推荐