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

在Kaggle训练PyTorch Transformer遇CUDA设备断言触发错误

解决Transformer翻译模型在大数据集下的CUDA断言错误

我花了两天排查仍未解决:在Kaggle上训练英语到阿拉伯语的Transformer翻译模型,用3万条句子对的小数据集完全正常,但用6万条以上数据集时,持续出现RuntimeError: CUDA error: device-side assert triggered错误,无法定位原因。

完整代码

"""
torch modules
with ofc  numpy  and pandas
"""

import pandas as pd
import numpy as np 

from torch import nn
import torch
from torchtext import data
from torch.nn  import functional as F
import torch.optim as  optim 
if torch.cuda.is_available():  
  dev = "cuda:0" 

  print("gpu up")
else:  
  dev = "cpu"  
device = torch.device(dev)

import random
SEED= 32

"""
regex and the tokenizers
"""

import re
from spacy.tokenizer import Tokenizer
from spacy.lang.en import English
from spacy.lang.ar import Arabic
from nltk.translate.bleu_score import sentence_bleu

enNLP = English()
arNLP = Arabic()

enTokenizer = Tokenizer(enNLP.vocab)
arTokenizer =  Tokenizer(arNLP.vocab)

df = pd.read_csv("/kaggle/input/translation-with-transformers/opus-ted.txt",encoding="utf-8",delimiter="\t\t",names=["eng","ar"])

"""
defining the tokenizers for arabic and english  

creating the fields for the dataset from torchtext 
that class is the simple way I could find for turning a df into a torch dataset

نهها and ببدأ are just arbitrary words for init and end of sentence tokens  
for some reason when I choose an arabic word for the unknown token  the vocab doesn't replace words that are not in the vocab  
"""

def myTokenizerEN(x):
 return  [word.text for word in 
          enTokenizer(re.sub(r"\s+\s+"," ",re.sub(r"[\.\'\`\"\r+\n+]"," ",x.lower())).strip())]
def myTokenizerAR(x):
 return  [word.text for word in 
          arTokenizer(re.sub(r"\s+\s+"," ",re.sub(r"[\.\'\`\"\r+\n+]"," ",x.lower())).strip())]

SRC = data.Field(tokenize=myTokenizerEN,batch_first=False,init_token="<sos>",eos_token="<eos>")
TARGET = data.Field(tokenize=myTokenizerAR,batch_first=False,tokenizer_language="ar",init_token="ببدأ",eos_token="نهها")

class DataFrameDataset(data.Dataset):

    def __init__(self, df, src_field, target_field, is_test=False, **kwargs):
        fields = [('eng', src_field), ('ar',target_field)]
        examples = []
        for i, row in df.iterrows():
            eng = row.eng 
            ar = row.ar
            examples.append(data.Example.fromlist([eng, ar], fields))

        super().__init__(examples, fields, **kwargs)

        
torchdataset = DataFrameDataset(df,SRC,TARGET)


train_data, valid_data = torchdataset.split(split_ratio=0.8, random_state = random.seed(SEED))

SRC.build_vocab(train_data,min_freq=2)
TARGET.build_vocab(train_data,min_freq=2)  


"""
we are using batches for validation and test set because of memory usage we can't pass the whole set at once

try lowering the batch size if you are out of memory 
"""
BATCH_SIZE = 64

train_iterator,valid_iterator = data.BucketIterator.splits(
    (train_data,valid_data), 
    batch_size = BATCH_SIZE,
    device = device,
    sort=False,
    sort_within_batch=False,
    shuffle=True)

#No. of unique tokens in text
src_vocab_size  = len(SRC.vocab)
print("Size of english vocabulary:",src_vocab_size)

#No. of unique tokens in label
trg_vocab_size =len(TARGET.vocab)
print("Size of arabic vocabulary:",trg_vocab_size)

num_heads = 8
num_encoder_layers = 3
num_decoder_layers = 3

max_len= 227 #227
embedding_size= 256 #256
src_pad_idx =SRC.vocab.stoi["<pad>"]


model = TranslateTransformer(
    embedding_size,
    src_vocab_size,
    trg_vocab_size,
    src_pad_idx,
    num_heads,
    num_encoder_layers,
    num_decoder_layers,
    max_len
).to(device)

loss_track = []
loss_validation_track= []

"""
I'm using adagrad because it assigns bigger updates to less 
frequently updated weights so thought it could be useful for 
words not used a lot.
"""

optimizer = optim.Adagrad(model.parameters(),lr = 0.003)
EPOCHS = 15

pad_idx = SRC.vocab.stoi["<pad>"]
criterion = nn.CrossEntropyLoss(ignore_index=pad_idx) 

for i in range(0,EPOCHS):
    stepLoss=[]
    model.train() # the training mode for the model (applies dropout and batchnorms)
    for batch  in train_iterator:
        input_sentence = batch.eng.to(device)
        trg = batch.ar.to(device)

        optimizer.zero_grad()
        out = model(input_sentence,trg[:-1])
        out = out.reshape(-1,trg_vocab_size)
        trg = trg[1:].reshape(-1)
        loss = criterion(out,trg)
        
        
        loss.backward()
        optimizer.step()
        stepLoss.append(loss.item())
        

    loss_track.append(np.mean(stepLoss))
    print("train crossentropy at epoch {} loss: ".format(i),np.mean(stepLoss))
    
    stepValidLoss=[]
    model.eval() # the evaluation mode for the model (doesn't apply dropout and batchNorm)
    for batch  in valid_iterator:
        input_sentence = batch.eng.to(device)
        trg = batch.ar.to(device)

        optimizer.zero_grad()
        out = model(input_sentence,trg[:-1])
        out = out.reshape(-1,trg_vocab_size)
        trg = trg[1:].reshape(-1)
        loss = criterion(out,trg)
        
        stepValidLoss.append(loss.item())
  
    loss_validation_track.append(np.mean(stepValidLoss))
    print("validation crossentropy at epoch {} loss: ".format(i),np.mean(stepValidLoss))

报错堆栈

---------------------------------------------------------------------------
RuntimeError                              Traceback (most recent call last)
<ipython-input-30-c9f694cb9d66> in <module>
     25     num_decoder_layers,
     26     max_len
---> 27 ).to(device)

/opt/conda/lib/python3.7/site-packages/torch/nn/modules/module.py in to(self, *args, **kwargs)
    441             return t.to(device, dtype if t.is_floating_point() else None, non_blocking)
    442 
---> 443         return self._apply(convert)
    444 
    445     def register_backward_hook(self, hook):

/opt/conda/lib/python3.7/site-packages/torch/nn/modules/module.py in _apply(self, fn)
    201     def _apply(self, fn):
    202         for module in self.children():
---> 203             module._apply(fn)
    204 
    205         def compute_should_use_set_data(tensor, tensor_applied):

/opt/conda/lib/python3.7/site-packages/torch/nn/modules/module.py in _apply(self, fn)
    223                 # `with torch.no_grad()`:
    224                 with torch.no_grad():
---> 225                     param_applied = fn(param)
    226                 should_use_set_data = compute_should_use_set_data(param, param_applied)
    227                 if should_use_set_data:

/opt/conda/lib/python3.7/site-packages/torch/nn/modules/module.py in convert(t)
    439             if convert_to_format is not None and t.dim() == 4:
    440                 return t.to(device, dtype if t.is_floating_point() else None, non_blocking, memory_format=convert_to_format)
---> 441             return t.to(device, dtype if t.is_floating_point() else None, non_blocking)
    442 
    443         return self._apply(convert)

RuntimeError: CUDA error: device-side assert triggered

可能的原因及解决方法

  • 目标序列pad索引错误
    代码中损失函数的ignore_index用了源语言的pad索引SRC.vocab.stoi["<pad>"],但目标序列的pad索引应该是TARGET.vocab.stoi["<pad>"],两者大概率不同。大数据集下目标序列pad token出现次数更多,直接触发断言错误。修改代码:

    pad_idx = TARGET.vocab.stoi["<pad>"]
    criterion = nn.CrossEntropyLoss(ignore_index=pad_idx)
    
  • 词汇表维度不匹配
    大数据集下词汇表规模变大,需确认TranslateTransformer类的输出层维度是否严格等于传入的trg_vocab_size,如果输出层维度小于实际目标词汇表长度,会导致索引越界。

  • 最大序列长度不足
    大数据集中可能存在更长的句子,max_len=227可能无法覆盖所有序列长度,导致模型处理时出现索引错误。先统计数据集句子的最大长度再调整:

    # 统计英文和阿拉伯语句子的最大token长度
    max_src_len = max(len(myTokenizerEN(s)) for s in df['eng'])
    max_trg_len = max(len(myTokenizerAR(s)) for s in df['ar'])
    # max_len设为两者最大值加2(包含sos和eos token)
    max_len = max(max_src_len, max_trg_len) + 2
    
  • CUDA内存溢出
    大数据集+大batch size可能导致CUDA内存不足,有时会以断言错误的形式表现。尝试降低BATCH_SIZE(比如改成32或16),或者启用梯度累积:

    accum_steps = 2  # 每2个batch更新一次参数
    for batch_idx, batch in enumerate(train_iterator):
        input_sentence = batch.eng.to(device)
        trg = batch.ar.to(device)
        
        out = model(input_sentence, trg[:-1])
        out = out.reshape(-1, trg_vocab_size)
        trg = trg[1:].reshape(-1)
        loss = criterion(out, trg)
        loss = loss / accum_steps  # 均分损失
        
        loss.backward()
        
        if (batch_idx + 1) % accum_steps == 0:
            optimizer.step()
            optimizer.zero_grad()
    
  • 无效数据导致Tokenizer异常
    大数据集中可能存在空句子或格式异常的句子,导致Tokenizer输出空序列。构建数据集前过滤无效数据:

    # 过滤空句子
    df = df[(df['eng'].str.strip() != "") & (df['ar'].str.strip() != "")]
    # 过滤tokenize后为空的句子
    valid_indices = []
    for idx, row in df.iterrows():
        if len(myTokenizerEN(row['eng'])) > 0 and len(myTokenizerAR(row['ar'])) > 0:
            valid_indices.append(idx)
    df = df.loc[valid_indices]
    

内容的提问来源于stack exchange,提问作者Abdalnassir Ghzawi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 16:59:51