树中所有有向直路径S函数求和的代码性能优化问询
树中所有有向直路径S函数值之和的优化方案
给定一棵含n个顶点的树,每个顶点拥有特殊值Cᵥ。长度k≥1的直路径定义为顶点序列v₁, v₂, ..., vₖ,其中相邻顶点由边连接且所有顶点互不相同;当k=1时,单个顶点也视为直路径。定义函数S(v₁, v₂, ..., vₖ) = Cᵥ₁ - Cᵥ₂ + Cᵥ₃ - Cᵥ₄ + …。需计算树中所有有向直路径的S函数值之和,结果对10⁹+7取模。
你提供的原代码如下:
def S(path): total, negative_one_pow = 0, 1 for node in path: total += (values[node - 1] * negative_one_pow) negative_one_pow *= -1 return total def search(graph): global total for node in range(1, n + 1): queue = [(node, [node])] visited = set() while queue: current_node, path = queue.pop(0) if current_node in visited: continue visited.add(current_node) total += S(path) for neighbor in graph[current_node]: queue.append((neighbor, [*path, neighbor])) n = int(input()) values = list(map(int, input().split())) graph = {i: [] for i in range(1, n + 1)} total = 0 for i in range(n - 1): a, b = map(int, input().split()) graph[a].append(b) graph[b].append(a) search(graph) print(total % 1000000007)
这段代码的问题在于枚举所有路径并逐个计算S值,时间复杂度实际为O(n³)(路径总数O(n²),每条路径计算S值需遍历节点),完全无法处理n≥1000的大规模树。以下是针对性的优化方案:
核心优化思路:数学推导+树形DP
不再枚举所有路径,而是计算每个节点的Cᵥ在所有路径中的贡献值,直接累加得到总和。
1. 节点贡献分析
对于每个节点u,只需计算:
cnt_odd:u作为路径第奇数位(第1、3、5...位)的路径总数cnt_even:u作为路径第偶数位(第2、4、6...位)的路径总数
u对总和的贡献为Cᵥ * (cnt_odd - cnt_even),最终总和是所有节点贡献的累加。
2. 树形DP计算贡献数量
通过两次DFS遍历树,分别计算子树内和子树外的路径数量:
- 第一次后序DFS:计算每个节点的子树内,以u为终点的奇数/偶数长度路径数(
in_odd[u]、in_even[u]),同时统计子树大小。 - 第二次前序DFS:基于父节点的信息,计算子树外以u为终点的奇数/偶数长度路径数(
out_odd[u]、out_even[u])。
3. 优化后的代码实现
MOD = 10**9 + 7 def main(): import sys sys.setrecursionlimit(1 << 25) n = int(sys.stdin.readline()) values = list(map(int, sys.stdin.readline().split())) adj = [[] for _ in range(n+1)] for _ in range(n-1): a, b = map(int, sys.stdin.readline().split()) adj[a].append(b) adj[b].append(a) # 第一次DFS:计算子树内的路径数和子树大小 in_odd = [0]*(n+1) in_even = [0]*(n+1) size = [1]*(n+1) def dfs1(u, parent): in_odd[u] = 1 # 自身是长度1的奇数路径 in_even[u] = 0 for v in adj[u]: if v == parent: continue dfs1(v, u) size[u] += size[v] # 子节点的偶数路径加上当前节点变成奇数路径,反之亦然 in_odd[u] = (in_odd[u] + in_even[v]) % MOD in_even[u] = (in_even[u] + in_odd[v]) % MOD dfs1(1, -1) # 第二次DFS:计算子树外的路径数 out_odd = [0]*(n+1) out_even = [0]*(n+1) def dfs2(u, parent): for v in adj[u]: if v == parent: continue # 父节点u的总偶数路径 = 子树外偶数路径 + 子树内除v外的偶数路径 parent_even = (out_even[u] + (in_even[u] - in_odd[v]) % MOD) % MOD parent_odd = (out_odd[u] + (in_odd[u] - in_even[v]) % MOD) % MOD # v的子树外奇数路径 = 父节点的偶数路径(加上v后变为奇数) out_odd[v] = parent_even out_even[v] = parent_odd dfs2(v, u) dfs2(1, -1) total = 0 for u in range(1, n+1): cnt_odd = (in_odd[u] + out_odd[u]) % MOD cnt_even = (in_even[u] + out_even[u]) % MOD contribution = (values[u-1] * ((cnt_odd - cnt_even) % MOD)) % MOD total = (total + contribution) % MOD # 处理取模后的负数情况 total = (total + MOD) % MOD print(total) if __name__ == "__main__": main()
4. 优化效果
时间复杂度从O(n³)降至O(n),可以轻松处理n=1e5级别的树,完全解决大规模数据的性能问题。
内容的提问来源于stack exchange,提问作者xaller
相关产品推荐
相关产品推荐

