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
相关产品推荐
相关产品推荐

