如何在BucketIterator中迭代分词后的PyTorch Multi30k数据集并解决sort_key报错
PyTorch+torchtext机器翻译任务BucketIterator sort_key报错解决方案
报错的核心原因是新旧版torchtext的工具链混用:你使用了torchtext原生的新版Multi30k数据集,但搭配的是torchtext.legacy体系下的Field、BucketIterator,两者格式不兼容:新版Multi30k返回PyTorch原生Dataset实例,没有legacy体系要求的sort_key属性,且你没有将定义的分词规则、Field与数据集做绑定,导致迭代器初始化失败。
修改步骤
- 第一步:修改Multi30k导入路径,改用legacy体系下的数据集实现,和后续工具链匹配
- 第二步:加载数据集时传入
fields参数,将你定义的英文、德文Field和数据集的两个语种字段绑定 - 第三步:BucketIterator初始化时显式传入
sort_key参数,指定排序规则(通常按源语言句子长度排序,减少padding冗余)
修改后可运行代码
import spacy # 改1:从legacy路径导入Multi30k from torchtext.legacy.datasets import Multi30k from torchtext.legacy.data import Field, BucketIterator spacy_eng = spacy.load("en_core_web_sm") spacy_ger = spacy.load("de_core_news_sm") def tokenize_eng(text): return [tok.text for tok in spacy_eng.tokenizer(text)] def tokenize_ger(text): return [tok.text for tok in spacy_ger.tokenizer(text)] english = Field(sequential=True, use_vocab=True, tokenize=tokenize_eng, lower=True, init_token='<sos>', eos_token='<eos>') german = Field(sequential=True, use_vocab=True, tokenize=tokenize_ger, lower=True, init_token='<sos>', eos_token='<eos>') # 改2:添加fields参数,和数据集的语种顺序对应,这里language_pair是('en','de'),所以fields顺序是英文、德文 train, valid, test = Multi30k(root=".data", split=('train', 'valid', 'test'), language_pair=('en', 'de'), fields = [('en', english), ('de', german)]) # 现在可以正常构建词表 english.build_vocab(train, max_size=10000, min_freq=2) german.build_vocab(train, max_size=10000, min_freq=2) # 改3:添加sort_key参数,按英文句子长度排序(如果源语言是德文就改成len(x.de)),同时开启batch内排序 train_data, valid_data, test_data = BucketIterator.splits((train, valid, test), batch_size=64, device='cuda', sort_key=lambda x: len(x.en), sort_within_batch=True)
如果不想用legacy工具链,可以改用新版torchtext的分词、词表工具配合PyTorch原生DataLoader实现同等功能,上述修改为最小改动方案,可直接解决当前AttributeError问题。
内容的提问来源于stack exchange,提问作者Devanshu
相关产品推荐
相关产品推荐

