将TensorFlow项目转PyTorch时创建TensorDataset遇'int'不可调用错误
解决PyTorch TensorDataset创建时的TypeError: 'int' object is not callable
这个错误大概率不是输入形状的问题,而是变量名冲突或导入错误导致的,以下是具体排查和修复步骤:
1. 检查TensorDataset的导入与变量重名
- 首先确认代码开头正确导入了TensorDataset:
如果漏了这行,或者之前不小心把from torch.utils.data import TensorDatasetTensorDataset赋值成了整数(比如TensorDataset = 10),就会导致调用时触发int不可调用的错误。 - 可以在创建
train_data前加一行验证:
正常输出应该是print(type(TensorDataset))<class 'type'>,如果输出是<class 'int'>,说明存在变量重名,删掉或重命名那个占用TensorDataset名称的变量即可。
2. 确认输入数据为PyTorch张量
虽然你提到形状正确,但如果输入数据是TensorFlow张量或numpy数组,少数情况下可能引发隐性问题。可以强制转换为PyTorch张量再传入:
import torch train_data = TensorDataset( torch.as_tensor(self.x_train[idx_train]), torch.as_tensor(self.covariates_train[idx_train]), torch.as_tensor(self.y_train[idx_train]) )
torch.as_tensor会自动处理numpy数组或TensorFlow张量的转换,同时尽量避免不必要的数据拷贝。
3. 检查索引变量idx_train
确保idx_train是数组、列表或PyTorch张量类型,而非单个整数。如果idx_train是单个整数,切片后会得到张量的单个元素而非批量张量,但你的形状输出显示是(180, ...),所以这个可能性较低,但可以快速验证:
print(type(idx_train), idx_train.shape if hasattr(idx_train, 'shape') else len(idx_train))
确认输出的长度/形状是180即可。
内容的提问来源于stack exchange,提问作者johnnyb
相关产品推荐
相关产品推荐

