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 1011type seg20 []struct {12 l, r int13 mask uint14 todo uint15}16 17func (t seg20) do(o int, v uint) {18 t[o].mask = v19 t[o].todo = v20}21 22func (t seg20) spread(o int) {23 if v := t[o].todo; v > 0 {24 t.do(o<<1, v)25 t.do(o<<1|1, v)26 t[o].todo = 027 }28}29 30func (t seg20) build(a []uint, o, l, r int) {31 t[o].l, t[o].r = l, r32 if l == r {33 t[o].mask = a[l-1]34 return35 }36 m := (l + r) >> 137 t.build(a, o<<1, l, m)38 t.build(a, o<<1|1, m+1, r)39 t.maintain(o)40}41 42func (t seg20) update(o, l, r int, v uint) {43 if l <= t[o].l && t[o].r <= r {44 t.do(o, v)45 return46 }47 t.spread(o)48 m := (t[o].l + t[o].r) >> 149 if l <= m {50 t.update(o<<1, l, r, v)51 }52 if m < r {53 t.update(o<<1|1, l, r, v)54 }55 t.maintain(o)56}57 58func (t seg20) maintain(o int) {59 t[o].mask = t[o<<1].mask | t[o<<1|1].mask60}61 62func (t seg20) query(o, l, r int) uint {63 if l <= t[o].l && t[o].r <= r {64 return t[o].mask65 }66 t.spread(o)67 m := (t[o].l + t[o].r) >> 168 if r <= m {69 return t.query(o<<1, l, r)70 }71 if l > m {72 return t.query(o<<1|1, l, r)73 }74 return t.query(o<<1, l, r) | t.query(o<<1|1, l, r)75}76 77func cf620E(_r io.Reader, _w io.Writer) {78 in := bufio.NewReader(_r)79 out := bufio.NewWriter(_w)80 defer out.Flush()81 82 var n, m, dfn, op, v, w int83 Fscan(in, &n, &m)84 a := make([]int, n)85 for i := range a {86 Fscan(in, &a[i])87 }88 g := make([][]int, n)89 for i := 1; i < n; i++ {90 Fscan(in, &v, &w)91 v--92 w--93 g[v] = append(g[v], w)94 g[w] = append(g[w], v)95 }96 97 b := make([]uint, n)98 nodes := make([]struct{ l, r int }, n)99 var dfs func(int, int) int100 dfs = func(v, fa int) (size int) {101 b[dfn] = 1 << a[v]102 dfn++103 nodes[v].l = dfn104 for _, w := range g[v] {105 if w != fa {106 sz := dfs(w, v)107 size += sz108 }109 }110 nodes[v].r = nodes[v].l + size111 size++112 return113 }114 dfs(0, -1)115 116 t := make(seg20, 2<<bits.Len(uint(n-1)))117 t.build(b, 1, 1, n)118 for ; m > 0; m-- {119 Fscan(in, &op, &v)120 o := nodes[v-1]121 if op == 1 {122 Fscan(in, &w)123 t.update(1, o.l, o.r, 1<<w)124 } else {125 Fprintln(out, bits.OnesCount(t.query(1, o.l, o.r)))126 }127 }128}129 130131