You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

寻求基于NumPy高效识别二维数组中含3的连续块的Pythonic实现方案

寻求基于NumPy高效识别二维数组中含3的连续块的Pythonic实现方案

我先通过以下代码生成笛卡尔积数组:

import numpy as np
import itertools as it

a = np.array([[j for j in range(1,5)] for i in range(3)])
table = np.array(list(it.product(*a)))

然后将每行中第一个3出现位置后的所有元素设为3:

for row in table:
    if len(np.where(row==3)[0])>0:
        row[np.where(row==3)[0][0]:]=3

最后尝试识别全是3的块,用了嵌套循环,效率很低:

nr, nc = table.shape
blocks = []
bh = nr
for i in range(nc-1):
    bh = (int)(bh/4)
    for j in range((int)(nr/bh)):
        temp = table[j*bh:(j+1)*bh,i:]
        if temp[0,0]==3 and temp[-1,-1]==3:
            blocks.append([[j*bh+1,i+1],[j*bh+bh,nc]])
            temp[:,:]=0

得到结果[[[33, 1], [48, 3]], [[9, 2], [12, 3]], [[25, 2], [28, 3]], [[57, 2], [60, 3]]],但这个方法很不高效,求更Pythonic的实现。

嘿,我来帮你把这段代码优化得更Pythonic,同时充分利用NumPy的向量化能力大幅提升效率!咱们一步步拆解优化:

第一步:简化初始数组生成

你原来的代码通过列表推导和itertools.product生成笛卡尔积数组,其实可以更简洁直接:

import numpy as np
import itertools as it

# 直接生成三个[1,2,3,4]的笛卡尔积,转成NumPy数组
table = np.array(list(it.product(range(1, 5), repeat=3)))

这和你原来的代码效果完全一致,但省去了中间数组a的定义,逻辑更清晰。

第二步:向量化处理“第一个3出现后全部设为3”的逻辑

你原来的逐行循环效率不高,换成NumPy的向量化操作,能一次性完成所有行的处理:

nr, nc = table.shape

# 找到每行第一个3的索引,没有3的行设为列数(表示不会触发替换)
first_3_idx = np.argmax(table == 3, axis=1)
# 修正没有3的行:argmax在全False时会返回0,所以我们把这些行的索引设为nc
no_3_mask = ~(table == 3).any(axis=1)
first_3_idx[no_3_mask] = nc

# 生成列索引矩阵,用于和first_3_idx比较
col_indices = np.arange(nc)[np.newaxis, :]
# 对所有列索引 >= 第一个3索引的位置,设为3
table[col_indices >= first_3_idx[:, np.newaxis]] = 3

这段代码用矩阵运算替代了Python级别的循环,数组规模越大,效率提升越明显。

第三步:高效识别全3块

你原来的嵌套循环逻辑可以通过NumPy的布尔矩阵和连续区间分割来优化,避免低效的循环嵌套:

# 生成布尔矩阵:full_3_from_col[i,j]表示第i行从第j列到末尾全是3
# 从右往左累积判断,只要有一个非3,后续列就为False,这个操作是向量化的,很快
full_3_from_col = np.cumprod((table == 3)[:, ::-1], axis=1)[:, ::-1].astype(bool)

blocks = []
# 标记已处理的区域,避免重复统计
processed = np.zeros((nr, nc), dtype=bool)

for col in range(nc):
    # 找到当前列中,满足从该列到末尾全是3且未被处理的行
    valid_rows = np.where(full_3_from_col[:, col] & ~processed[:, col])[0]
    if len(valid_rows) == 0:
        continue
    
    # 把连续的行分割成块:通过diff找出行号不连续的断点
    break_points = np.where(np.diff(valid_rows) != 1)[0] + 1
    row_blocks = np.split(valid_rows, break_points)
    
    for rb in row_blocks:
        start_row = rb[0]
        end_row = rb[-1]
        # 转换为1-based坐标(和你的输出格式完全一致)
        blocks.append([[start_row + 1, col + 1], [end_row + 1, nc]])
        # 标记该区域为已处理,避免重复统计
        processed[start_row:end_row+1, col:] = True

这段代码先快速生成每行从各列开始是否全为3的布尔矩阵,然后对每列找出连续的有效行块,最后记录坐标并标记已处理区域。逻辑清晰,完全利用NumPy的高效运算,比原来的嵌套循环快得多。

运行这段优化后的代码,你会得到和原来完全一致的blocks结果,但代码更简洁、更符合Pythonic风格,处理效率也有显著提升。

备注:内容来源于stack exchange,提问作者Ted Black

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.21 13:53:06