結果
問題 | No.2949 Product on Tree |
ユーザー |
|
提出日時 | 2024-09-23 07:28:36 |
言語 | PyPy3 (7.3.15) |
結果 |
AC
|
実行時間 | 1,450 ms / 2,000 ms |
コード長 | 558 bytes |
コンパイル時間 | 260 ms |
コンパイル使用メモリ | 81,652 KB |
実行使用メモリ | 251,540 KB |
最終ジャッジ日時 | 2024-09-23 07:29:33 |
合計ジャッジ時間 | 49,833 ms |
ジャッジサーバーID (参考情報) |
judge3 / judge1 |
(要ログイン)
ファイルパターン | 結果 |
---|---|
sample | AC * 3 |
other | AC * 46 |
ソースコード
import sys sys.setrecursionlimit(2*10**5) import pypyjit pypyjit.set_param('max_unroll_recursion=-1') mod=998244353 n=int(input()) a=list(map(int,input().split())) g=[[]for i in range(n)] for i in range(n-1): u,v=map(int,input().split()) u-=1 v-=1 g[u].append(v) g[v].append(u) ans=0 def dfs(p,prev): global ans sm1=0 sm2=0 for e in g[p]: if e==prev:continue v=dfs(e,p) sm1=(sm1+v)%mod sm2=(sm2+v*v)%mod ans=(ans+a[p]*sm1)%mod ans=(ans+a[p]*(sm1*sm1-sm2)*pow(2,mod-2,mod))%mod return a[p]*(1+sm1)%mod dfs(0,-1) print(ans)