You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

树中所有有向直路径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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 20:05:31