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

微调Facebook BART模型时遇CUDA及索引错误求助

微调BART模型文本分类时的GPU/CPU错误排查与解决

我正在基于Hugging Face Trainer微调Facebook BART模型实现文本分类,核心代码如下:

训练参数与模型初始化

training_args = TrainingArguments(
    output_dir=model_directory,      # output directory
    num_train_epochs=1,              # total number of training epochs - 3
    per_device_train_batch_size=4,  # batch size per device during training - 16
    per_device_eval_batch_size=16,   # batch size for evaluation - 64
    warmup_steps=50,                # number of warmup steps for learning rate scheduler - 500
    weight_decay=0.01,               # strength of weight decay
    logging_dir=model_logs,          # directory for storing logs
    logging_steps=10,
)

model = BartForSequenceClassification.from_pretrained("facebook/bart-large-mnli") # bart-large-mnli

trainer = Trainer(
    model=model,                          # the instantiated 🤗 Transformers model to be trained
    args=training_args,                   # training arguments, defined above
    compute_metrics=new_compute_metrics,  # a function to compute the metrics
    train_dataset=train_dataset,          # training dataset
    eval_dataset=val_dataset              # evaluation dataset
)

Tokenizer代码

from transformers import BartTokenizerFast
tokenizer = BartTokenizerFast.from_pretrained('facebook/bart-large-mnli')

GPU集群错误(Databricks Runtime 12.2 LTS ML,Standard_NC24s_v3)

调用trainer.train()时出现以下错误:

RuntimeError: Caught RuntimeError in replica 1 on device 1.
Original Traceback (most recent call last):
File "/databricks/python/lib/python3.9/site-packages/torch/nn/parallel/parallel_apply.py", line 61, in _worker
output = module(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/transformers/models/bart/modeling_bart.py", line 1496, in forward
outputs = self.model(
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/transformers/models/bart/modeling_bart.py", line 1222, in forward
encoder_outputs = self.encoder(
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/transformers/models/bart/modeling_bart.py", line 846, in forward
layer_outputs = encoder_layer(
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/transformers/models/bart/modeling_bart.py", line 323, in forward
hidden_states, attn_weights, _ = self.self_attn(
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/transformers/models/bart/modeling_bart.py", line 191, in forward
query_states = self.q_proj(hidden_states) * self.scaling
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
return forward_call(*input, **kwargs)
File "/databricks/python/lib/python3.9/site-packages/torch/nn/modules/linear.py", line 114, in forward
return F.linear(input, self.weight, self.bias)
RuntimeError: CUDA error: CUBLAS_STATUS_NOT_INITIALIZED when calling cublasCreate(handle)

解决方法

  • 显式初始化CUDA设备:在模型加载后添加model = model.to('cuda'),或在TrainingArguments中指定device='cuda:0';同时在代码开头加入import torch; torch.cuda.init()确保CUBLAS库正确初始化
  • 适配多GPU分布式训练:在TrainingArguments中设置ddp_find_unused_parameters=False,因为BART部分参数在文本分类任务中可能未被激活,导致分布式训练报错
  • 进一步降低显存负载:启用gradient_accumulation_steps=4减少单步显存占用,或先换用轻量版facebook/bart-base-mnli验证训练流程正确性

CPU集群错误(Databricks Runtime 12.1 ML,Standard_L8s)

偶尔出现以下错误(错误数字会变化):

IndexError: Target 11 is out of bounds.

解决方法

  • 匹配模型类别数:facebook/bart-large-mnli默认是3分类模型,如果你的任务是多分类(比如12类),加载模型时必须指定num_labels参数,示例:
    model = BartForSequenceClassification.from_pretrained("facebook/bart-large-mnli", num_labels=12)
    
  • 清洗数据集异常标签:遍历训练/验证集,过滤掉标签值大于等于num_labels的异常样本
  • 确认标签格式正确性:确保数据集的标签列是整数类型,无字符串或格式错误

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 04:12:00