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

PyTorch神经机器翻译报错:维度1序列长度不匹配(8/9)

PyTorch神经机器翻译代码报错:ValueError: expected sequence of length 8 at dim 1 (got 9)

问题详情

运行基于PyTorch的神经机器翻译代码时触发上述报错,此前代码运行正常。已确认translation_src与translation_target序列长度一致,自定义padding和collate函数逻辑无误,自行重写collate内的padding逻辑后问题仍存在。

相关代码

class TranslationDataset(Dataset):
    def __init__(self, dataset):
        self.dataset = dataset

    def __len__(self):
        return len(self.dataset)

    def __getitem__(self, idx):
        src_encoded=self.dataset[idx]['translation_src']
        trg_encoded=self.dataset[idx]['translation_trg']

        # Determine the maximum sequence length
        max_len = max(len(src_encoded), len(trg_encoded))

        # Pad the sequences to have the same length
        src_encoded = src_encoded + [0]*(max_len - len(src_encoded))
        trg_encoded = trg_encoded + [0]*(max_len - len(trg_encoded))
        return (
            torch.tensor(src_encoded),
            torch.tensor(trg_encoded),
        )

train_ds = TranslationDataset(data['train'])
val_ds = TranslationDataset(data['test'])

def pad_collate_fn(batch):
    src_sentences,trg_sentences=[],[]
    for sample in batch:
        src_sentences+=[sample[0]]
        trg_sentences+=[sample[1]]

    src_sentences = pad_sequence(src_sentences, batch_first=True, padding_value=0)
    trg_sentences = pad_sequence(trg_sentences, batch_first=True, padding_value=0)

    return src_sentences, trg_sentences

def chunk(indices, chunk_size):
    return torch.split(torch.tensor(indices), chunk_size)

class CustomBatchSampler(Sampler):
    def __init__(self, dataset, batch_size):

        # Dataset is already sorted so just chunk indices
        # into batches of indices for sampling
        self.batch_size=batch_size
        self.indices=range(len(dataset))
        self.batch_of_indices=list(chunk(self.indices, self.batch_size))
        self.batch_of_indices = [batch.tolist() for batch in self.batch_of_indices]

    def __iter__(self):
        random.shuffle(self.batch_of_indices)
        return iter(self.batch_of_indices)

    def __len__(self):
        return len(self.batch_of_indices)
     

custom_batcher_train = CustomBatchSampler(train_ds, config['BATCH_SIZE'])
custom_batcher_val = CustomBatchSampler(val_ds, config['BATCH_SIZE'])

# example-use
dummy_batcher = CustomBatchSampler(train_ds, 3)
dummy_dl=DataLoader(train_ds, collate_fn=pad_collate_fn , batch_sampler=dummy_batcher, pin_memory=True)
for x ,y  in dummy_dl:
    print('Shapes: ')
    print('-'*10)
    print(x.size())
    print(y.size())
    print()

    print('e.g. src batch (see there is minimal/no padding):')
    print('-'*10)
    print(x.numpy())

    break

报错栈信息

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
<ipython-input-103-b82fc7198b80> in <cell line: 30>()
     28 dummy_batcher = CustomBatchSampler(train_ds, 3)
     29 dummy_dl=DataLoader(train_ds, collate_fn=pad_collate_fn , batch_sampler=dummy_batcher, pin_memory=True)
---> 30 for x ,y  in dummy_dl:
     31     print('Shapes: ')
     32     print('-'*10)

4 frames
/usr/local/lib/python3.10/dist-packages/torch/utils/data/dataloader.py in __next__(self)
    628                 # TODO(https://github.com/pytorch/pytorch/issues/76750)
    629                 self._reset()  # type: ignore[call-arg]
---> 630             data = self._next_data()
    631             self._num_yielded += 1
    632             if self._dataset_kind == _DatasetKind.Iterable and \

/usr/local/lib/python3.10/dist-packages/torch/utils/data/dataloader.py in _next_data(self)
    672     def _next_data(self):
    673         index = self._next_index()  # may raise StopIteration
---> 674         data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    675         if self._pin_memory:
    676             data = _utils.pin_memory.pin_memory(data, self._pin_memory_device)

/usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/fetch.py in fetch(self, possibly_batched_index)
     47         if self.auto_collation:
     48             if hasattr(self.dataset, "__getitems__") and self.dataset.__getitems__:
---> 49                 data = self.dataset.__getitems__(possibly_batched_index)
     50             else:
     51                 data = [self.dataset[idx] for idx in possibly_batched_index]

/usr/local/lib/python3.10/dist-packages/datasets/arrow_dataset.py in __getitems__(self, keys)
   2805     def __getitems__(self, keys: List) -> List:
   2806         """Can be used to get a batch using a list of integers indices."""
-> 2807         batch = self.__getitem__(keys)
   2808         n_examples = len(batch[next(iter(batch))])
   2809         return [{col: array[i] for col, array in batch.items()} for i in range(n_examples)]

<ipython-input-101-977a8cf27fd7> in __getitem__(self, idx)
     17         trg_encoded = trg_encoded + [0]*(max_len - len(trg_encoded))
     18         return [
---> 19             torch.tensor(src_encoded),
     20             # torch.tensor(trg_encoded),
     21         ]

ValueError: expected sequence of length 8 at dim 1 (got 9)

报错原因分析

报错根源在于TranslationDataset的__getitem__方法未兼容批量索引的情况:

  • 当使用自定义BatchSampler时,PyTorch的DataLoader会尝试调用数据集的__getitems__方法批量获取样本
  • 你使用的datasets.ArrowDataset实现了__getitems__,会将批量索引传入__getitem__,此时self.dataset[idx]['translation_src']返回的是一组序列而非单个序列
  • 原代码中针对单个序列的padding逻辑被错误地应用在批量序列上,导致生成的src_encoded维度混乱,无法转换为张量

解决方案

方案1:修改Dataset的__getitem__方法,兼容批量索引

class TranslationDataset(Dataset):
    def __init__(self, dataset):
        self.dataset = dataset

    def __len__(self):
        return len(self.dataset)

    def __getitem__(self, idx):
        # 处理批量索引
        if isinstance(idx, list):
            samples = self.dataset[idx]
            src_list = samples['translation_src']
            trg_list = samples['translation_trg']
            processed_samples = []
            for src, trg in zip(src_list, trg_list):
                max_len = max(len(src), len(trg))
                src_padded = src + [0]*(max_len - len(src))
                trg_padded = trg + [0]*(max_len - len(trg))
                processed_samples.append((torch.tensor(src_padded), torch.tensor(trg_padded)))
            return processed_samples
        # 处理单个索引
        else:
            src_encoded = self.dataset[idx]['translation_src']
            trg_encoded = self.dataset[idx]['translation_trg']
            max_len = max(len(src_encoded), len(trg_encoded))
            src_encoded = src_encoded + [0]*(max_len - len(src_encoded))
            trg_encoded = trg_encoded + [0]*(max_len - len(trg_encoded))
            return (torch.tensor(src_encoded), torch.tensor(trg_encoded))

方案2:关闭ArrowDataset的批量获取功能

让DataLoader逐个调用__getitem__,无需修改Dataset逻辑:

train_ds = TranslationDataset(data['train'])
train_ds.dataset.__getitems__ = None  # 关闭批量获取
val_ds = TranslationDataset(data['test'])
val_ds.dataset.__getitems__ = None

方案3:将padding逻辑完全移到collate函数中

Dataset仅返回原始序列,避免处理批量索引的复杂情况:

class TranslationDataset(Dataset):
    def __init__(self, dataset):
        self.dataset = dataset

    def __len__(self):
        return len(self.dataset)

    def __getitem__(self, idx):
        # 直接返回未padding的原始序列
        src_encoded = self.dataset[idx]['translation_src']
        trg_encoded = self.dataset[idx]['translation_trg']
        return (torch.tensor(src_encoded), torch.tensor(trg_encoded))

def pad_collate_fn(batch):
    src_sentences, trg_sentences = [], []
    for sample in batch:
        src_sentences.append(sample[0])
        trg_sentences.append(sample[1])
    
    # 先保证每个样本的src和trg长度一致
    processed_src, processed_trg = [], []
    for src, trg in zip(src_sentences, trg_sentences):
        max_len = max(len(src), len(trg))
        src_padded = torch.cat([src, torch.zeros(max_len - len(src), dtype=src.dtype)])
        trg_padded = torch.cat([trg, torch.zeros(max_len - len(trg), dtype=trg.dtype)])
        processed_src.append(src_padded)
        processed_trg.append(trg_padded)
    
    # 再对整个batch做统一padding
    src_batch = pad_sequence(processed_src, batch_first=True, padding_value=0)
    trg_batch = pad_sequence(processed_trg, batch_first=True, padding_value=0)
    return src_batch, trg_batch

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 21:40:54