决策单调性优化 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\),有:
时间复杂度 \(\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)=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\),有:
即所谓「交叉优于包含」。
满足四边形不等式是满足决策单调性的充分不必要条件,证明是简单的。
分治
考虑如何去计算 \(\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;
}

浙公网安备 33010602011771号