使用numpy.sum()与numpy.where()条件求和计算类别频率
实现方案
原代码存在两处基础错误需要先修正:
- 没有将嵌套列表格式的
labels转为numpy数组,无法直接调用numpy的统计函数 - 总行数
N的取值逻辑错误,原写法取了第一行标签列表,并非样本总数量
核心实现逻辑
要求的正负类频率按列统计,可通过numpy.sum()配合轴参数实现统计逻辑,numpy.where()用于快速标记负类位置:
- 统计每列正类(值为1)数量:调用
np.sum()时指定axis=0,即可按列累加所有元素值,得到的结果就是每个类别(每列)的正样本总数 - 统计每列负类(值为0)数量:用
np.where(labels_np == 0, 1, 0)将所有值为0的位置替换为1、值为1的位置替换为0,再按axis=0求和,即可得到每个类别的负样本总数 - 两类频率分别用对应列的样本总数除以总行数N即可,输出自动为浮点型numpy数组;也可以直接用
1 - positive_frequencies计算负类频率,结果完全一致。
修正后完整代码
import numpy as np def compute_class_freqs(): """ Compute positive and negative frequences for each class. Returns: positive_frequencies (np.array): array of positive frequences for each class, size (num_classes) negative_frequencies (np.array): array of negative frequences for each class, size (num_classes) """ ### START CODE HERE (REPLACE INSTANCES OF 'None' with your code) ### labels = [[0,1,0],[1,1,1],[0,1,1]] # 转换为numpy数组 labels_np = np.array(labels) print(labels) # 修正总行数计算逻辑 N = labels_np.shape[0] # 按列求和统计每列1的数量,除以总行数得到正类频率 positive_frequencies = np.sum(labels_np, axis=0) / N # 用np.where标记0值位置,按列求和除以总行数得到负类频率 negative_frequencies = np.sum(np.where(labels_np == 0, 1, 0), axis=0) / N ### END CODE HERE ### return positive_frequencies, negative_frequencies
运行结果验证
传入示例标签时,返回结果为:
positive_frequencies = array([0.33333333, 1. , 0.66666667])negative_frequencies = array([0.66666667, 0. , 0.33333333])
完全符合按列统计频率的要求,返回值为浮点类型numpy数组。
内容的提问来源于stack exchange,提问作者python-dude153
相关产品推荐
相关产品推荐

