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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 21:45:08