邮局加强。跑太慢了,求优化,说好的 2.75s??? 直接挂成 5.45s...
#include <bits/stdc++.h>
using namespace std;
#define ll long long
#define rint register ll
ll N, M;
ll ans;
ll arr[500005], sum[500005], cnt[500005], f[500005];
struct pp{
ll u, v;
} q[500005];
ll calc(ll i, ll j)
{
ll mid = (i + j + 1) >> 1;
return (sum[j] - sum[mid]) - arr[mid] * (j - mid) + arr[mid] * (mid - i) - (sum[mid] - sum[i]);
}
ll binary_search(ll i, ll j)
{
ll l = j, r = N + 1;
while(l < r - 1)
{
ll mid = (l + r) >> 1;
if(f[i] + calc(i, mid) < f[j] + calc(j, mid)) l = mid;
else r = mid;
}
return r;
}
bool check(ll x)
{
ll head = 1, tail = 0;
q[++ tail] = {0, N + 1};
for(rint i = 1; i <= N; ++ i)
{
while(head < tail && q[head].v <= i) ++ head;
f[i] = f[q[head].u] + calc(q[head].u, i) + x;
cnt[i] = cnt[q[head].u] + 1;
while(head < tail && binary_search(q[tail].u, i) <= q[tail - 1].v) -- tail;
q[tail].v = binary_search(q[tail].u, i);
q[++ tail] = {i, N + 1};
}
return cnt[N] >= M;
}
signed main()
{
scanf("%lld %lld", &N, &M);
for(rint i = 1; i <= N; ++ i)
scanf("%lld", &arr[i]);
sort(arr + 1, arr + N + 1);
for(rint i = 1; i <= N; ++ i)
sum[i] = sum[i - 1] + arr[i];
ll l = 0, r = 1e9 + 5;
while(l < r - 1)
{
ll mid = (l + r) >> 1;
if(check(mid)) l = mid;
else r = mid;
}
check(l);
cout << f[N] - M * l << endl;
return 0;
}
#include <bits/stdc++.h>
using namespace std;
#define ll long long
#define rint register ll
ll N, M;
ll ans;
ll arr[500005], sum[500005], cnt[500005], f[500005];
struct pp{
ll u, v;
} q[500005];
ll calc(ll i, ll j)
{
ll mid = (i + j + 1) >> 1;
return (sum[j] - sum[mid]) - arr[mid] * (j - mid) + arr[mid] * (mid - i) - (sum[mid] - sum[i]);
}
ll binary_search(ll i, ll j)
{
ll l = j, r = N + 1;
while(l < r - 1)
{
ll mid = (l + r) >> 1;
if(f[i] + calc(i, mid) < f[j] + calc(j, mid)) l = mid;
else r = mid;
}
return r;
}
bool check(ll x)
{
ll head = 1, tail = 0;
q[++ tail] = {0, N + 1};
for(rint i = 1; i <= N; ++ i)
{
while(head < tail && q[head].v <= i) ++ head;
f[i] = f[q[head].u] + calc(q[head].u, i) + x;
cnt[i] = cnt[q[head].u] + 1;
while(head < tail && binary_search(q[tail].u, i) <= q[tail - 1].v) -- tail;
q[tail].v = binary_search(q[tail].u, i);
q[++ tail] = {i, N + 1};
}
return cnt[N] >= M;
}
signed main()
{
scanf("%lld %lld", &N, &M);
for(rint i = 1; i <= N; ++ i)
scanf("%lld", &arr[i]);
sort(arr + 1, arr + N + 1);
for(rint i = 1; i <= N; ++ i)
sum[i] = sum[i - 1] + arr[i];
ll l = 0, r = 1e9 + 5;
while(l < r - 1)
{
ll mid = (l + r) >> 1;
if(check(mid)) l = mid;
else r = mid;
}
check(l);
cout << f[N] - M * l << endl;
return 0;
}