基于PyTorch/Lightning的电影推荐系统索引越界问题求助
问题定位与修复方案
错误核心原因
这个IndexError是0-based索引越界导致的:张量维度0的长度为6040,合法索引范围是0~6039,但代码中调用了索引6040。结合MovieLens 1M数据集的用户总数正好是6040(原始ID为1~6040),大概率是ID映射未做0-base转换,或序列填充值设置错误。
针对性修复方案
1. 用户/物品ID的0-base映射修正
MovieLens 1M的用户、物品ID均从1开始编号,直接用原始ID作为嵌入层索引的话,当ID=6040时,会超出num_embeddings=6040的嵌入层索引范围(最大合法索引为6039)。
修复代码:
在数据预处理阶段统一对ID做减1处理:
# 假设加载后的数据存储在DataFrame中 data['user_id'] = data['user_id'] - 1 data['movie_id'] = data['movie_id'] - 1
同时确保嵌入层的num_embeddings参数与修正后的ID总数一致(比如用户嵌入层设为num_embeddings=6040,此时0~6039正好覆盖所有用户)。
2. 序列填充与掩码逻辑修正
如果构建用户交互序列时,填充值设为了用户总数6040,LSTM处理序列时会将该填充值当作索引传入嵌入层,触发越界。
修复代码:
将填充值改为合法范围内的数值(比如0),并在嵌入层指定padding_idx:
# 构建序列时用0作为填充值 sequences = torch.nn.utils.rnn.pad_sequence(sequence_list, batch_first=True, padding_value=0) # 嵌入层设置padding_idx,自动忽略填充位置的计算 user_embedding = torch.nn.Embedding(num_embeddings=6040, embedding_dim=64, padding_idx=0)
同时修正掩码张量的生成逻辑,确保标记出有效序列位置:
mask = (sequences != 0).float() # 0为填充值,掩码1表示有效位置
3. 循环索引逻辑修正
如果手动遍历序列时,循环计数器超出了0-base范围(比如从1开始遍历到序列长度),也会触发越界。
修复代码:
确保循环使用0-base索引:
# 错误示例(会取到等于序列长度的索引) for i in range(1, len(sequences)): token = sequences[i] # 正确写法 for i in range(len(sequences)): token = sequences[i]
快速检查清单
- 确认用户/物品ID的取值范围为
0~N-1(N为对应实体总数) - 检查嵌入层
num_embeddings是否与实体总数一致,padding_idx设置正确 - 排查序列数据中的填充值是否在合法索引范围内
- 验证所有索引操作的变量(循环计数器、掩码索引等)未超出0-base范围
内容的提问来源于stack exchange,提问作者Vaggelis Papadiotis
相关产品推荐
相关产品推荐

