如何用二叉树实现栈前k元素的O(log n)时间复杂度部分求和
自定义栈的O(log n) ksum实现方案
问题背景
需要实现一个支持以下功能的自定义栈:
- 从栈顶添加/移除元素
- 修改指定位置元素的值
- 计算栈顶前k个元素的和,时间复杂度需达到O(log n)
当前基于ArrayList的ksum实现为O(n),不符合要求,需通过求和二叉树(叶子节点为栈元素,父节点为子节点和)来优化。
解决方案代码实现
我们通过维护一棵以数组存储的完全二叉求和树来实现所有操作的O(log n)复杂度,以下是修改后的完整代码:
import java.util.Iterator; import java.util.PriorityQueue; import java.util.Stack; public class UltraFast implements UltraStack { protected Stack<Integer> stack; protected PriorityQueue<Integer> maxHeap; protected long[] sumTree; // 求和二叉树,数组形式存储 protected int treeCapacity; // 求和树容量,始终为2的幂次 public UltraFast() { stack = new Stack<>(); maxHeap = new PriorityQueue<>((a, b) -> Integer.compare(b, a)); treeCapacity = 1; sumTree = new long[treeCapacity]; } @Override public void push(int x) { stack.push(x); maxHeap.offer(x); // 扩容求和树:当栈大小超过当前树容量的一半时,将容量翻倍 if (stack.size() > treeCapacity / 2) { treeCapacity *= 2; long[] newTree = new long[treeCapacity]; System.arraycopy(sumTree, 0, newTree, 0, sumTree.length); sumTree = newTree; } // 将新元素加入求和树的对应叶子节点,并向上更新父节点的和 int leafPos = treeCapacity / 2 + (stack.size() - 1); sumTree[leafPos] = x; int parentPos = (leafPos - 1) / 2; while (parentPos >= 0) { sumTree[parentPos] = sumTree[2 * parentPos + 1] + sumTree[2 * parentPos + 2]; parentPos = (parentPos - 1) / 2; } } @Override public Integer pop() { if (stack.isEmpty()) { return null; } int top = stack.pop(); maxHeap.remove(top); // 找到栈顶元素在求和树中的叶子节点,置为0后向上更新父节点 int leafPos = treeCapacity / 2 + stack.size(); sumTree[leafPos] = 0; int parentPos = (leafPos - 1) / 2; while (parentPos >= 0) { sumTree[parentPos] = sumTree[2 * parentPos + 1] + sumTree[2 * parentPos + 2]; parentPos = (parentPos - 1) / 2; } // 可选缩容:当栈大小小于当前树容量的1/4时,将容量减半(优化空间占用) if (stack.size() < treeCapacity / 4 && treeCapacity > 1) { treeCapacity /= 2; long[] newTree = new long[treeCapacity]; // 复制有效叶子节点到新树 System.arraycopy(sumTree, treeCapacity, newTree, treeCapacity / 2, stack.size()); // 重新计算新树的非叶子节点和 for (int i = treeCapacity / 2 - 1; i >= 0; i--) { newTree[i] = newTree[2 * i + 1] + newTree[2 * i + 2]; } sumTree = newTree; } return top; } @Override public Integer set(int i, int x) { if (i < 0 || i >= stack.size()) { return null; } int oldValue = stack.get(i); stack.set(i, x); maxHeap.remove(oldValue); maxHeap.offer(x); // 更新求和树中对应叶子节点的值,并向上更新父节点 int leafPos = treeCapacity / 2 + i; sumTree[leafPos] = x; int parentPos = (leafPos - 1) / 2; while (parentPos >= 0) { sumTree[parentPos] = sumTree[2 * parentPos + 1] + sumTree[2 * parentPos + 2]; parentPos = (parentPos - 1) / 2; } return oldValue; } @Override public long ksum(int k) { if (k <= 0) { return 0; } int actualK = Math.min(k, stack.size()); // 转换为栈中元素的索引区间:从栈底侧的size()-actualK到栈顶的size()-1 int stackLeft = stack.size() - actualK; int stackRight = stack.size() - 1; // 转换为求和树中的叶子节点位置 int treeLeft = treeCapacity / 2 + stackLeft; int treeRight = treeCapacity / 2 + stackRight; // 递归查询区间和 return querySum(treeLeft, treeRight, 0, 0, treeCapacity - 1); } // 递归查询求和树中[ql, qr]区间的和,当前节点node对应树的区间[nodeL, nodeR] private long querySum(int ql, int qr, int node, int nodeL, int nodeR) { if (qr < nodeL || ql > nodeR) { return 0; } if (ql <= nodeL && nodeR <= qr) { return sumTree[node]; } int mid = (nodeL + nodeR) / 2; long leftSum = querySum(ql, qr, 2 * node + 1, nodeL, mid); long rightSum = querySum(ql, qr, 2 * node + 2, mid + 1, nodeR); return leftSum + rightSum; } @Override public Integer get(int i) { if (i < 0 || i >= stack.size()) { return null; } return stack.get(i); } @Override public Integer max() { return maxHeap.isEmpty() ? null : maxHeap.peek(); } @Override public int size() { return stack.size(); } @Override public Iterator<Integer> iterator() { return stack.iterator(); } }
关键实现说明
- 求和树结构:采用数组存储完全二叉树,容量始终保持为2的幂次,方便计算叶子节点位置和区间查询。栈中第
i个元素(栈底为0,栈顶为size()-1)对应求和树的叶子节点位置为treeCapacity/2 + i。 - push操作:元素入栈后,将其加入求和树的对应叶子节点,然后向上遍历更新所有父节点的和,时间复杂度O(log n)。
- pop操作:弹出栈顶元素后,将求和树中对应叶子节点置为0,向上更新父节点和。可选缩容操作优化空间,时间复杂度O(log n)。
- set操作:修改栈中元素后,更新求和树对应叶子节点的值,向上更新父节点和,时间复杂度O(log n)。
- ksum操作:通过递归查询求和树的指定区间(栈顶k个元素对应的叶子节点区间)来计算和,时间复杂度O(log n)。
内容的提问来源于stack exchange,提问作者fil1423
相关产品推荐
相关产品推荐

