40 分代码
#include <bits/stdc++.h>
#define int long long
using namespace std;
int n, m, f[1000005], w[1000005], before[1000005], last[1000005], s[3000005], h[3000005], ans;
void print() {
for (int i = 1; i <= 2 * n - 1; i++)
cout << s[i] << "\t" << h[i];
cout << endl;
}
int jisuan(int id, int l, int r, int R, int x)
{
if (R < l) return 0;
int mid = ((l + r) >> 1);
if (l == r)
{
if (R >= r)
{
return x + h[id];
}
return 0;
}
if (R > r)
{
return x + h[id];
}
if (R <= mid)
{
return jisuan(id * 2, l, mid, R, x);
}
if (R > mid && R <= r)
{
return max(jisuan(id * 2, l, mid, R, x + s[id * 2 + 1]),jisuan(id * 2 + 1, mid + 1, r, R, x));
}
}
int add(int id, int l, int r, int x, int k) {
if (l == x && r == x)
{
s[id] += k;
h[id] = max(0ll,s[id]);
return h[id];
}
h[id] = 0;
if (x <= (l + r >> 1)) h[id] = max(h[id], add(id * 2, l, l + r >> 1, x, k) + s[id * 2 + 1]);
if (x >= ((l + r >> 1) + 1)) h[id] = max(h[id], add(id * 2 + 1, (l + r >> 1) + 1, r, x, k));
s[id] = s[id * 2] + s[id * 2 + 1];
return h[id];
}
signed main() {
cin >> n >> m;
for (int i = 1; i <= n; i++)
{
cin >> f[i];
}
for (int i = 1; i <= m; i++)
{
cin >> w[i];
}
for (int i = 1; i <= n; i++)
{
before[i] = last[f[i]];
last[f[i]] = i;
}
register const int N = n;
n = 1 << (int)(ceil(log2(n)));
for (int i = 1; i <= N; i++)
{
add(1,1,n,i,w[f[i]]);
if (before[i])
{
add(1,1,n,before[i],-2 * w[f[i]]);
if (before[before[i]])
{
add(1,1,n,before[before[i]],w[f[i]]);
}
}
ans = max(ans,jisuan(1,1,n,i,0));
}
cout << ans;
return 0;
}