如何将itertools生成的powerset转换为列对齐的numpy数组
优化实现方案
这里提供两种可直接使用的实现,第一种是在原有逻辑基础上做性能优化,第二种是内存和速度都最优的位掩码实现:
方案1:预存索引映射优化原有逻辑
核心优化点是提前把元素和对应列索引的映射存成字典,避免每次循环都重复调用np.where查找位置,代码如下:
from itertools import combinations, chain import numpy as np def columnizePowerset(powerset, items): # 提前构建元素到列索引的映射,查找复杂度O(1) item2idx = {val: idx for idx, val in enumerate(items)} n_rows = len(powerset) n_cols = len(items) # 初始化空字符串数组 res = np.full((n_rows, n_cols), fill_value='', dtype=str) # 直接遍历幂集,不需要额外转object类型数组 for row_idx, subset in enumerate(powerset): for val in subset: res[row_idx, item2idx[val]] = val return res def powerset(iterable): s = list(iterable) return tuple(chain.from_iterable(combinations(s, r) for r in range(1, len(s)+1))) if __name__ == "__main__": items = ('a', 'b', 'c') combos = powerset(items) res = columnizePowerset(combos, items) print(res)
和原有实现相比,去掉了多余的数组转换操作和重复的查找逻辑,当items元素较多时,速度提升非常明显。
方案2:位掩码实现(效率最高)
幂集的每个非空子集,刚好对应从1到2^n - 1的整数的二进制表示,每一位为1就代表对应位置的元素在子集中存在,完全不需要提前生成幂集,直接按位判断赋值即可,内存和速度都最优:
import numpy as np def get_columnized_powerset(items): n = len(items) total_rows = 2 ** n - 1 res = np.full((total_rows, n), fill_value='', dtype=str) for row_idx in range(1, total_rows + 1): # 遍历每一位判断是否为1 for col_idx in range(n): if row_idx & (1 << col_idx): res[row_idx - 1, col_idx] = items[col_idx] # 如果需要和原幂集完全一致的输出顺序(按子集长度从小到大排序),打开下面一行注释即可 # res = res[np.argsort((res != '').sum(axis=1))] return res if __name__ == "__main__": items = ('a', 'b', 'c') res = get_columnized_powerset(items) print(res)
两种方案的输出都完全符合要求:
[['a' '' ''] ['' 'b' ''] ['a' 'b' ''] ['' '' 'c'] ['a' '' 'c'] ['' 'b' 'c'] ['a' 'b' 'c']]
内容的提问来源于stack exchange,提问作者Craig Nathan
相关产品推荐
相关产品推荐

