如何在PyTorch正弦预测任务中正确实现Dataset的__getitem__方法
现有实现的核心问题
- 索引校验逻辑错误:合法索引范围应为
0 <= idx < self.size,你当前的判断条件仅在idx > self.size时触发报错,会漏过idx == self.size、负索引等非法情况 - 返回值类型不符合要求:方法签名标注返回
Tensor类型,但实际返回Python原生的int和float类型,无法直接输入到PyTorch模型中参与训练 - 输入特征采样不合理:直接用整数索引作为x输入,采样粒度过粗,且取值范围过大,不利于模型拟合正弦函数的周期性规律
正确实现代码
import torch import math from torch.utils.data import Dataset class SinusDataset(Dataset): def __init__(self, size: int = 1000, x_range: tuple = (0, 2 * math.pi)): self.size = size self.x_min, self.x_max = x_range # 小数据集推荐提前生成所有数据,避免每次getitem重复计算 self.x = torch.linspace(self.x_min, self.x_max, steps=self.size, dtype=torch.float32) self.y = torch.sin(self.x) def __len__(self): return self.size def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]: # PyTorch框架会自动处理负索引、批量索引合法性,无需额外手动校验 return self.x[idx], self.y[idx]
改动说明
- 新增
x_range参数控制输入x的取值范围,默认覆盖0到2π的完整正弦周期,更适合模型拟合规律 - 提前在
__init__中用torch.linspace生成均匀分布的浮点型x值,采样更均匀,且直接生成张量类型,避免后续重复类型转换 - 移除冗余的索引校验逻辑,保留PyTorch原生支持的负索引、批量索引特性
- 严格保证返回值为
torch.Tensor类型,符合方法签名要求,可直接输入模型参与训练
如果需要随机采样的数据增强效果,也可以改成动态生成样本的版本:
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]: x = torch.rand(1, dtype=torch.float32) * (self.x_max - self.x_min) + self.x_min y = torch.sin(x) return x.squeeze(), y.squeeze()
内容的提问来源于stack exchange,提问作者MatthiasDunkel
相关产品推荐
相关产品推荐

