在NumPy数组中直接替换字典值实现独热编码为何出现内容截断?
我有一个包含单词的数组,想要执行独热编码。输入示例为:AI DSA DSA AI ML ML AI DS DS AI C AI ML ML C。
我的代码如下:
import numpy as np def apply_one_hot_encoding(X): dic = {} k = sorted(list(set(X))) for i in range(len(k)): arr = ['0' for i in range(len(k))] arr[i] = '1' dic[k[i]] = ''.join(arr) for i in range(len(X)): t = dic[X[i]] X[i] = t return X if __name__ == "__main__": X = np.array(list(input().split())) one_hot_encoded_array = apply_one_hot_encoding(X) for i in one_hot_encoded_array: print(*i)
我预期的输出是每行对应5位编码(比如1 0 0 0 0),但实际得到的输出每行仅3位(如1 0 0)。如果把t值追加到另一个列表并返回该列表,就能得到正确结果。请问为何直接替换时赋值会被截断为仅3个字符?
问题根源:NumPy数组的固定数据类型限制
你遇到的问题核心在于NumPy数组的 dtype 特性:
当你用np.array(list(input().split()))创建数组时,NumPy会自动推断元素的数据类型。你的输入里最长的单词是DSA(3个字符),所以NumPy会把数组的 dtype 设置为<U3(表示最多容纳3个字符的Unicode字符串)。
当你把长度为5的独热编码字符串(比如"10000")赋值给这个数组的元素时,NumPy会自动截断字符串到3个字符,只保留前3位,这就是你看到输出只有1 0 0的原因。
而如果用普通Python列表处理,列表对元素的长度没有限制,所以追加到新列表就能得到完整的编码。
解决方案
这里有几个可行的解决办法:
1. 改用普通Python列表处理
直接去掉NumPy数组的包装,用原生列表存储输入:
if __name__ == "__main__": X = list(input().split()) # 这里不再转成numpy数组 one_hot_encoded_array = apply_one_hot_encoding(X) for i in one_hot_encoded_array: print(*i)
这样修改后,列表元素可以是任意长度的字符串,不会出现截断问题。
2. 显式指定NumPy数组的字符串长度
如果必须使用NumPy数组,创建时指定足够大的字符串 dtype,确保能容纳独热编码的长度:
if __name__ == "__main__": X = np.array(list(input().split()), dtype='<U5') # 指定最多5个字符 one_hot_encoded_array = apply_one_hot_encoding(X) for i in one_hot_encoded_array: print(*i)
这里<U5表示允许存储最多5个字符的Unicode字符串,刚好匹配你需要的5位编码。如果以后编码长度变化,也可以调整这个数值(比如<U10)。
3. 返回新列表(你已经验证有效的方法)
在函数内部创建一个新列表,将编码后的字符串添加进去,而不是修改原数组:
def apply_one_hot_encoding(X): dic = {} k = sorted(list(set(X))) for i in range(len(k)): arr = ['0' for i in range(len(k))] arr[i] = '1' dic[k[i]] = ''.join(arr) # 创建新列表存储结果 result = [] for item in X: result.append(dic[item]) return result
这种方法避开了修改原NumPy数组的限制,同样能得到正确的结果。
内容的提问来源于stack exchange,提问作者Akula Sai Sridhar

