使用生成器创建过滤Lambda字典时遇到的异常问题
解决Python循环中Lambda捕获变量的问题,以及Numpy数组按最大值索引分组的优化方案
一、先解决你遇到的Lambda字典异常问题
你碰到的是Python里lambda延迟绑定的典型坑——lambda里的变量i不是在创建时就固定下来,而是等到调用时才会去查找当前作用域的i值。循环结束后i已经变成了3,所以所有lambda实际判断的都是x == 3,自然和预期不符。
有两种简单的解决办法:
方法1:给Lambda添加默认参数绑定当前值
默认参数的赋值是在lambda创建时完成的,这样每个lambda都会持有自己创建时的i值,不会共享循环结束后的最终值:
otv_points = {str(i): lambda x, i=i: x == i for i in range(4)} print(otv_points["0"](0)) # 输出True print(otv_points["0"](3)) # 输出False
方法2:用functools.partial替代Lambda
如果觉得默认参数写法不太直观,可以用partial提前绑定目标值,代码可读性更强:
from functools import partial def is_equal(x, target): return x == target otv_points = {str(i): partial(is_equal, target=i) for i in range(4)}
二、更高效的Numpy数组分组方案
其实你没必要用lambda字典做过滤,Numpy本身提供了更高效的向量操作,直接按最大值索引分组会简单得多:
方案1:直接用布尔掩码过滤(保留原始顺序)
import numpy as np # 原始数组 arr = np.array([[1,2,3], [1,7,2], [1,7,8]]) # 获取每个子数组最大值的索引(axis=1表示按行计算) max_indices = np.argmax(arr, axis=1) result = {} # 遍历所有可能的列索引 for col_idx in range(arr.shape[1]): # 生成布尔掩码,筛选出最大值索引等于当前列的子数组 mask = max_indices == col_idx result[col_idx] = arr[mask].tolist() # 转成列表,需要Numpy数组就去掉.tolist() print(result) # 输出:{0: [], 1: [[1, 7, 2]], 2: [[1, 2, 3], [1, 7, 8]]}
方案2:用itertools.groupby分组(适合不需要保留原始顺序的场景)
如果不要求保留原始子数组的顺序,可以先排序再分组:
import numpy as np from itertools import groupby arr = np.array([[1,2,3], [1,7,2], [1,7,8]]) max_indices = np.argmax(arr, axis=1) # 按最大值索引排序,groupby需要连续的相同键 sorted_idx = np.argsort(max_indices) sorted_arr = arr[sorted_idx] sorted_max_idx = max_indices[sorted_idx] result = {} for idx, group in groupby(zip(sorted_max_idx, sorted_arr), key=lambda x: x[0]): result[idx] = [sub.tolist() for _, sub in group] # 补充没有元素的索引 for col_idx in range(arr.shape[1]): if col_idx not in result: result[col_idx] = [] print(result)
内容的提问来源于stack exchange,提问作者Aerocobra
相关产品推荐
相关产品推荐

