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

numpy数组随机拆分后索引数量不符,问题排查与解决

问题原因分析

你的问题出在np.random.randint这个函数上——它生成的随机索引可能包含重复值!

举个简单例子:如果原数组有10行,你用randint选3个索引,结果可能是[2,5,2],这里2出现了两次。当你用列表推导式[i for i in range(...) if i not in ind]时,只会排除一次2,所以剩下的行数是10-2=8,而不是预期的10-3=7。对应到你的场景,ind里有重复的索引,导致实际被排除的唯一索引数量少于5000,所以rest的长度就变成了46928 - (5000 - 重复次数),也就是你看到的42192。

正确的解决方法

这里有两种更可靠的方式来实现无重复的随机行拆分:

方法1:用np.random.permutation打乱索引后拆分

这种方法先把所有索引打乱,再直接切分,逻辑简单直观:

import numpy as np

# 获取所有行的索引并打乱
shuffled_indices = np.random.permutation(input_matrix.shape[0])
# 拆分出5000个测试集索引和剩下的训练集索引
test_indices = shuffled_indices[:5000]
train_indices = shuffled_indices[5000:]

# 提取对应的子数组
test_array = input_matrix[test_indices]
train_array = input_matrix[train_indices]

方法2:用np.random.choice指定不重复采样

直接用choice的replace=False参数,确保生成的索引没有重复:

import numpy as np

# 生成5000个不重复的随机索引
test_indices = np.random.choice(input_matrix.shape[0], size=5000, replace=False)
# 生成剩下的索引:用集合差集的方式更高效
all_indices = np.arange(input_matrix.shape[0])
train_indices = np.setdiff1d(all_indices, test_indices)

# 提取子数组
test_array = input_matrix[test_indices]
train_array = input_matrix[train_indices]

这两种方法都能保证train_array的形状是(41928,28,28),test_array是(5000,28,28),完全符合你的需求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:48:23