RT。
思路是主席树维护出区间不同颜色数,单调队列维护三元组求解。
复杂度是 O(nklog2n) 的,比正常做法多了一个 log,极限数据大概 4,5 秒。
#include <cstdio>
#include <cstring>
#include <cassert>
#include <iostream>
#include <algorithm>
using std::max;
using std::cin;
using std::cout;
const int MAXN = 3.5e4 + 5;
int n, k;
int a[MAXN << 1];
int lst[MAXN << 1];
int f[MAXN][55];
struct Str {
int l, r, p;
Str() {}
Str(int L, int R, int P) {
l = L, r = R, p = P;
}
} que[MAXN << 1], *hd, *tl;
namespace PreTree {
const int L = 1, R = MAXN;
int rt[MAXN << 1], tot;
struct Seg {
int ls, rs;
int val;
Seg() { ls = rs = val = 0; }
} tr[MAXN * 300];
void pushup(int u) {
tr[u].val = tr[ tr[u].ls ].val + tr[ tr[u].rs ].val;
}
void modify(int& u, int v, int pos, int val, int l, int r) {
u = ++tot, tr[u] = tr[v];
if (l == r) {
tr[u].val += val; return;
}
int mid = l + r >> 1;
if (pos <= mid) modify(tr[u].ls, tr[v].ls, pos, val, l, mid);
else modify(tr[u].rs, tr[v].rs, pos, val, mid + 1, r);
pushup(u);
}
void modify(int ver, int pre, int pos, int val) {
modify(rt[ver], rt[pre], pos, val, L, R);
}
int query(int u, int ql, int qr, int l, int r) {
if (!u) return 0;
if (l >= ql && r <= qr) return tr[u].val;
int mid = l + r >> 1;
if (qr <= mid) return query(tr[u].ls, ql, qr, l, mid);
else if (ql > mid) return query(tr[u].rs, ql, qr, mid + 1, r);
else return query(tr[u].ls, ql, qr, l, mid) + query(tr[u].rs, ql, qr, mid + 1, r);
}
int query(int ver, int l, int r) {
return query(rt[ver], l, r, L, R);
}
}
using PreTree::modify;
using PreTree::query;
bool chk(int i, int j, int p1, int p2) {
return f[p1][j - 1] + query(i, p1 + 1, i) <= f[p2][j - 1] + query(i, p2 + 1, i);
}
int main() {
std::ios::sync_with_stdio(false), cin.tie(nullptr);
cin >> n >> k;
for (int i = 1; i <= n; ++i) cin >> a[i];
for (int i = 1; i <= n; ++i) {
modify(i, i - 1, i, 1);
if (lst[ a[i] ]) modify(i, i, lst[ a[i] ], -1);
lst[ a[i] ] = i; f[i][1] = query(i, 1, i);
}
for (int j = 2; j <= k; ++j) {
hd = tl = que; *tl++ = Str(j, n, j - 1);
for (int i = j; i <= n; ++i) {
int p = (*hd).p;
f[i][j] = f[p][j - 1] + query(i, p + 1, i);
while (hd != tl && (*hd).r <= i) ++hd;
(*hd).l = i + 1;
while (hd != tl && chk((*(tl - 1)).l, j, (*(tl - 1)).p, i)) --tl;
if (hd == tl) *tl++ = Str(i + 1, n, i);
else {
int l = (*(tl - 1)).l, r = (*(tl - 1)).r, res = (*(tl - 1)).r + 1;
while (l <= r) {
int mid = l + r >> 1;
if (chk(mid, j, (*(tl - 1)).p, i)) res = mid, r = mid - 1;
else l = mid + 1;
}
(*(tl - 1)).r = res - 1;
if (res != n + 1) *tl++ = Str(res, n, i);
}
}
}
cout << f[n][k] << '\n';
return 0;
}