基于多主机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))}
最优流程方案
高层核心思路
- 离线预令牌化+缓存:提前完成文本到令牌的转换,将结果以高效格式存储回GCS,避免训练时重复计算
- 并行化数据处理:利用多进程/TPU并行能力处理数据集分片,最大化GCS读取吞吐量
- 简化训练时加载逻辑:直接读取预处理后的令牌数据,配合原生批量生成工具,避免手动维护缓冲区和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
相关产品推荐
相关产品推荐

