풀이
입력은 정수 1..M에서 정의된 함수 f를 주며, 각 쿼리는 시작값 x에 함수를 K번 적용한 값, 즉 f^K(x)를 묻습니다. 함수를 한 번씩 적용하면 쿼리마다 최대 500,000번의 연산이 필요할 수 있으므로, 이진 리프팅으로 함수의 거듭제곱을 미리 계산합니다.
up[b][x]를 x에 f를 정확히 2^b번 적용한 결과라고 정의합니다. 기본 행은 up[0][x] = f(x)입니다. 길이가 2^(b-1)인 점프를 연속해서 두 번 하면 길이 2^b인 점프가 되므로 다음 점화식을 얻습니다.
up[b][x] = up[b - 1][up[b - 1][x]]
쿼리에서는 K의 각 비트를 확인합니다. 비트 b가 켜져 있으면 현재 값을 up[b][현재 값]으로 갱신합니다. 선택한 점프 길이의 합이 K이므로 정확히 f^K(x)를 계산합니다. K = 0이면 켜진 비트가 없어 답은 그대로 x이고, K = 1이면 기본 행만 사용합니다. 표는 20개 행이며, K <= 500000에서 가장 높은 비트는 18번 비트이므로 모든 쿼리를 처리할 때 인덱스 범위를 넘지 않습니다.
전처리 시간과 메모리는 Kmax = 500000일 때 O(M log Kmax)입니다. 각 쿼리는 O(log Kmax) 시간에 처리합니다.
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
import java.io.BufferedInputStream;
import java.io.IOException;
public class Main {
private static final int LEVELS = 20;
public static void main(String[] args) throws IOException {
FastScanner input = new FastScanner();
int m = input.nextInt();
int[][] up = new int[LEVELS][m + 1];
for (int x = 1; x <= m; x++) {
up[0][x] = input.nextInt();
}
for (int bit = 1; bit < LEVELS; bit++) {
for (int x = 1; x <= m; x++) {
up[bit][x] = up[bit - 1][up[bit - 1][x]];
}
}
int queryCount = input.nextInt();
StringBuilder answer = new StringBuilder();
for (int query = 0; query < queryCount; query++) {
int k = input.nextInt();
int value = input.nextInt();
for (int bit = 0; bit < LEVELS; bit++) {
if ((k & (1 << bit)) != 0) {
value = up[bit][value];
}
}
answer.append(value).append('\n');
}
System.out.print(answer);
}
private static class FastScanner {
private final BufferedInputStream in = new BufferedInputStream(System.in);
private final byte[] buffer = new byte[1 << 16];
private int pointer;
private int length;
int nextInt() throws IOException {
int c;
do {
c = read();
} while (c <= ' ');
int value = 0;
while (c > ' ') {
value = value * 10 + c - '0';
c = read();
}
return value;
}
private int read() throws IOException {
if (pointer == length) {
length = in.read(buffer);
pointer = 0;
if (length == -1) {
return -1;
}
}
return buffer[pointer++];
}
}
}