微调Facebook BART模型时遇CUDA及索引错误求助
我正在基于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 callingcublasCreate(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

