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

torch.randint()函数在get_batch中无法生成随机值的问题求助

问题分析与解决方案

核心问题:Python列表与PyTorch张量的索引不兼容

你的代码存在两个关键问题:

  1. data 是Python列表,而 ix 是PyTorch张量。用张量元素作为Python列表的索引时,虽然标量张量可能被隐式转换为整数,但这个过程可能出现未预期行为(比如部分环境下转换失败,导致索引默认取0)。
  2. 直接对Python列表切片后用 torch.stack 拼接,不仅效率低,还容易因类型不匹配引发异常。

修复方案

方案1:将data转为PyTorch张量(推荐)

把整个 data 转为PyTorch张量后,利用张量的广播索引高效获取batch,彻底避免Python列表的索引问题:

def get_batch(split):
    # 将Python列表转为PyTorch张量,根据数据类型选择合适的dtype(如torch.long/torch.float)
    data = torch.tensor([包含200万个数字的列表], dtype=torch.long)
    ix = torch.randint(len(data) - block_size, (batch_size,))
    print(ix)
    # 利用广播生成索引矩阵,批量获取x和y
    x = data[ix[:, None] + torch.arange(block_size)]
    y = data[ix[:, None] + torch.arange(block_size) + 1]
    return x, y

方案2:将ix转为Python整数列表

如果需要保留Python列表形式的 data,可以先把 ix 转为整数列表,确保索引是Python原生整数:

def get_batch(split):
    data = [包含200万个数字的列表]
    ix = torch.randint(len(data) - block_size, (batch_size,)).tolist()  # 转为整数列表
    print(ix)
    # 每个切片转为张量后再堆叠
    x = torch.stack([torch.tensor(data[i:i+block_size]) for i in ix])
    y = torch.stack([torch.tensor(data[i+1:i+block_size+1]) for i in ix])
    return x, y

额外验证点

  • 检查 data 列表中对应 ix 索引位置的元素是否确实非0(排除因 data 本身数据问题导致全0张量的可能)。
  • 确保 block_size 和 batch_size 的定义在函数作用域内可访问(比如是全局变量或函数参数)。

内容的提问来源于stack exchange,提问作者Shanu Jha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 16:34:53