You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.01 12:15:03