如何在HuggingFace Transformers中配置TPU并行/FSDP解决显存不足
Kaggle TPU 上用 FSDP 微调 Open Llama 3B V2 的配置方案
前置依赖检查
确保环境安装了最新版的 transformers、accelerate、torch_xla,Kaggle TPU笔记本通常预装了适配版本,若有问题可手动更新:
pip install --upgrade transformers accelerate torch_xla[tpu] -f https://storage.googleapis.com/libtpu-releases/index.html
FSDP核心配置要点
针对Open Llama 3B V2的结构,需重点配置以下FSDP参数实现全参数分片:
- 采用
FULL_SHARD分片策略:将模型参数、梯度、优化器状态全部分散到各TPU核心 - 指定自动包裹层:Open Llama基于Llama架构,需将
LlamaLayer设为自动分片的最小单元 - 可选参数卸载:显存仍不足时开启参数卸载到CPU,牺牲少量速度换取显存空间
完整代码调整示例
import torch import torch_xla.core.xla_model as xm import torch_xla.distributed.fsdp as xla_fsdp from transformers import ( AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer, DataCollatorForLanguageModeling ) from transformers.models.llama.modeling_llama import LlamaLayer # 加载模型与Tokenizer model_name = "openlm-research/open_llama_3b_v2" tokenizer = AutoTokenizer.from_pretrained(model_name) tokenizer.pad_token = tokenizer.eos_token # 补充pad token,避免报错 model = AutoModelForCausalLM.from_pretrained(model_name) # 配置FSDP策略 fsdp_config = { "auto_wrap_policy": xla_fsdp.AutoWrapPolicy( transformer_layer_cls={LlamaLayer} ), "sharding_strategy": xla_fsdp.ShardingStrategy.FULL_SHARD, "offload_params": True, # 显存紧张时开启,否则设为False "sync_module_states": True, # 确保各TPU核心参数初始同步 "forward_prefetch": True # 预取前向计算数据,提升速度 } # 用XLA FSDP包裹模型 model = xla_fsdp.XLAFSDPPolicy( model, fsdp_config=fsdp_config, ) # 设置训练参数 training_args = TrainingArguments( output_dir="./open_llama_finetune", per_device_train_batch_size=4, # 单核心batch size,总batch=单核心*TPU核心数 gradient_accumulation_steps=2, # 梯度累积弥补batch大小 num_train_epochs=3, logging_steps=10, save_steps=100, fp16=True, # TPU必须开启混合精度,大幅节省显存 dataloader_num_workers=4, report_to="none", remove_unused_columns=False, ) # 数据处理(示例,替换为你的数据集) def tokenize_function(examples): return tokenizer(examples["text"], truncation=True, max_length=512) # 假设你的原始数据集为dataset,先做tokenization tokenized_dataset = dataset.map(tokenize_function, batched=True) # 分布式采样器适配TPU train_sampler = torch.utils.data.distributed.DistributedSampler( tokenized_dataset, num_replicas=xm.xrt_world_size(), rank=xm.get_ordinal(), shuffle=True, ) # 初始化Trainer trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_dataset, data_collator=DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False), train_sampler=train_sampler, ) # 启动训练 trainer.train()
常见问题处理
- 显存仍不足:进一步调小
per_device_train_batch_size,配合gradient_accumulation_steps保持有效batch规模;或确保fp16=True已开启 - 层包裹错误:若模型结构有自定义修改,需替换对应Transformer层类,Open Llama必须用
LlamaLayer - TPU同步问题:开启
sync_module_states=True可避免各核心参数初始化不一致
内容的提问来源于stack exchange,提问作者starblasters8
相关产品推荐
相关产品推荐

