fork download
  1. #include<bits/stdc++.h>
  2. using namespace std;
  3.  
  4. #define ll long long
  5.  
  6. int main(){
  7.  
  8. int tt = 1;
  9. cin >> tt;
  10.  
  11. while(tt--){
  12. ll n;
  13. cin >> n;
  14.  
  15. vector<ll> a(n);
  16. for(ll i = 0; i < n; i++) cin >> a[i];
  17.  
  18. vector<vector<ll>> v(n);
  19. for(ll i = 1; i < n; i++){
  20. ll x, y;
  21. cin >> x >> y;
  22. x--;
  23. y--;
  24. v[x].push_back(y);
  25. v[y].push_back(x);
  26. }
  27.  
  28. ll ans = 0;
  29. vector<ll> ss(n, 1);
  30.  
  31. function<void(ll, ll)> dfs = [&](ll u, ll par){
  32.  
  33. vector<ll> val;
  34. for(auto& z : v[u]) if(z != par){
  35. dfs(z, u);
  36. val.push_back(ss[z]);
  37. ss[u] += ss[z];
  38. }
  39.  
  40. val.push_back(n - ss[u]);
  41. assert(accumulate(val.begin(), val.end(), 0ll) == n - 1);
  42.  
  43. if((ll)sqrtl(a[u]) * (ll)sqrtl(a[u]) < a[u]) return;
  44.  
  45. ll sum = 0;
  46. ll pairs = 0;
  47. ll triplets = 0;
  48.  
  49. for(auto& z : val){
  50. triplets += pairs * z;
  51. pairs += sum * z;
  52. sum += z;
  53. }
  54.  
  55. ans += pairs;
  56. ans += triplets;
  57. };
  58.  
  59. dfs(0, -1);
  60.  
  61. cout << ans << '\n';
  62. }
  63.  
  64. return 0;
  65. }
Success #stdin #stdout 0s 5320KB
stdin
4
5
1 1 1 1 1
1 2
2 3
2 4
4 5
10
1 2 3 4 5 6 7 8 9 10
1 3
2 6
6 7
5 4
8 3
3 4
4 6
9 1
10 2
6
12 6 3 18 9 2
3 4
4 5
2 6
6 1
4 2
8
3 16 9 1 8 16 4 9
2 1
3 1
4 3
3 5
6 3
4 7
8 1
stdout
10
48
0
40