使用torch.cat()拼接张量时报错:'tuple'对象不支持元素赋值
问题:使用torch.cat()拼接PyTorch张量时遇到'tuple' object does not support item assignment错误
我尝试使用torch.cat()拼接PyTorch张量,但遇到错误提示:'tuple' object does not support item assignment。以下是我的代码:
inputs = tokenizer.encode_plus(txt, add_special_tokens=False, return_tensors="pt") input_id_chunks = inputs["input_ids"][0].split(510) mask_chunks = inputs["attention_mask"][0].split(510) print(type(input_id_chunks)) for i in range(len(input_id_chunks)): print(type(input_id_chunks[i])) print(input_id_chunks[i]) input_id_chunks[i] = torch.cat([ torch.Tensor([101]), input_id_chunks[i], torch.Tensor([102]) ])
打印结果显示input_id_chunks类型为tuple,其元素为torch.Tensor,但执行时触发TypeError: 'tuple' object does not support item assignment错误。单独测试torch.cat()可以正常运行,不清楚原代码哪里有问题。
问题原因
torch.Tensor.split()方法返回的是不可变的元组(tuple),元组不支持直接修改元素值,所以循环中尝试给input_id_chunks[i]赋值会触发错误。
解决方案
把元组转换为可变的列表(list),即可正常修改元素:
修改后的代码
inputs = tokenizer.encode_plus(txt, add_special_tokens=False, return_tensors="pt") # 将split返回的元组转为列表 input_id_chunks = list(inputs["input_ids"][0].split(510)) mask_chunks = list(inputs["attention_mask"][0].split(510)) print(type(input_id_chunks)) for i in range(len(input_id_chunks)): print(type(input_id_chunks[i])) print(input_id_chunks[i]) # 保持张量类型一致,避免浮点/整型不匹配问题 cls_token = torch.tensor([101], dtype=input_id_chunks[i].dtype) sep_token = torch.tensor([102], dtype=input_id_chunks[i].dtype) input_id_chunks[i] = torch.cat([cls_token, input_id_chunks[i], sep_token])
额外说明
原代码中torch.Tensor([101])创建的是浮点型张量,而tokenizer输出的input_ids是整型(通常是torch.long类型),拼接时会自动转换类型,但显式指定dtype可以避免潜在的类型问题。
内容的提问来源于stack exchange,提问作者Chi-Yuan Li
相关产品推荐
相关产品推荐

