問題
$N$頂点の木があり、頂点$v$には色$C_v$($1\le C_v\le N$)が塗られている。$Q$個のクエリが与えられ、各クエリは頂点対$(u,v)$で、$u$から$v$への単純パス上に現れる相異なる色の種類数を答えよ。全クエリはオフラインで処理してよい。
木を「オイラーツアー」で長さ$2N$の列に展開し(各頂点を入場時・退場時の2回訪問)、パスクエリ$(u,v)$を区間$[L,R]$に変換する。区間内で「奇数回登場する頂点だけが実際にパス上にある」というパリティ判定を使い、通常のMo's Algorithmと同じadd/removeの枠組みで処理する。
入力形式
N Q
C_1 C_2 ... C_N
a_1 b_1
:
a_{N-1} b_{N-1}
u_1 v_1
:
u_Q v_Q
制約
$2 \le N \le 2\times10^4$
$1 \le Q \le 2\times10^4$
$1 \le C_i \le N$
辺は木を構成する
入出力例
入力例1
7 3
1 2 1 3 2 1 2
1 2
1 3
2 4
2 5
3 6
3 7
4 7
5 6
4 5
出力例1
3
2
2
1問目 (4,7): パス 4-2-1-3-7、色 3,2,1,1,2 → {1,2,3} で3種類。2問目 (5,6): パス 5-2-1-3-6、色 2,2,1,1,1 → {1,2} で2種類。3問目 (4,5): パス 4-2-5、色 3,2,2 → {2,3} で2種類。
概念図
ヒント(段階的開示)
ヒント1: 方向性
各クエリごとに愚直にLCAを求めパスをたどると$O(N)$かかり、$Q$個で$O(NQ)$となり間に合わない。配列の区間クエリで使うMo's Algorithmの考え方(クエリをソートして尺取り的にadd/removeする)を木に応用できないか考えよう。
ヒント2: アプローチ
木をDFSして各頂点を「最初に訪れたとき」と「部分木を訪れ終えて戻るとき」の2回、長さ$2N$の配列(オイラーツアー)に記録する。クエリ$(u,v)$($\text{first}[u]\le\text{first}[v]$に並べ替え)に対し$l=\text{LCA}(u,v)$とすると、$l=u$なら区間$[\text{first}[u],\text{first}[v]]$、$l\ne u$なら区間$[\text{last}[u],\text{first}[v]]$を使い、後者では$l$の色を別途1つ追加でカウントする。
ヒント3: 誘導(コード骨格)
def toggle(v):
c = color[v]
if active[v]:
cnt[c] -= 1
if cnt[c] == 0: distinct -= 1
active[v] = False
else:
cnt[c] += 1
if cnt[c] == 1: distinct += 1
active[v] = True
# クエリを (left, right, extra_lca_or_-1, index) にしてMo's Algorithmでソート
# block = sqrt(2N)、答え = distinct + (extraがありcnt[color[extra]]==0なら+1)
LCAはダブリング(Binary Lifting)で$O(\log N)$にしておく必要がある。
模範解答 (Python)
import sys
from math import isqrt
def solve():
data = sys.stdin.buffer.read().split()
idx = 0
n = int(data[idx]); idx += 1
q = int(data[idx]); idx += 1
color = [0] + [int(data[idx + i]) for i in range(n)]; idx += n
adj = [[] for _ in range(n + 1)]
for _ in range(n - 1):
a = int(data[idx]); idx += 1
b = int(data[idx]); idx += 1
adj[a].append(b)
adj[b].append(a)
queries = []
for _ in range(q):
u = int(data[idx]); idx += 1
v = int(data[idx]); idx += 1
queries.append((u, v))
LOG = max(1, (n).bit_length() + 1)
parent = [[0] * (n + 1) for _ in range(LOG)]
depth = [0] * (n + 1)
tour = []
first = [0] * (n + 1)
last = [0] * (n + 1)
visited = [False] * (n + 1)
order_stack = [1]
it_stack = [iter(adj[1])]
visited[1] = True
parent[0][1] = 0
first[1] = len(tour); tour.append(1)
while it_stack:
u = order_stack[-1]
it = it_stack[-1]
advanced = False
for v in it:
if not visited[v]:
visited[v] = True
parent[0][v] = u
depth[v] = depth[u] + 1
first[v] = len(tour); tour.append(v)
order_stack.append(v)
it_stack.append(iter(adj[v]))
advanced = True
break
if not advanced:
uu = order_stack.pop()
it_stack.pop()
last[uu] = len(tour); tour.append(uu)
for k in range(1, LOG):
for v in range(1, n + 1):
parent[k][v] = parent[k - 1][parent[k - 1][v]]
def lca(u, v):
if depth[u] < depth[v]:
u, v = v, u
diff = depth[u] - depth[v]
for k in range(LOG):
if (diff >> k) & 1:
u = parent[k][u]
if u == v:
return u
for k in range(LOG - 1, -1, -1):
if parent[k][u] != parent[k][v]:
u = parent[k][u]
v = parent[k][v]
return parent[0][u]
qs = []
for query_idx, (u, v) in enumerate(queries):
if first[u] > first[v]:
u, v = v, u
l = lca(u, v)
if l == u:
left, right, extra = first[u], first[v], -1
else:
left, right, extra = last[u], first[v], l
qs.append((left, right, extra, query_idx))
block = max(1, isqrt(2 * n))
qs.sort(key=lambda t: (t[0] // block, t[1] if (t[0] // block) % 2 == 0 else -t[1]))
active = [False] * (n + 1)
cnt_color = [0] * (n + 1)
distinct = 0
def toggle(v):
nonlocal distinct
c = color[v]
if active[v]:
cnt_color[c] -= 1
if cnt_color[c] == 0:
distinct -= 1
active[v] = False
else:
cnt_color[c] += 1
if cnt_color[c] == 1:
distinct += 1
active[v] = True
cur_l, cur_r = 0, -1
ans = [0] * q
for left, right, extra, query_idx in qs:
while cur_r < right:
cur_r += 1
toggle(tour[cur_r])
while cur_l > left:
cur_l -= 1
toggle(tour[cur_l])
while cur_r > right:
toggle(tour[cur_r])
cur_r -= 1
while cur_l < left:
toggle(tour[cur_l])
cur_l += 1
res = distinct
if extra != -1 and cnt_color[color[extra]] == 0:
res += 1
ans[query_idx] = res
sys.stdout.write("\n".join(map(str, ans)) + "\n")
solve()
計算量: $O((N+Q)\sqrt{N}\log N)$(Mo's Algorithm本体$O((N+Q)\sqrt N)$ + LCA計算$O(Q\log N)$)。
Step-by-Step 解説
1オイラーツアーの構築
各頂点を入場時・退場時の2回、長さ$2N$の配列に記録する(再帰深さ対策で反復DFSを使う)。
各頂点を入場時・退場時の2回、長さ$2N$の配列に記録する(再帰深さ対策で反復DFSを使う)。
2LCAの前計算(ダブリング)
$O(N\log N)$の前計算で任意の2頂点のLCAを$O(\log N)$で求める。
$O(N\log N)$の前計算で任意の2頂点のLCAを$O(\log N)$で求める。
3クエリを区間へ変換
$l=u$なら$[\text{first}[u],\text{first}[v]]$、そうでなければ$[\text{last}[u],\text{first}[v]]$+$l$の色を別途加算。
$l=u$なら$[\text{first}[u],\text{first}[v]]$、そうでなければ$[\text{last}[u],\text{first}[v]]$+$l$の色を別途加算。
4Mo's Algorithmでの尺取り
ブロックサイズ$\sqrt{2N}$でソートし、左右ポインタを動かしながら`toggle`する。
ブロックサイズ$\sqrt{2N}$でソートし、左右ポインタを動かしながら`toggle`する。
5答えの算出
`distinct`に、必要なら$l$の色(未アクティブなら)を加えたものが答え。
`distinct`に、必要なら$l$の色(未アクティブなら)を加えたものが答え。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| 頂点を1回しかツアーに記録しない | 部分木クエリ用のEuler Tourと混同 | 各頂点を入場・退場の2回記録し、長さ$2N$の配列にする |
| $l=u$の場合も$l$を別途加算してしまう | パリティの仕組みを誤解 | $l=u$のときは区間内に$l$が自然に1回含まれるため追加加算は不要 |
| 再帰DFSでツアー構築しRecursionErrorになる | Pythonの再帰上限に達する | 明示的なスタックを使った反復DFSで実装する |
| ブロックサイズを$\sqrt N$にしてしまう | 配列長が$N$でなく$2N$であることを見落とす | ツアー長$2N$に対し$\sqrt{2N}$程度をブロックサイズにする |
次のステップ
- 発展: パス上の色の最頻値・中央値などdistinctカウント以外の集約値クエリへの拡張
- 発展: 頂点の色の更新も混ざる場合(時間軸を追加した3次元Mo's Algorithm)
- 次回予告: モンゴメリ乗算(Montgomery Multiplication・REDCアルゴリズムによる高速mod演算)