如何基于条件分组拆分Numpy数组?优化重复条件代码
嘿,这个问题我之前也碰到过,重复写条件确实挺冗余的!给你分享个更简洁、更Pythonic的方案,只需要定义一次条件就能搞定,而且用纯Numpy操作效率也很高:
方法一:Numpy原生操作(推荐大数据组)
import numpy as np a = np.array([1, 2, 3, 4, 10, 9, 8, 7, 7, 7, 8, 9, 10]) # 只定义一次条件,避免重复计算 condition = np.logical_and(a > 2, a < 8) # 给条件数组前后补False,处理首尾符合条件的边界情况 padded_condition = np.concatenate([[False], condition, [False]]) # 找到条件发生变化的索引位置(False→True 或 True→False) change_indices = np.where(padded_condition[1:] != padded_condition[:-1])[0] # 按变化点拆分数组,提取所有符合条件的块(奇数位的拆分结果) result = np.split(a, change_indices)[1::2] print(result) # 输出:[array([3, 4]), array([7, 7, 7])]
逻辑解释:
- 只定义一次
condition,完全避免重复计算 - 给条件数组前后补
False,是为了兼容数组开头/结尾就符合条件的场景(比如如果数组第一个元素就满足条件,补False后能正确捕捉到起始分割点) np.where定位条件变化的位置,np.split按这些位置切开数组后,符合条件的块刚好在奇数索引的位置(偶数索引是不符合条件的片段)
方法二:用itertools.groupby(代码更简洁,适合小数组)
如果你的数组规模不大,也可以用Python标准库的groupby,同样只需要写一次判断逻辑:
import numpy as np from itertools import groupby a = np.array([1, 2, 3, 4, 10, 9, 8, 7, 7, 7, 8, 9, 10]) # 这里把判断逻辑写在key里,只定义一次 result = [np.array(list(group)) for is_valid, group in groupby(a, key=lambda x: 2 < x < 8) if is_valid] print(result) # 输出:[array([3, 4]), array([7, 7, 7])]
注意:
这个方法代码更短,但groupby是纯Python循环,处理超大Numpy数组时,效率会比纯Numpy操作低一些,所以根据你的数据规模选就行。
内容的提问来源于stack exchange,提问作者JeffP
相关产品推荐
相关产品推荐

