SDOI 2016 征途 (斜率优化)

· · 题解

题目描述:

Pine开始了从S地到T地的征途。

从S地到T地的路可以划分成n段,相邻两段路的分界点设有休息站。

Pine计划用m天到达T地。除第m天外,每一天晚上Pine都必须在休息站过夜。所以,一段路必须在同一天中走完。

Pine希望每一天走的路长度尽可能相近,所以他希望每一天走的路的长度的方差尽可能小。

帮助Pine求出最小方差是多少。

设方差是v,可以证明, v\times m^2是一个整数。为了避免精度误差,输出结果时输出 v\times m^2

输入输出格式

输入格式:

第一行两个数 nm

第二行 n 个数,表示 n 段路的长度

输出格式:

一个数,v\times m^2

题解

拿到题之后没有思路,于是先来推一波式子。

最原始的方差表达形式为:

s^2 = \frac{(\bar{v} - v_1)^2 + (\bar{v} - v_2)^2 + ... + (\bar{v} - v_m)^2}{m}

整理得到:

s^2 = \frac{(\sum_{i = 1} ^ m v_i - v_1)^2 + (\sum_{i = 1} ^ m v_i - v_2)^2 + ... + (\sum_{i = 1} ^ m v_i - v_m)^2}{m}

二项式展开之后可得:

s^2 = \frac{m\times\frac{(\sum_{i = 1} ^ m v_i)^2}{m ^ 2} - 2\times\frac{\sum_{i = 1} ^ m v_i}{m}\times\sum_{i= 1}^m v_i^2 + \sum_{i = 1} ^m v_i ^ 2}{m}

最后整理得到最终形式:

s ^ 2 = \frac{-\frac{\sum_{i = 1} ^ m v_i ^2}{m} +\sum_{i = 1} ^ m v_i ^ 2}{m}

最后将m ^ 2代入可得:

s ^ 2 \times m ^ 2 = -(\sum_{i = 1} ^ m v_i) ^ 2 + m \times \sum_{i = 1} ^ m v_i ^ 2

此时可以发现前一项是一个常数,扔掉不理他→_→

那么问题转化为如何求\sum_{i = 1} ^ m v_i ^ 2

很明显可以有很朴素的状态f[i][j]表示考虑到第i段,走了j天时的总量。

然鹅这样的DP是O(n ^ 3)的,所以考虑斜率优化。

为了方便,我们设计f[i]g[i]为当前考虑的时间段和之前的时间段的状态,那么我们就可以愉快地再大力推一波式子了→_→:

f[i] = g[i] + (sum[i] - sum[j]) ^ 2

一顿操作后可以得到斜率式的形式:

2 \times sum[i] \times sum[j] + f[i] - sum[i] ^ 2 = g[j] + sum[j] ^ 2

然后套斜率优化的套路就可以惹。

另外需要注意初始化g[i]

代码

#include <iostream>
#include <cstdio>
#include <cstring>
#include <algorithm>

#define MAXN 3005

typedef long long ll;

using namespace std;

int N, M;
int head, tail;
int q[MAXN];
ll sum[MAXN];
ll f[MAXN], g[MAXN];

inline ll read_int()
{
    ll ret = 0, f = 1; char c = getchar();
    while(c < '0' || c > '9') {if(c == '-') f = -1; c = getchar();}
    while(c >= '0' && c <= '9') {ret = (ret << 1) + (ret << 3) + int(c - 48); c = getchar();}
    return ret * f;
}

inline ll pow(ll x) {return x * x;}

inline double X(int i) {return sum[i];}

inline double Y(int j) {return g[j] + pow(sum[j]);}

inline double slope(int i, int j) {return (Y(i) - Y(j)) / (X(i) - X(j));}

void init()
{
    N = read_int(), M = read_int();
    for(int i = 1; i <= N; i++)
    {
        sum[i] = read_int();
        sum[i] += sum[i - 1];
        g[i] = pow(sum[i]);
    }
}

void dp()
{
    for(int l = 1; l < M; l++)
    {
        head = 1, tail = 1;
        q[1] = l;
        for(int i = l + 1; i <= N; i++)
        {
            while(head < tail && slope(q[head], q[head + 1]) < 2 * sum[i])
                head++;
            f[i] = g[q[head]] + (pow(sum[i] - sum[q[head]]));
            while(head < tail && slope(q[tail], q[tail - 1]) > slope(q[tail], i))
                tail--;
            q[++tail] = i;
        }
        memcpy(g, f, sizeof(f));
    }
    printf("%lld\n", f[N] * M - pow(sum[N]));
}

int main()
{
    init();
    dp();
    return 0;
}