优化查找满足异或不等式的索引对的Python代码
优化异或对计数的O(n²)算法方案
给定两个长度为n的整数数组a和b,要统计满足a_i ⊕ a_j ≥ b_i ⊕ b_j的索引对(i,j)数量,原嵌套循环的O(n²)复杂度在大数据量下性能瓶颈明显,这里提供基于字典树(Trie)的高效优化方案,时间复杂度可降至O(n * W)(W为整数的二进制位数,通常为32或64)。
核心思路
问题的关键在于逐位比较a_i⊕a_j和b_i⊕b_j的二进制大小:从最高位到最低位,第一个不同的位决定了两者的大小关系。我们可以用字典树存储每个元素对(a_i, b_i)的二进制位信息,遍历每个元素时,通过字典树快速统计符合条件的元素数量。
实现步骤
- 字典树节点设计:每个节点包含子节点分支(对应
(a_i的当前位, b_i的当前位)的4种组合)和经过该节点的元素计数。 - 插入操作:将每个
(a_i, b_i)的二进制位从最高位到最低位插入字典树,更新路径上的节点计数。 - 查询操作:对于每个
(a_j, b_j),遍历字典树,逐位计算a_i⊕a_j和b_i⊕b_j的位值:- 若
a_i⊕a_j的位大于b_i⊕b_j的位,直接累加该分支下的所有元素计数; - 若两位相等,继续深入该分支比较下一位;
- 若
a_i⊕a_j的位更小,跳过该分支。
- 若
Python代码实现
class TrieNode: def __init__(self): self.children = {} # 键为(a_bit, b_bit)元组,值为子节点 self.count = 0 # 经过当前节点的元素总数 def insert(node, a, b, bit): """将(a, b)的二进制位从高位到低位插入字典树""" if bit < 0: node.count += 1 return a_bit = (a >> bit) & 1 b_bit = (b >> bit) & 1 key = (a_bit, b_bit) if key not in node.children: node.children[key] = TrieNode() insert(node.children[key], a, b, bit - 1) node.count += 1 def query(node, a_j, b_j, bit): """查询满足a_i⊕a_j >= b_i⊕b_j的元素数量""" if bit < 0: return node.count # 所有位都相等,满足条件 total = 0 a_j_bit = (a_j >> bit) & 1 b_j_bit = (b_j >> bit) & 1 for (a_bit, b_bit), child in node.children.items(): x_xor = a_bit ^ a_j_bit y_xor = b_bit ^ b_j_bit if x_xor > y_xor: total += child.count elif x_xor == y_xor: total += query(child, a_j, b_j, bit - 1) # x_xor < y_xor 时不满足条件,跳过 return total def count_valid_pairs(a, b): root = TrieNode() max_bit = 31 # 针对32位有符号整数,64位则改为63 # 先插入所有元素到字典树 for ai, bi in zip(a, b): insert(root, ai, bi, max_bit) # 统计所有符合条件的数对 result = 0 for aj, bj in zip(a, b): result += query(root, aj, bj, max_bit) return result
复杂度说明
- 时间复杂度:每个元素插入和查询都需要遍历W位二进制位,总时间为O(n*W),W通常为32或64,远低于O(n²)。
- 空间复杂度:字典树的节点数最多为n*W,空间开销可控,适合处理百万级别的数组。
内容的提问来源于stack exchange,提问作者Maxviz
相关产品推荐
相关产品推荐

