RT
这是我的代码
#include<bits/stdc++.h>
using namespace std;
#define int long long
const int M=17;
int n,m,r,c;
int a[M][M];
int S;
int sit[60005],js=0;
int s[M][M][20005],h[M][20005];
int f[M][M][20005];
signed main(){
cin>>n>>m>>r>>c;
S=(1<<m)-1;
for(int i=1;i<=n;i++){
for(int j=1;j<=m;j++){
scanf("%lld",&a[i][j]);
}
}
for(int i=0;i<=S;i++){
int sum=0;
for(int j=0;j<m;j++)if((i&(1<<j)))sum++;
if(sum==c)sit[++js]=i;
}
for(int i=1;i<=js;i++){
for(int j=1;j<=n;j++){
int last=0;
for(int k=0;k<m;k++){
if((sit[i]&(1<<k))){
if(last!=0)h[j][i]+=abs(a[j][k+1]-a[j][last]);
last=k+1;
}
}
for(int k=1;k<j;k++){
for(int l=0;l<m;l++){
if((sit[i]&(1<<l))){
s[k][j][i]+=abs(a[j][l+1]-a[k][l+1]);
}
}
}
}
}
memset(f,0x3f,sizeof f);
for(int i=1;i<=js;i++)f[1][1][i]=h[1][i],f[0][0][i]=0;
for(int k=1;k<=js;k++){
for(int i=2;i<=n;i++){
for(int j=0;j<=i&&j<=r;j++){
for(int l=0;l<i;l++){
f[i][j][k]=min(f[i][j][k],f[l][j-1][k]+h[i][k]+s[l][i][k]);
}
}
}
}
int ans=114514114514;
for(int i=1;i<=n;i++){
for(int j=1;j<=js;j++){
ans=min(ans,f[i][r][j]);
}
}
cout<<ans;
return 0;
}