不知道是自己哪里写挂了还是思路不对/dk。
#include <iostream>
using namespace std;
const int MAX_SIZE = 2e7;
#define int long long
namespace FastIO
{
template <typename T = int>
inline T read() {
T s = 0, w = 1;
char c = getchar();
while (!isdigit(c)) {
if (c == '-') w = -1;
c = getchar();
}
while (isdigit(c)) s = (s << 1) + (s << 3) + (c ^ 48), c = getchar();
return s * w;
}
template <typename T>
inline void read(T &s) {
s = 0;
int w = 1;
char c = getchar();
while (!isdigit(c)) {
if (c == '-') w = -1;
c = getchar();
}
while (isdigit(c)) s = (s << 1) + (s << 3) + (c ^ 48), c = getchar();
s = s * w;
}
template <typename T, typename... Args>
inline void read(T &x, Args &...args) {
read(x), read(args...);
}
template <typename T>
inline void write(T x, char ch) {
if (x < 0) x = -x, putchar('-');
static char stk[25];
int top = 0;
do {
stk[top++] = x % 10 + '0', x /= 10;
} while (x);
while (top) putchar(stk[--top]);
putchar(ch);
return;
}
} // namespace FastIO
using namespace FastIO;
struct SegTree {
int lc, rc;
int sum;
int counter;
};
SegTree seg[(MAX_SIZE << 1) + 10];
int tot;
int build() {
++tot;
seg[tot].sum = seg[tot].lc = seg[tot].rc = seg[tot].counter = 0;
return tot;
}
void insert(int p, int l, int r, int val, int data) {
if (l == r) {
seg[p].counter += data;
seg[p].sum += data * (val - 1);
return;
}
int mid = (l + r) >> 1;
if (val <= mid) {
if (!seg[p].lc) seg[p].lc = build();
insert(seg[p].lc, l, mid, val, data);
} else {
if (!seg[p].rc) seg[p].rc = build();
insert(seg[p].rc, mid + 1, r, val, data);
}
seg[p].counter = seg[seg[p].lc].counter + seg[seg[p].rc].counter;
seg[p].sum = seg[seg[p].lc].sum + seg[seg[p].rc].sum;
return;
}
int GetAns(int tl, int tr, int siz, int p) {
if (!p || !siz) return 0;
if (tl == tr) {
return seg[p].sum;
}
int val = 0;
int mid = (tl + tr) >> 1;
if (siz >= seg[seg[p].rc].counter) {
val += seg[seg[p].rc].sum;
val += GetAns(tl, mid, siz - seg[seg[p].rc].counter, seg[p].lc);
} else {
val += GetAns(mid + 1, tr, siz, seg[p].rc);
}
return val;
}
const int MAX_NUM = 1e9 + 10;
const int MAX = 5.1e5;
int ori[MAX];
signed main() {
tot = 0;
int root = build();
int n, k, Q;
read(n, k, Q);
while (Q--) {
int x = read();
int y = read();
if (ori[x] != 0) {
insert(root, 1, MAX_NUM, ori[x] + 1, -1);
}
ori[x] = y;
insert(root, 1, MAX_NUM, ori[x] + 1, 1);
int ans = GetAns(1, MAX_NUM, k, root);
printf("%lld\n", ans);
}
return 0;
}