代码:
#include<bits/stdc++.h>
using namespace std;
int n,m,r[20],h[20],mins=1e9;
int S(){
int s=0;
for(int i=1;i<=n;i++){
s+=r[i]*r[i]-r[i+1]*r[i+1]+2*r[i]*h[i];
}
return s;
}
void dfs(int now,int rr,int hh,int s){
if(s<(m-now+2)*(m-now+1)/2) return;
if(now>m){
if(s==0) mins=min(mins,S());
return;
}
for(int i=m-now+1;i<=rr;i++){
for(int j=m-now+1;j<=rr;j++){
r[now]=i;
h[now]=j;
dfs(now+1,i-1,j-1,s-i*i*j);
}
}
}
int main(){
scanf("%d%d",&n,&m);
int p=sqrt(n);
dfs(1,p,n,n);
if(mins==1e9) printf("0");
else printf("%d",mins);
return 0;
}