CF833B The Bakery 解题报告

题目链接

题意

有一个有 \(n\) 个元素的序列 \(a\)\(a_i\) 代表该位置上的种类号。你需要将序列化分为 \(k\) 段,使得每一段的种类数之和最大。

\(n \le 35000,\ k \le \min(n,\ 50)\)

\(a_i \le n\)

思路

朴素的dp:

考虑用 \(dp_{i,j}\) 表示序列 \(1\)\(i\) 之间划分 \(j\) 次的最大价值。

于是有状态转移方程:

\(dp_{i,j} = \sum\max(dp_{k, j - 1} + cnt_{k + 1, i})\)

其中 \(cnt_{i, j}\) 表示从 \(i\)\(j\) 之间的种类数。

复杂度 \(O(n ^ 3 k)\),显然不能被接受。

考虑优化

用上一层的划分最大值建树,所以一共要建 \(k\) 棵树。

发现每个位置贡献的种类数都有范围。

然后存进线段树里就可以了。

代码

#include<bits/stdc++.h>
using namespace std;
const int N = 35005;
int dp[51][N], a[N], pre[N], sum[N], n, k;
struct node
{
    int data, lzy;
}w[N * 4];
int InRange(int l, int r, int L, int R)
{
    return (l >= L) && (r <= R);
}
int OutofRange(int l, int r, int L, int R)
{
    return (r < L) || (l > R);
}
void maketag(int u, int x, int l, int r)
{
    w[u].lzy += x;
    w[u].data += x;
}
void pushup(int u)
{
    w[u].data = max(w[u * 2].data, w[u * 2 + 1].data);
}
void pushdown(int u, int l, int r)
{
    if(!w[u].lzy)
    {
        return ;
    }
    int mid = l + r >> 1;
    maketag(u * 2, w[u].lzy, l, mid);
    maketag(u * 2 + 1, w[u].lzy, mid + 1, r);
    w[u].lzy = 0;
}
void build(int u, int l, int r, int x)
{
	if(l == r)
	{
		w[u].data = dp[x][l - 1];
		return ;
	}
	int mid = l + r >> 1;
	build(u * 2, l, mid, x);
	build(u * 2 + 1, mid + 1, r, x);
	pushup(u);
}
void update(int u, int l, int r, int L, int R, int x)
{
    if(InRange(l, r, L, R))
    {
        maketag(u, x, l, r);
        return ;
    }
    if(OutofRange(l, r, L, R))
    {
        return ;
    }
    int mid = l + r >> 1;
    pushdown(u, l, r);
    update(u * 2, l, mid, L, R, x);
    update(u * 2 + 1, mid + 1, r, L, R, x);
    pushup(u);
}
int query(int u, int l, int r, int L, int R)
{
    if(InRange(l, r, L, R))
    {
        return w[u].data;
    }
    if(OutofRange(l, r, L, R)) return -1e9;
    int mid = l + r >> 1;
    pushdown(u, l, r);
    return max(query(u * 2, l, mid, L, R), query(u * 2 + 1, mid + 1, r, L, R));
}
 
signed main()
{
	cin >> n >> k;
	for(int i = 1; i <= n; i ++)
	{
		cin >> a[i];
		pre[i] = sum[a[i]] + 1;
		sum[a[i]] = i;
	}
    memset(dp[0], -0x3f, sizeof(dp[0]));
    dp[0][0] = 0;
	for(int i = 1; i <= k; i ++)
	{
		memset(w, 0, sizeof(w));
		build(1, 1, n, i - 1);
		for(int j = 1; j <= n; j ++)
		{
			update(1, 1, n, pre[j], j, 1);
			dp[i][j] = query(1, 1, n, 1, j);
		}
	}
	cout << dp[k][n];
}
posted @ 2026-06-05 14:19  Luowj  阅读(7)  评论(0)    收藏  举报