You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于多主机TPU的LLM高效数据集处理优化方案问询

大语言模型训练数据加载与令牌化优化方案(TPUv4-32 + JAX/Flax)

问题背景

计划在TPUv4-32上使用JAX/Flax训练大语言模型,数据集为Red-Pajama-v2,存储于挂载的Google Cloud Storage桶中,包含5000个.json.gz分片,路径为~/folder-for-bucket/red_pajama/****/en_head.json.gz,每个文件为JSON行格式,文本内容存储在raw_content字段。采用HuggingFace的LlamaTokenizerFast,模型上下文长度为1024 tokens,批量大小设为512。当前实现的令牌流方案存在批次加载过慢问题,脚本受输入瓶颈限制,需优化数据集加载、令牌化及批量迭代流程。

当前实现代码

# ---------- tokenizer ----------
tokenizer = LlamaTokenizerFast.from_pretrained("meta-llama/Llama-2-7b-hf")
tokenizer.pad_token = tokenizer.eos_token

# ---------- streaming dataset ----------
pattern = os.path.join(args.data_dir, "*", "en_head.json.gz")
raw_stream = load_dataset("json", data_files=pattern, split="train", streaming=True)
raw_stream = raw_stream.shard(jax.process_count(), jax.process_index())

# ---------- fast batched tokenizer ----------

def token_stream():
    buf = []
    for ex in raw_stream:
        buf.append(ex["raw_content"])
        if len(buf) >= DOCS_PER_CHUNK:
            for ids in tokenizer(buf, add_special_tokens=False)["input_ids"]:
                yield from ids + [tokenizer.eos_token_id]
            buf.clear()
    # flush remaining docs
    if buf:
        for ids in tokenizer(buf, add_special_tokens=False)["input_ids"]:
            yield from ids + [tokenizer.eos_token_id]

# ---------- token → batch iterator ----------

def batch_iter(global_bsz: int):
    ts, buf = token_stream(), []
    while True:
        buf.extend(itertools.islice(ts, seq_len + 1 - len(buf)))
        if len(buf) < seq_len + 1:
            continue
        seq = np.asarray(buf[:seq_len], dtype=np.int32)
        buf  = buf[seq_len:]
        yield {"input_ids": np.tile(seq[None, :], (global_bsz, 1))}

最优流程方案

高层核心思路

  1. 离线预令牌化+缓存:提前完成文本到令牌的转换,将结果以高效格式存储回GCS,避免训练时重复计算
  2. 并行化数据处理:利用多进程/TPU并行能力处理数据集分片,最大化GCS读取吞吐量
  3. 简化训练时加载逻辑:直接读取预处理后的令牌数据,配合原生批量生成工具,避免手动维护缓冲区和token流拼接

具体优化实现

1. 离线预处理脚本(关键部分)

from datasets import load_dataset
import jax

# 初始化分词器
tokenizer = LlamaTokenizerFast.from_pretrained("meta-llama/Llama-2-7b-hf")
tokenizer.pad_token = tokenizer.eos_token
seq_len = 1024

# 加载原始数据集(非流式,利用多进程并行读取)
pattern = "gs://your-bucket-path/red_pajama/*/en_head.json.gz"
raw_dataset = load_dataset("json", data_files=pattern, split="train")

# 按TPU进程数分片,预处理阶段完成分片避免训练时开销
raw_dataset = raw_dataset.shard(jax.process_count(), jax.process_index())

# 批量令牌化并分割为固定长度序列
def tokenize_and_chunk(examples):
    # 批量处理文本
    tokenized = tokenizer(examples["raw_content"], add_special_tokens=False)
    # 拼接所有token并添加eos标记
    concatenated_tokens = []
    for ids in tokenized["input_ids"]:
        concatenated_tokens.extend(ids + [tokenizer.eos_token_id])
    # 分割为1025长度的块(input取前1024,label取后1024)
    chunked_tokens = [concatenated_tokens[i:i+seq_len+1] for i in range(0, len(concatenated_tokens)-seq_len, seq_len)]
    return {
        "input_ids": [chunk[:seq_len] for chunk in chunked_tokens],
        "labels": [chunk[1:] for chunk in chunked_tokens]
    }

# 并行执行令牌化,保存为Parquet格式(GCS友好的高效存储格式)
tokenized_dataset = raw_dataset.map(
    tokenize_and_chunk,
    batched=True,
    batch_size=1000,
    num_proc=8,  # 根据机器核心数调整
    remove_columns=["raw_content"]
)

# 保存预处理结果到GCS
tokenized_dataset.save_to_disk("gs://your-bucket-path/tokenized_red_pajama")

2. 训练时数据加载代码

from datasets import load_from_disk

# 加载预令牌化后的数据集
tokenized_dataset = load_from_disk("gs://your-bucket-path/tokenized_red_pajama")

# 转换为JAX格式,直接生成训练批次
tokenized_dataset = tokenized_dataset.with_format("jax")
train_loader = tokenized_dataset.shuffle(seed=42).batch(512)

# 迭代训练
for batch in train_loader:
    # batch包含input_ids和labels,直接传入训练步骤
    train_step(batch)

额外性能优化点

  • GCS读取优化:使用gcsfuse挂载时添加--max-conns-per-host=100和--implicit-dirs参数提升并行读取速度;或直接使用Datasets的GCS原生支持,避免挂载开销
  • TPU并行适配:利用JAX的pmap或Flax的TrainState自动处理数据并行,确保每个TPU核心获取独立批次
  • 流式加载备选方案:若必须使用流式模式,改用tf.data.Dataset从GCS读取文件,配合tf.data.experimental.AUTOTUNE优化并行度后转换为JAX数组

内容的提问来源于stack exchange,提问作者innerproduct

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 17:29:52