fork download
  1. // ROOT : DRAGON3012009 : WA in Real Life
  2. #include <bits/stdc++.h>
  3. #define FOR(i,l,r) for(int i = l ; i <= r ; i ++)
  4. #define FORD(i,r,l) for(int i = r ; i >= l ; i --)
  5. #define REP(i, a ) for(int i = 0 ; i < a ; i ++ )
  6. #define compare(v) sort((v).begin(), (v).end()); (v).erase(unique((v).begin(), (v).end()), (v).end());
  7. #define ll long long
  8. #define el "\n"
  9. #define fi first
  10. #define se second
  11. #define _ROOT_ int main()
  12. #define M 1000000007
  13. #define MAXN 1000001
  14. #define Bit(i) (1LL << i )
  15. #define INF (1ll<<30)
  16. #define NAME "file"
  17. #define debug(a) cout << #a << " = " << a << endl;
  18. using namespace std;
  19.  
  20. ll n, m, q ;
  21. ll a[MAXN ] ;
  22. ll sz[MAXN ] ;
  23. vector<ll> adj[MAXN ] ;
  24.  
  25. ll Power(ll a, ll b ) { ll res = 1; while(b){ if(b&1) res = a*res%M; b>>=1; a=a*a%M; } return res; }
  26. ll add(ll a, ll b ) { return a + b >= M ? a + b - M : a + b; }
  27. ll mul(ll a, ll b ) { return 1LL * (a%M) * (b%M) % M; }
  28. ll sub(ll a, ll b ) { return a - b < 0 ? a - b + M : a - b; }
  29. ll divi(ll a, ll b) { return 1LL * a * Power(b, M - 2 ) % M; }
  30.  
  31. bool check(ll a ) {
  32. ll t = sqrt(a ) ;
  33. return t * t == a ;
  34. }
  35.  
  36. void dfs(ll u , ll p ,ll &ans ) {
  37. sz[u] = 1 ;
  38. for(ll v : adj[u]) if(v != p ) {
  39. dfs(v , u , ans ) ;
  40. sz[u] += sz[v] ;
  41. }
  42. vector<ll> val ;
  43. if(check(a[u])) {
  44. for(ll v : adj[u]) if(v != p ) val.push_back(sz[v]) ;
  45. val.push_back(n - sz[u]) ;
  46.  
  47. ll sum = 0 , pairr = 0 , trip = 0 ;
  48.  
  49. for(ll v : val ) {
  50. trip = add(trip , mul(pairr , v )) ;
  51. pairr = add(pairr , mul(sum , v )) ;
  52. sum = add(sum , v ) ;
  53. }
  54. ans = add(ans , pairr ) ;
  55. ans = add(ans , trip ) ;
  56. }
  57. }
  58.  
  59. void init() {
  60. cin >> n ;
  61. FOR(i , 2 , n ) {
  62. ll x, y ; cin >> x >> y ;
  63. adj[x].push_back(y) ;
  64. adj[y].push_back(x) ;
  65. }
  66. FOR(i , 1 , n ) cin >> a[i] ;
  67.  
  68. }
  69.  
  70. void solve() {
  71. ll ans = 0 ;
  72. dfs(1 , 1 , ans ) ;
  73. FOR(i , 1 , n ) adj[i].clear() ;
  74. cout << ans << el ;
  75. }
  76.  
  77.  
  78. _ROOT_ {
  79. // freopen(NAME".inp", "r", stdin);
  80. // freopen(NAME".out", "w", stdout) ;
  81. ios_base::sync_with_stdio(0);
  82. cin.tie(0);
  83. cout.tie(0);
  84. int t = 1;// cin >> t ;
  85. while(t--) {
  86. init();
  87. solve();
  88. }
  89. return (0&0);
  90. }
  91.  
Success #stdin #stdout 0.01s 28792KB
stdin
Standard input is empty
stdout
0