torch.randint()函数在get_batch中无法生成随机值的问题求助
问题分析与解决方案
核心问题:Python列表与PyTorch张量的索引不兼容
你的代码存在两个关键问题:
data是Python列表,而ix是PyTorch张量。用张量元素作为Python列表的索引时,虽然标量张量可能被隐式转换为整数,但这个过程可能出现未预期行为(比如部分环境下转换失败,导致索引默认取0)。- 直接对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
相关产品推荐
相关产品推荐

