Use this to learn the idea, then write your own version.
1package main2 3import (4 "bufio"5 . "fmt"6 "io"7 "math/bits"8)9 1011func cf2236G(in io.Reader, _w io.Writer) {12 out := bufio.NewWriter(_w)13 defer out.Flush()14 var T, n, q, v, w, ans, t int15 mem := [21]int{}16 for Fscan(in, &T); T > 0; T-- {17 Fscan(in, &n, &q)18 a := make([]int, n+1)19 for i := 1; i <= n; i++ {20 Fscan(in, &a[i])21 }22 g := make([][]int, n+1)23 for range n - 1 {24 Fscan(in, &v, &w)25 g[v] = append(g[v], w)26 g[w] = append(g[w], v)27 }28 29 mx := bits.Len(uint(n))30 pa := make([][17]int, n+1)31 dep := make([]int, n+1)32 topDep := make([]int, n+1)33 sum := make([]int, n+1)34 lastNZ := make([]int, n+1)35 st := []int{}36 var dfs func(int, int, int, int, int, int)37 dfs = func(v, fa, or, idx, topD, nz int) {38 pa[v][0] = fa39 for i := range mx - 1 {40 pa[v][i+1] = pa[pa[v][i]][i]41 }42 43 if a[v] > 0 {44 st = append(st, v)45 for or&a[v] > 0 {46 or &^= a[st[idx]]47 topD = dep[st[idx]] + 148 idx++49 }50 or |= a[v]51 }52 topDep[v] = topD53 dep[v] = dep[fa] + 154 sum[v] = sum[fa] + dep[v] - topD + 155 56 lastNZ[v] = nz57 if a[v] > 0 {58 nz = v59 }60 61 for _, w := range g[v] {62 if w != fa {63 dfs(w, v, or, idx, topD, nz)64 }65 }66 if a[v] > 0 {67 st = st[:len(st)-1]68 }69 }70 dfs(1, 0, 0, 0, 0, 0)71 72 uptoDep := func(v, d int) int {73 for k := uint32(dep[v] - d); k > 0; k &= k - 1 {74 v = pa[v][bits.TrailingZeros32(k)]75 }76 return v77 }78 getLCA := func(v, w int) int {79 v = uptoDep(v, dep[w])80 if v == w {81 return v82 }83 for i := mx - 1; i >= 0; i-- {84 pv, pw := pa[v][i], pa[w][i]85 if pv != pw {86 v, w = pv, pw87 }88 }89 return pa[v][0]90 }91 up := func(v, dLca int) (int, int) {92 x := v93 for i := mx - 1; i >= 0; i-- {94 p := pa[x][i]95 if topDep[p] > dLca {96 x = p97 }98 }99 if topDep[x] > dLca {100 x = pa[x][0]101 }102 sz := dep[x] - dLca + 1103 return x, sum[v] - sum[x] + sz*(sz+1)/2104 }105 106 for range q {107 Fscan(in, &v, &w)108 if dep[v] < dep[w] {109 v, w = w, v110 }111 lca := getLCA(v, w)112 dLca := dep[lca]113 v, ans = up(v, dLca)114 if w == lca {115 Fprintln(out, ans)116 continue117 }118 w, t = up(w, dLca)119 ans += t - 1120 121 nodes := mem[:0]122 for x := w; dep[x] > dLca; x = lastNZ[x] {123 nodes = append(nodes, x)124 }125 126 or := a[lca]127 for x := v; dep[x] > dLca; x = lastNZ[x] {128 or |= a[x]129 }130 131 depV := dep[v]132 depW := dLca + 1133 for i := len(nodes) - 1; i >= 0; i-- {134 w := nodes[i]135 ans += (depV - dLca) * (dep[w] - depW)136 depW = dep[w]137 for dep[v] > dLca && or&a[w] > 0 {138 or &^= a[v]139 depV = dep[v] - 1140 v = lastNZ[v]141 }142 or |= a[w]143 if i == 0 {144 ans += depV - dLca145 }146 }147 Fprintln(out, ans)148 }149 }150}151 152153