NumPy核心归约函数:从max、sum到argmax的向量化性能优化与实战避坑
2026/8/15 6:33:51 网站建设 项目流程

1. 从一次数据清洗的“翻车”说起:为什么你需要搞懂这七个函数?

前几天帮一个做量化分析的朋友处理一组股票分钟级数据,他需要快速找出每只股票在当天交易时段内的最高价、最低价,并定位这些极值出现的时间点,最后还要计算每只股票的总成交额。听起来是个简单的任务,对吧?我一开始也是这么想的,顺手就写了个循环,对每只股票的数据数组用max()min()去遍历。结果,面对几百只股票、上万条分钟数据,脚本跑了快十分钟还没结束,朋友在旁边等得直皱眉。

那一刻我意识到,问题不在于算法复杂,而在于工具没选对。在Python的数据科学栈里,尤其是NumPy的世界里,“写循环”往往是性能瓶颈的第一个信号。我立刻停掉脚本,把循环全部替换成了NumPy的向量化操作:np.max(axis=0),np.argmin(axis=1),再加上一个np.sum()。同样的计算,眨眼间就完成了,耗时不到一秒。朋友看着刷新出来的结果,说了句:“这才是专业工具该有的样子。”

这个故事引出了今天要彻底讲清楚的七个NumPy函数:np.max(),np.argmax(),np.maximum(),np.min(),np.argmin(),np.minimum(),np.sum()。它们被并称为NumPy的“归约函数”或“元素级比较函数”,是处理数值型数组的基石。很多初学者觉得它们简单,看一眼文档就过了,但实际用起来,尤其是在多维数组、轴(axis)操作、广播(broadcasting)机制以及处理缺失值(NaN)时,到处都是坑。

如果你满足于用Python原生的max()sum(),那你可能永远无法体会NumPy向量化计算带来的百倍甚至千倍的性能提升。更重要的是,不理解这些函数在轴向上的细微差别,在处理图像数据(高度、宽度、通道)、时间序列(样本、时间步、特征)或者任何多维张量时,你很容易得到形状错误或者完全不符合预期的结果。本文将不仅告诉你每个函数怎么用,更会深入它们的设计逻辑、性能差异和使用禁忌,让你真正从“会用”到“精通”。

2. 核心三兄弟:求最值与定位(np.max/min(),np.argmax/argmin()

这是最常被用到的一组函数,它们的核心任务是从一个数组里找出极值。但“找出来”这个动作,在NumPy的语境下,因为有了“轴”的概念,变得需要仔细斟酌。

2.1np.max()np.min():不仅仅是找到那个数

np.max(a, axis=None, keepdims=False)np.min()的函数签名完全一致。它们的核心作用是沿指定轴计算最大值/最小值,并返回一个降维后的新数组。理解“沿指定轴”和“降维”是关键。

轴(axis)的直观理解: 对于一个二维数组(矩阵),我们可以把它想象成一个Excel表格。

  • axis=0:沿着行的方向(垂直向下),对每一进行操作。可以理解为“跨行求列统计”。
  • axis=1:沿着列的方向(水平向右),对每一进行操作。可以理解为“跨列求行统计”。
import numpy as np # 创建一个3x4的二维数组 arr = np.array([[1, 2, 8, 4], [9, 5, 3, 7], [2, 6, 1, 5]]) print("原始数组:\n", arr) print("形状:", arr.shape) # (3, 4) # 不指定axis,返回全局最大值/最小值 print("全局最大值 np.max(arr):", np.max(arr)) # 输出:9 print("全局最小值 np.min(arr):", np.min(arr)) # 输出:1 # axis=0: 沿着第0轴(行方向)压缩,对每一列求最大值 # 结果数组的形状是 (4,),因为3行被压缩掉了。 col_max = np.max(arr, axis=0) print("\n每列的最大值 np.max(arr, axis=0):", col_max) # 输出:[9 6 8 7] # 计算过程:第一列 max(1,9,2)=9;第二列 max(2,5,6)=6;第三列 max(8,3,1)=8;第四列 max(4,7,5)=7 # axis=1: 沿着第1轴(列方向)压缩,对每一行求最大值 # 结果数组的形状是 (3,),因为4列被压缩掉了。 row_max = np.max(arr, axis=1) print("每行的最大值 np.max(arr, axis=1):", row_max) # 输出:[8 9 6] # 计算过程:第一行 max(1,2,8,4)=8;第二行 max(9,5,3,7)=9;第三行 max(2,6,1,5)=6

keepdims参数:维持维度信息的神器这是一个极易被忽略但极其重要的参数。默认keepdims=False,意味着执行归约操作后,被操作的轴会消失(维度从n降到1再被移除)。但有时我们需要保持数组的维度,以便进行后续的广播运算。

# 假设我们要计算每个数据点与其所在行最大值的差值 arr = np.array([[1, 2, 8], [9, 5, 3]]) # 错误做法:直接相减,形状不匹配会报错 # diff = arr - np.max(arr, axis=1) # 错误!arr形状(2,3),max结果形状(2,) # 正确做法:使用 keepdims=True row_max_keep = np.max(arr, axis=1, keepdims=True) print("row_max (keepdims=False):", np.max(arr, axis=1)) # 形状 (2,) print("row_max (keepdims=True):\n", row_max_keep) # 形状 (2, 1) print("row_max_keep的形状:", row_max_keep.shape) # 现在可以广播相减了 diff = arr - row_max_keep print("每个元素减去其所在行的最大值:\n", diff) # 输出: # [[-7 -6 0] # [ 0 -4 -6]]

注意:在深度学习框架(如PyTorch, TensorFlow)中,类似的归约操作也几乎都提供了keepdimkeepdims参数,其设计思路一脉相承。养成使用它的习惯,能避免很多形状不匹配的错误。

2.2np.argmax()np.argmin():找到“位置”比找到“值”更重要

如果说max/min告诉你考最高分是多少,那么argmax/argmin就告诉你考最高分的是哪个学生。它们返回的是沿指定轴的最大值/最小值的索引

arr = np.array([[1, 2, 8, 4], [9, 5, 3, 7], [2, 6, 1, 5]]) # 全局最大值的索引(将数组展平后的一维索引) print("全局最大值索引 np.argmax(arr):", np.argmax(arr)) # 输出:4 # 因为展平后数组为 [1,2,8,4,9,5,3,7,2,6,1,5],最大值9在索引4的位置 # 每列最大值的索引(行索引) col_argmax = np.argmax(arr, axis=0) print("每列最大值所在的**行**索引 np.argmax(arr, axis=0):", col_argmax) # 输出:[1 2 0 1] # 解释:第0列最大值9在第1行;第1列最大值6在第2行;第2列最大值8在第0行;第3列最大值7在第1行。 # 每行最大值的索引(列索引) row_argmax = np.argmax(arr, axis=1) print("每行最大值所在的**列**索引 np.argmax(arr, axis=1):", row_argmax) # 输出:[2 0 1] # 解释:第0行最大值8在第2列;第1行最大值9在第0列;第2行最大值6在第1列。

一个经典应用场景:从One-hot编码中获取类别标签在分类任务中,神经网络的输出通常是每个类别的概率分布(如softmax后的结果),形状为(batch_size, num_classes)。我们需要得到每个样本预测的类别,即概率最大的那个索引。

# 模拟一个批量大小为3,共5个类别的网络输出概率 batch_probabilities = np.array([[0.1, 0.2, 0.05, 0.6, 0.05], [0.7, 0.1, 0.1, 0.05, 0.05], [0.05, 0.8, 0.1, 0.03, 0.02]]) print("网络输出概率(形状: 3x5):\n", batch_probabilities) # 我们需要的是每行(每个样本)最大概率的列索引(类别编号) predicted_classes = np.argmax(batch_probabilities, axis=1) print("预测的类别索引:", predicted_classes) # 输出:[3 0 1] # 样本0预测为第3类,样本1预测为第0类,样本2预测为第1类。

实操心得:argmax返回的是第一个遇到的最值索引当数组中存在多个相同的最大值或最小值时,argmaxargmin默认返回第一个遇到的索引。这是一个重要的边界条件。

arr = np.array([2, 5, 5, 1, 5]) print("数组:", arr) print("np.argmax(arr):", np.argmax(arr)) # 输出:1,而不是2或4 print("np.argmin(arr):", np.argmin(arr)) # 输出:3

如果你的业务逻辑要求获取所有最值的位置,你就不能直接用argmax,而需要结合布尔索引:np.where(arr == np.max(arr))

3. 元素级的较量:np.maximum()np.minimum()

这组函数与前两组有本质区别。np.maximum(x1, x2, /, out=None, *, where=True, ...)不是从一个数组里找最大值,而是对两个数组进行逐元素比较,取每个位置上较大的那个值np.minimum同理。它执行的是元素级(element-wise)操作,不进行归约,输出数组的形状由输入数组的广播规则决定。

3.1 基础用法与广播机制

# 最基本的逐元素比较 a = np.array([1, 5, 3]) b = np.array([2, 3, 6]) result = np.maximum(a, b) print("a:", a) print("b:", b) print("np.maximum(a, b):", result) # 输出:[2 5 6] # 计算过程:max(1,2)=2; max(5,3)=5; max(3,6)=6 # 与单个数值比较(广播的典型应用) arr = np.array([[-1, 5, -3], [4, -2, 0]]) clipped = np.maximum(arr, 0) # 将数组中所有小于0的值截断(clip)为0 print("原始数组:\n", arr) print("经过 np.maximum(arr, 0) ReLU激活后:\n", clipped) # 输出: # [[0 5 0] # [4 0 0]] # 这其实就是深度学习ReLU激活函数的朴素实现。

广播(Broadcasting)是理解maximum/minimum威力的关键。当两个数组形状不同时,NumPy会尝试通过广播机制将它们扩展为兼容的形状,然后再进行逐元素操作。

# 示例:一个2x3的矩阵与一个长度为3的行向量比较 matrix = np.array([[10, 20, 30], [40, 50, 60]]) row_vector = np.array([25, 15, 35]) # 形状 (3,) # 广播发生:row_vector 被“复制”到与 matrix 行数匹配 # 相当于变成了 [[25,15,35], [25,15,35]],再与 matrix 逐元素比较 result = np.maximum(matrix, row_vector) print("矩阵:\n", matrix) print("行向量:", row_vector) print("逐元素取大值:\n", result) # 输出: # [[25 20 35] # [40 50 60]] # 计算过程: # 第一行:max(10,25)=25; max(20,15)=20; max(30,35)=35 # 第二行:max(40,25)=40; max(50,15)=50; max(60,35)=60

3.2 高级应用:实现自定义的上下限截断(Clipping)

虽然NumPy提供了专门的np.clip()函数,但用maximumminimum组合可以实现同样的功能,并且逻辑更清晰。

def my_clip(arr, min_val, max_val): """手动实现数组范围截断,将元素限制在[min_val, max_val]区间内""" # 先确保不低于下限 temp = np.maximum(arr, min_val) # 再确保不高于上限 result = np.minimum(temp, max_val) return result data = np.array([1, 5, 10, 15, 20]) clipped_data = my_clip(data, 5, 15) print("原始数据:", data) print("截断到 [5, 15] 后:", clipped_data) # 输出:[ 5 5 10 15 15]

一个图像处理的真实案例:融合两张图片的亮部假设我们有两张同一场景不同曝光的照片,想合成一张高动态范围(HDR)效果的图片,一个简单策略是取每张图片在每个像素点上较亮的值。

# 模拟两张灰度图片(像素值范围0-255) # 假设image1整体偏暗但高光细节好,image2整体偏亮但暗部细节好 height, width = 100, 100 image1 = np.random.randint(0, 180, (height, width)).astype(np.float32) # 偏暗 image2 = np.random.randint(80, 255, (height, width)).astype(np.float32) # 偏亮 # 简单的“取亮”融合 fused_image = np.maximum(image1, image2) # 此时fused_image在每个像素位置都保留了两张图中更亮的那一个值 # 这只是一个非常初级的融合策略,真实的HDR算法要复杂得多。

4. 基石中的基石:np.sum()的深度剖析

np.sum()可能是NumPy中使用频率最高的函数之一。它的基本功能是求和,但结合axis,keepdims,dtype等参数,能演变出无数种用法。

4.1 轴向求和与维度理解

arr = np.array([[1, 2, 3], [4, 5, 6]]) print("二维数组:\n", arr) print("形状:", arr.shape) # (2, 3) # 全局求和 print("np.sum(arr):", np.sum(arr)) # 1+2+...+6 = 21 # 沿axis=0求和(跨行,对列求和) sum_axis0 = np.sum(arr, axis=0) print("沿axis=0求和 (跨行,列和):", sum_axis0) # [5 7 9] print("结果形状:", sum_axis0.shape) # (3,) # 计算:第一列 1+4=5;第二列 2+5=7;第三列 3+6=9 # 沿axis=1求和(跨列,对行求和) sum_axis1 = np.sum(arr, axis=1) print("沿axis=1求和 (跨列,行和):", sum_axis1) # [6 15] print("结果形状:", sum_axis1.shape) # (2,) # 计算:第一行 1+2+3=6;第二行 4+5+6=15 # 同时沿多个轴求和 sum_axis01 = np.sum(arr, axis=(0, 1)) # 等价于全局求和 print("同时沿axis=0和1求和:", sum_axis01) # 21

4.2dtype参数:防止溢出与精度控制

这是np.sum()一个至关重要但常被忽视的参数。NumPy数组有固定的数据类型(dtype),如int32,float64。在进行求和时,如果中间结果超过了该数据类型能表示的范围,就会发生溢出(对于整数)或精度损失(对于浮点数)。

# 整数溢出案例 large_int_arr = np.ones(1000000, dtype=np.int16) * 30000 # 每个元素是30000,100万个这样的数总和是300亿,远超int16的范围(-32768 ~ 32767) try: wrong_sum = np.sum(large_int_arr) # 使用默认的输入数组dtype(int16)进行计算 print("使用int16 dtype的求和结果(溢出):", wrong_sum) # 会得到一个错误的值 except Exception as e: print("可能出错:", e) # 正确做法:指定一个足够大的输出数据类型 correct_sum = np.sum(large_int_arr, dtype=np.int64) # 或者 np.float64 print("指定dtype=np.int64的求和结果:", correct_sum) # 浮点数精度案例 float_arr = np.full(1000000, 0.1, dtype=np.float32) # 单精度浮点数 sum_float32 = np.sum(float_arr) sum_float64 = np.sum(float_arr, dtype=np.float64) # 在累加过程中使用双精度 print("使用float32累加:", sum_float32) print("使用float64累加:", sum_float64) print("理论值应为:", 0.1 * 1000000) # 你会发现 sum_float64 的结果更接近理论值,因为双精度减少了累加过程中的舍入误差。

重要经验:在处理大规模数据求和时,尤其是整数数组,养成习惯指定dtype=np.float64或足够大的整数类型,可以避免许多隐蔽的数值计算错误。对于金融、科学计算等领域,这是必须遵守的准则。

4.3where参数:条件求和的利器

从NumPy 1.20版本开始,np.sum()增加了where参数,允许你只对满足条件的元素进行求和,这比先进行布尔索引再求和更高效、更优雅。

arr = np.array([1, -2, 3, -4, 5, 6]) # 传统方法:布尔索引 positive_sum_old = np.sum(arr[arr > 0]) print("传统方法(正数和):", positive_sum_old) # 1+3+5+6=15 # 新方法:使用 where 参数 positive_sum_new = np.sum(arr, where=arr > 0) print("使用where参数(正数和):", positive_sum_new) # 15 # 它更强大的地方在于处理多维数组和复杂条件 matrix = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) # 求矩阵中所有大于3且小于8的元素之和 conditional_sum = np.sum(matrix, where=(matrix > 3) & (matrix < 8)) print("矩阵中大于3且小于8的元素和:", conditional_sum) # 4+5+6+7=22

使用where参数,NumPy在内部进行优化,避免了创建中间布尔数组的开销,对于大型数组性能提升明显。

5. 性能对比与实战避坑指南

知道怎么用只是第一步,知道什么时候用哪个,以及如何避免踩坑,才是进阶的关键。

5.1 向量化操作 vs Python循环:性能天壤之别

让我们用数据量化一下性能差异。我们将对一个包含1000万个随机数的数组求最大值和总和。

import numpy as np import time # 生成测试数据 data = np.random.randn(10_000_000) # 1000万个随机数 print("数据量:1000万浮点数") # 方法1: 使用Python内置的max和sum(在列表上) py_list = data.tolist() start = time.time() py_max = max(py_list) py_sum = sum(py_list) py_time = time.time() - start print(f"Python内置函数耗时: {py_time:.4f} 秒") # 方法2: 使用NumPy的向量化函数 start = time.time() np_max = np.max(data) np_sum = np.sum(data) np_time = time.time() - start print(f"NumPy向量化函数耗时: {np_time:.4f} 秒") print(f"NumPy比Python快 {py_time / np_time:.1f} 倍") print("结果一致性检查 - 最大值:", np.allclose(py_max, np_max)) print("结果一致性检查 - 总和:", np.allclose(py_sum, np_sum))

在我的测试环境中,NumPy版本通常比纯Python循环快50到100倍以上。这个差距随着数据量的增大而急剧扩大。核心原因在于:NumPy的核心算法是用C语言实现的,并且在内存中连续存储数据,能够充分利用现代CPU的SIMD(单指令多数据流)指令集进行并行计算。而Python的循环是解释执行的,每个操作都有巨大的开销。

避坑指南1:对NumPy数组,永远不要用Python的max()/min()/sum()直接对NumPy数组调用Python内置函数,首先会导致数组被隐式转换为Python列表(如果可能),或者产生非预期行为,性能损失巨大。务必使用np.max(),np.sum()等。

5.2 处理NaN值:沉默的“杀手”

NaN(Not a Number)是浮点数计算中一个特殊值,表示未定义或不可表示的结果。NumPy的归约函数默认行为是遇到NaN时,整个计算结果也会是NaN。

arr_with_nan = np.array([1.0, 2.0, np.nan, 4.0, 5.0]) print("包含NaN的数组:", arr_with_nan) print("np.sum(arr_with_nan):", np.sum(arr_with_nan)) # 输出:nan print("np.max(arr_with_nan):", np.max(arr_with_nan)) # 输出:nan print("np.min(arr_with_nan):", np.min(arr_with_nan)) # 输出:nan

这常常导致数据处理流程意外中断。为了解决这个问题,NumPy提供了以nan开头的特殊函数:

print("np.nansum(arr_with_nan):", np.nansum(arr_with_nan)) # 忽略NaN求和:1+2+4+5=12.0 print("np.nanmax(arr_with_nan):", np.nanmax(arr_with_nan)) # 忽略NaN求最大值:5.0 print("np.nanmin(arr_with_nan):", np.nanmin(arr_with_nan)) # 忽略NaN求最小值:1.0 print("np.nanargmax(arr_with_nan):", np.nanargmax(arr_with_nan)) # 忽略NaN的最大值索引:4

避坑指南2:处理真实数据(尤其是从文件读取的)时,先检查是否存在NaN,或直接使用np.nan*系列函数。可以使用np.isnan(arr).any()来检查。

5.3 轴(axis)参数的“负索引”与高维数组

对于二维数组,axis=0axis=1还算直观。但对于三维及以上的高维数组(例如图像数据(高度,宽度,通道)或批量时间序列(批量大小,时间步长,特征数)),轴的方向就容易混乱。这时可以使用负索引来从后往前指定轴。

# 创建一个3维数组,模拟一个批量大小为2,3x4像素,RGB三通道的图片数据 # 形状:(batch_size, height, width, channels) -> (2, 3, 4, 3) image_batch = np.random.randint(0, 256, size=(2, 3, 4, 3), dtype=np.uint8) print("图像批次形状:", image_batch.shape) # (2, 3, 4, 3) # 需求1:计算每张图片所有像素在R通道上的平均值 # 我们需要对每张图片的高度(height)和宽度(width)求和,即压缩 axis=1 和 axis=2 # 方法A:分别指定两个轴 mean_r_per_image_a = np.mean(image_batch[..., 0], axis=(1, 2)) # ... 是省略号,表示所有前面的维度 print("每张图片R通道均值(方法A):", mean_r_per_image_a.shape) # (2,) # 方法B:使用负索引,从最后一个维度往前数。通道是axis=-1,宽度是axis=-2,高度是axis=-3。 # 对高度和宽度求和,就是 axis=(-3, -2) mean_r_per_image_b = np.mean(image_batch[..., 0], axis=(-3, -2)) print("每张图片R通道均值(方法B):", mean_r_per_image_b.shape) # (2,) # 需求2:计算所有图片、所有像素的每个通道的总和 # 即压缩前三个轴 (batch, height, width),保留通道轴 (axis=-1) sum_per_channel = np.sum(image_batch, axis=(0, 1, 2)) # 等价于 axis=(0, 1, 2) print("所有图片各通道像素总和:", sum_per_channel.shape) # (3,)

避坑指南3:在处理高维数组时,画一个维度的草图,或者使用负索引来指代“最后几个轴”,可以让代码更清晰、更不容易出错。

5.4initial参数:求和的起点

np.sum()还有一个initial参数,用于指定求和的初始值。这在某些场景下非常有用,比如处理空数组。

empty_arr = np.array([]) try: print(np.sum(empty_arr)) # 对空数组求和,默认返回0.0(对于浮点类型) except Exception as e: print(e) # 但如果你希望空数组的和是一个特定的值,比如在计算连乘的对数似然时,空集的和应该是0(加性单位元) # 或者,你想确保结果从一个基数开始累加 arr = np.array([10, 20, 30]) print("从100开始累加:", np.sum(arr, initial=100)) # 输出:100 + 10+20+30 = 160

这个参数在函数式编程或特定数学场景下能提供更精确的控制。

6. 综合实战:用这七个函数解决一个真实问题

假设你是一家电商公司的数据分析师,你有一份销售数据,是一个三维数组sales_data,形状为(产品类别数, 月份数, 地区数),例如(5, 12, 10)表示5个品类、12个月、10个地区的销售额。

你的任务是:

  1. 找出全年销售额最高的单个“品类-月份-地区”组合及其销售额和具体位置。
  2. 找出每个品类,在哪个地区全年(12个月加总)销售额最低。
  3. 计算每个地区,所有品类在各个月份的销售额总和。
  4. 为了制作热力图,需要将每个“品类-月份”组合的数据,与全年的月平均销售额进行比较,生成一个突出显示高于平均月份的数据矩阵。

让我们一步步用NumPy函数来解决。

import numpy as np # 1. 生成模拟数据 np.random.seed(42) # 确保结果可复现 n_categories = 5 n_months = 12 n_regions = 10 sales_data = np.random.randint(1000, 50000, size=(n_categories, n_months, n_regions)).astype(np.float32) # 随机插入一些NaN,模拟数据缺失 nan_mask = np.random.rand(*sales_data.shape) < 0.01 # 大约1%的数据为NaN sales_data[nan_mask] = np.nan print(f"销售数据形状: {sales_data.shape}") # (5, 12, 10) # 任务1: 找出全年销售额最高的单个组合(忽略NaN) # 使用 nanmax 避免NaN影响,再使用 unravel_index 将一维索引转换为多维索引 max_value = np.nanmax(sales_data) print(f"\n1. 全年最高单笔销售额: {max_value:.2f}") # 找到这个最大值在所有维度中的位置 flat_index = np.nanargmax(sales_data) # 展平后的一维索引 cat_idx, month_idx, region_idx = np.unravel_index(flat_index, sales_data.shape) print(f" 位置: 品类[{cat_idx}], 月份[{month_idx+1}], 地区[{region_idx}]") # 任务2: 找出每个品类,在哪个地区全年销售额最低 # 步骤:a) 对月份轴(axis=1)求和,得到每个品类-地区的全年总额,形状(5,10) # b) 对每个品类(axis=0的每个元素),找总额最小的地区(axis=1的方向) yearly_sales_per_cat_region = np.nansum(sales_data, axis=1) # 形状 (5, 10) print(f"\n2. 每个品类-地区的全年销售额总和形状: {yearly_sales_per_cat_region.shape}") # 对每个品类,找销售额最低的地区索引 worst_region_per_category = np.nanargmin(yearly_sales_per_cat_region, axis=1) # 形状 (5,) print(f" 每个品类销售额最低的地区索引: {worst_region_per_category}") # 同时可以拿到最低的销售额值 worst_sales_value = np.nanmin(yearly_sales_per_cat_region, axis=1) for cat in range(n_categories): print(f" 品类{cat}: 最差地区[{worst_region_per_category[cat]}], 销售额{worst_sales_value[cat]:.2f}") # 任务3: 计算每个地区,所有品类在各个月份的销售额总和 # 需要对品类轴(axis=0)和月份轴(axis=1)求和,保留地区轴(axis=2) # 使用 keepdims=True 方便后续如果需要广播 total_sales_per_region = np.nansum(sales_data, axis=(0, 1), keepdims=True) # 形状 (1, 1, 10) # 为了打印好看,去掉多余的维度 total_sales_per_region = total_sales_per_region.squeeze() print(f"\n3. 每个地区的总销售额:") for region in range(n_regions): print(f" 地区{region}: {total_sales_per_region[region]:.2f}") # 任务4: 生成高于月平均销售额的突出显示矩阵 # 步骤:a) 计算每个月份(跨品类和地区)的平均销售额,形状 (12,) # b) 将原始数据与每个月的平均值进行比较 # 注意:比较时需要考虑NaN,我们用 nanmean monthly_avg_sales = np.nanmean(sales_data, axis=(0, 2)) # 压缩品类和地区轴,形状 (12,) print(f"\n4. 各月份平均销售额: {monthly_avg_sales}") # 为了比较,需要将 monthly_avg_sales 广播到和 sales_data 匹配的形状 (5,12,10) # 首先用 keepdims=True 得到形状 (1,12,1),然后利用广播 monthly_avg_sales_expanded = np.nanmean(sales_data, axis=(0, 2), keepdims=True) # 形状 (1,12,1) # 逐元素比较,生成布尔矩阵 highlight_mask = sales_data > monthly_avg_sales_expanded # 这个 highlight_mask 是一个布尔数组,True表示该位置销售额高于其所在月份的平均值 print(f" 突出显示矩阵(True表示高于当月平均)中True的比例: {np.mean(highlight_mask):.2%}") # 我们可以进一步,将高于平均的值保留,低于平均的设为0,得到一个“亮点”矩阵 highlighted_sales = np.where(highlight_mask, sales_data, 0) print(f" ‘亮点’矩阵(仅保留高于月平均的值)的总和: {np.nansum(highlighted_sales):.2f}")

通过这个综合案例,你可以看到,仅仅七个基础函数,通过灵活组合和轴向操作,就能高效解决一个看似复杂的多维度数据分析问题。关键在于清晰地定义“沿着哪个轴进行压缩”,以及熟练运用广播机制来对齐数据形状。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询