关于Numpy实现One-hot编码输出结果的逻辑疑问
搞懂Numpy生成One-Hot编码的索引逻辑
别急,这完全是预期行为!咱们一步步拆解这段代码的工作原理,你马上就能明白1的位置分布规律了。
核心:Numpy的整数数组索引
这段代码里最关键的就是Y_one_hot[Y.flatten(), np.arange(m)] = 1这一行,它用了Numpy的高级整数数组索引,逻辑其实很直白:
Y.flatten():把原始的Y数组转成一维后,每个元素的值就是要设置为1的位置的行索引np.arange(m):生成0到83的数组,每个元素是要设置为1的位置的列索引- 简单说:对于原始
Y中第i个元素(从0开始数),如果这个元素的值是k,那我们就把Y_one_hot里(k, i)这个坐标的位置设为1。
用你的实际例子验证
拿你提供的Y数组和Y_one_hot第一行(索引0)的输出来看:
- 先找
Y里所有等于0的元素的位置:
遍历你的Y数组,等于0的元素出现在这些索引位置:29、34、39、50、55、60 - 对应到
Y_one_hot的第0行,这些列索引(29、34、39、50、55、60)的位置就会被设为1,其他位置保持0。 - 看你给出的第一行输出:
[0 0 0 ... 1(索引29)... 1(索引34)... 1(索引39)... 1(索引50)... 1(索引55)... 1(索引60)... 0],完全和这个规律匹配!
为什么不同行的1数量不一样?
不同行里1的数量,其实就是原始Y数组中对应数字出现的次数:
- 比如第0行的1有6个,正好是
Y中数字0出现的次数 - 再看
Y里数字8出现的次数最多(你数一下,大概有15次左右),对应的Y_one_hot第8行的1数量也会是最多的,你可以自己验证一下。
简化版例子帮你加深理解
如果我们把例子缩小,更容易看清楚逻辑:
Y = np.array([2, 0, 1]) m = 3 vocab_size = 3 Y_one_hot = np.zeros((vocab_size, m)) Y_one_hot[Y.flatten(), np.arange(m)] = 1 print(Y_one_hot)
输出会是:
[[0. 1. 0.] [0. 0. 1.] [1. 0. 0.]]
看,Y[0]=2 → Y_one_hot[2,0] =1;Y[1]=0 → Y_one_hot[0,1]=1;Y[2]=1 → Y_one_hot[1,2]=1,完全对应得上。
所以回到你的代码,生成的Y_one_hot完全符合One-Hot编码的预期,1的分布不是随机的,完全由原始Y数组的元素决定。
内容的提问来源于stack exchange,提问作者D3181
相关产品推荐
相关产品推荐

