决策单调性优化 DP / 题解:Ciel and Gondolas

题目传送门

朴素 DP

设 \(f_{i,j}\) 表示前 \(i\) 个人划分为 \(j\) 组的最小答案。

记 \(a_{l,r}\) 表示 \([l,r]\) 划分为一个组的贡献,简单计算可以得到是对 \(u\) 的左上角为 \((l,l)\)、右下角为 \((r,r)\) 的子矩阵求和再除以 \(2\)。可以 \(\mathcal O\left(n^2\right)\) 预处理。

钦定 \(f_{0,0}=0\),有:

\[f_{i,j}=\min_{l\in[0,\min(i-1,k)]}\left(f_{l,j-1}+a_{l+1,i}\right) \]

时间复杂度 \(\mathcal O\left(kn^2\right)\)。

参考代码
//#include<bits/stdc++.h>
#include<algorithm>
#include<iostream>
#include<cstring>
#include<iomanip>
#include<cstdio>
#include<string>
#include<vector>
#include<cmath>
#include<ctime>
#include<deque>
#include<queue>
#include<stack>
#include<list>
using namespace std;
constexpr const int N=4000,K=800,inf=0x3f3f3f3f;
int n,k,u[N+1][N+1],a[N+1][N+1],f[N+1][N+1];
int main(){
	/*freopen("test.in","r",stdin);
	freopen("test.out","w",stdout);*/
	
	ios::sync_with_stdio(false);
	cin.tie(0);cout.tie(0);
	
	cin>>n>>k;
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			cin>>u[i][j];
			u[i][j]+=u[i][j-1];
		}
	}
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			u[i][j]+=u[i-1][j];
		}
	}
	for(int l=1;l<=n;l++){
		for(int r=l;r<=n;r++){
			a[l][r]=(u[r][r]-u[l-1][r]-u[r][l-1]+u[l-1][l-1])>>1;
		}
	}
	memset(f,0x3f,sizeof(f));
	f[0][0]=0;
	for(int i=1;i<=n;i++){
		for(int j=1;j<=i;j++){
			for(int l=0;l<i;l++){
				f[i][j]=min(f[i][j],f[l][j-1]+a[l+1][i]);
			}
		}
	}
	cout<<f[n][k]<<'\n';
	
	cout.flush();
	
	/*fclose(stdin);
	fclose(stdout);*/
	return 0;
}

决策单调性优化 DP

考虑一下 \(f_{i,j}\) 只与 \(f_{l,j-1}\) 和 \(a\) 相关,考虑暴力枚举 \(j\),将问题降为了一个一维问题。

形式化地,决策单调性优化 DP 在求解最优性问题:

\[f(i)=\min_{j\leq i}w(j,i) \]

当 \(f(i)=w(j_0,i)\) 时,我们称 \(f(i)\) 的决策点 \(\operatorname{op}(i)=j_0\)。

决策单调性即对于 \(i<j\),有 \(\operatorname{op}(i)\leq\operatorname{op}(j)\)。

四边形不等式

决策单调性是否存在可以通过一些方法判断,常用的有四边形不等式。即对于 \(a\leq b\leq c\leq d\),有:

\[w(a,c)+w(b,d)\leq w(a,d)+w(b,c) \]

即所谓「交叉优于包含」。

满足四边形不等式是满足决策单调性的充分不必要条件,证明是简单的。

分治

考虑如何去计算 \(\operatorname{op}(1),\operatorname{op}(2),\cdots,\operatorname{op}(n)\),从而优化 DP 计算。

我们想求解 \([l,r]\) 内的所有点的决策点,设区间内的决策点的区间为 \([\textit{opL},\textit{opR}]\)。

考虑分治,记 \(\textit{mid}=\dfrac{l+r}2\) 为区间中点,先暴力在 \([\textit{opL},\min(\textit{opR},\textit{mid}-1)]\) 内算出 \(\operatorname{op}(\textit{mid})\)。

因为决策单调性,所以 \(\forall i\in[l,\textit{mid}),\operatorname{op}(i)\leq\operatorname{op}(\textit{mid})\),\(\forall i\in(\textit{mid},r],\operatorname{op}(\textit{mid})\leq\operatorname{op}(i)\)。

一共有 \(\mathcal O(\log n)\) 层,单层计算 DP 时的决策点区间互不相交,因此单层复杂度为 \(\mathcal O(n)\)。总时间复杂度 \(\mathcal O(n\log n)\)。

void solve(int l,int r,int opL,int opR,int j){
	int mid=l+r>>1;
	for(int i=opL;i<=min(opR,mid-1);i++){
        //计算 DP
	}
	if(l<=mid-1){
		solve(l,mid-1,opL,op[mid],j);
	}
	if(mid+1<=r){
		solve(mid+1,r,op[mid],opR,j);
	}
}

AC 代码

时间复杂度 \(\mathcal O(n^2+kn\log n)\)。

参考代码
//#include<bits/stdc++.h>
#include<algorithm>
#include<iostream>
#include<cstring>
#include<iomanip>
#include<cstdio>
#include<string>
#include<vector>
#include<cmath>
#include<ctime>
#include<deque>
#include<queue>
#include<stack>
#include<list>
using namespace std;
constexpr const int N=4000,K=800,inf=0x3f3f3f3f;
int n,k,u[N+1][N+1],a[N+1][N+1],f[N+1][N+1];
int op[N+1];
void solve(int l,int r,int opL,int opR,int j){
	int mid=l+r>>1;
	for(int i=opL;i<=min(opR,mid-1);i++){
		if(f[i][j-1]+a[i+1][mid]<f[mid][j]){
			f[mid][j]=f[i][j-1]+a[i+1][mid];
			op[mid]=i;
		}
	}
	if(l<=mid-1){
		solve(l,mid-1,opL,op[mid],j);
	}
	if(mid+1<=r){
		solve(mid+1,r,op[mid],opR,j);
	}
}
int main(){
	/*freopen("test.in","r",stdin);
	freopen("test.out","w",stdout);*/
	
	ios::sync_with_stdio(false);
	cin.tie(0);cout.tie(0);
	
	cin>>n>>k;
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			cin>>u[i][j];
			u[i][j]+=u[i][j-1];
		}
	}
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			u[i][j]+=u[i-1][j];
		}
	}
	for(int l=1;l<=n;l++){
		for(int r=l;r<=n;r++){
			a[l][r]=(u[r][r]-u[l-1][r]-u[r][l-1]+u[l-1][l-1])>>1;
		}
	}
	memset(f,0x3f,sizeof(f));
	f[0][0]=0;
	for(int j=1;j<=k;j++){
		solve(1,n,0,n,j);
	}
	cout<<f[n][k]<<'\n';
	
	cout.flush();
	
	/*fclose(stdin);
	fclose(stdout);*/
	return 0;
}
卡常代码

加入了 AIGC 的快读。

//#include<bits/stdc++.h>
#include<algorithm>
#include<iostream>
#include<cstring>
#include<iomanip>
#include<cstdio>
#include<string>
#include<vector>
#include<cmath>
#include<ctime>
#include<deque>
#include<queue>
#include<stack>
#include<list>
using namespace std;

// ???????????“??”??(?? bool ?????,??????)
template <class T>
using EnableInt = std::enable_if_t<
    std::is_integral<T>::value &&
    !std::is_same<T, bool>::value &&
    !std::is_same<T, char>::value &&
    !std::is_same<T, signed char>::value &&
    !std::is_same<T, unsigned char>::value, int>;

// ============================ ?? ============================
class FastInput {
public:
    explicit FastInput(FILE* file = stdin) : fin_(file) {
        pin_ = pend_ = buf_;
        next();                       // ???????
    }

    FastInput(const FastInput&) = delete;
    FastInput& operator=(const FastInput&) = delete;

    // ????
    template <class T, EnableInt<T> = 0>
    FastInput& operator>>(T& x) {
        x = read_int<T>();
        return *this;
    }

    // ????(????)
    FastInput& operator>>(char& c) {
        while (cur_ != -1 && cur_ <= ' ') next();
        c = static_cast<char>(cur_);
        next();
        return *this;
    }

private:
    static constexpr size_t kBufSize = 1 << 20;   // 1MB

    FILE* fin_;
    char  buf_[kBufSize];
    char *pin_, *pend_;
    int   cur_;                                   // ????,-1 ?? EOF

    void next() {
        if (pin_ == pend_) {
            size_t n = fread(buf_, 1, kBufSize, fin_);
            pin_ = buf_;
            pend_ = buf_ + n;
            if (n == 0) { cur_ = -1; return; }
        }
        cur_ = static_cast<unsigned char>(*pin_++);
    }

    template <class T>
    T read_int() {
        // ??????????????(??????)
        while (cur_ != -1 && (cur_ < '0' || cur_ > '9') && cur_ != '-') next();

        using U = std::make_unsigned_t<T>;
        bool neg = false;
        if (cur_ == '-') { neg = true; next(); }

        U val = 0;
        while (cur_ >= '0' && cur_ <= '9') {
            val = val * 10 + static_cast<U>(cur_ - '0');
            next();
        }
        // ?????????,?? INT_MIN ????? UB
        return neg ? static_cast<T>(static_cast<U>(0) - val)
                   : static_cast<T>(val);
    }
};

// ============================ ?? ============================
class FastOutput {
public:
    explicit FastOutput(FILE* file = stdout) : fout_(file) { pout_ = buf_; }
    ~FastOutput() { flush(); }

    FastOutput(const FastOutput&) = delete;
    FastOutput& operator=(const FastOutput&) = delete;

    // ????
    template <class T, EnableInt<T> = 0>
    FastOutput& operator<<(T x) {
        write_int(x);
        return *this;
    }

    FastOutput& operator<<(char c) {
        put(c);
        return *this;
    }

    FastOutput& operator<<(const char* s) {
        while (*s) put(*s++);
        return *this;
    }

    void flush() {
        if (pout_ != buf_) {
            fwrite(buf_, 1, static_cast<size_t>(pout_ - buf_), fout_);
            pout_ = buf_;
        }
    }

private:
    static constexpr size_t kBufSize = 1 << 20;

    FILE* fout_;
    char  buf_[kBufSize];
    char* pout_;

    void put(char c) {
        if (pout_ == buf_ + kBufSize) flush();
        *pout_++ = c;
    }

    template <class T>
    void write_int(T x) {
        if (x == 0) { put('0'); return; }

        using U = std::make_unsigned_t<T>;
        bool neg = false;
        U u;
        if constexpr (std::is_signed<T>::value) {
            if (x < 0) { neg = true; u = static_cast<U>(0) - static_cast<U>(x); }
            else       { u = static_cast<U>(x); }
        } else {
            u = static_cast<U>(x);
        }

        char tmp[24];                 // uint64 ?? 20 ? + ??
        int  len = 0;
        while (u) { tmp[len++] = static_cast<char>('0' + u % 10); u /= 10; }

        if (neg) put('-');
        while (len > 0) put(tmp[--len]);
    }
};
#define cin in
#define cout out

constexpr const int N=4000,K=800,inf=0x3f3f3f3f;
int n,k,u[N+1][N+1],a[N+1][N+1],f[K+1][N+1];
int op[N+1];
void solve(int l,int r,int opL,int opR,int j){
	int mid=l+r>>1;
	for(int i=opL;i<=min(opR,mid-1);i++){
		if(f[j-1][i]+a[i+1][mid]<f[j][mid]){
			f[j][mid]=f[j-1][i]+a[i+1][mid];
			op[mid]=i;
		}
	}
	if(l<=mid-1){
		solve(l,mid-1,opL,op[mid],j);
	}
	if(mid+1<=r){
		solve(mid+1,r,op[mid],opR,j);
	}
}
int main(){
	/*freopen("test.in","r",stdin);
	freopen("test.out","w",stdout);*/
	
//	ios::sync_with_stdio(false);
//	cin.tie(0);cout.tie(0);
	
	FastInput  in(stdin);
    FastOutput out(stdout);
	
	cin>>n>>k;
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			cin>>u[i][j];
			u[i][j]+=u[i][j-1];
		}
	}
	for(int i=1;i<=n;i++){
		for(int j=1;j<=n;j++){
			u[i][j]+=u[i-1][j];
		}
	}
	for(int l=1;l<=n;l++){
		for(int r=l;r<=n;r++){
			a[l][r]=(u[r][r]-u[l-1][r]-u[r][l-1]+u[l-1][l-1])>>1;
		}
	}
	memset(f,0x3f,sizeof(f));
	f[0][0]=0;
	for(int j=1;j<=k;j++){
		solve(1,n,0,n,j);
	}
	cout<<f[k][n]<<'\n';
	
	cout.flush();
	
	/*fclose(stdin);
	fclose(stdout);*/
	return 0;
}
posted @ 2026-09-17 16:56  TH911  阅读(6)  评论(0)    收藏  举报