求助60pts
查看原帖
求助60pts
236416
_stOrz_楼主2023/9/8 16:25
#include <bits/stdc++.h>
#define int long long
using namespace std;

const int N = 10;

int a[705][20], v[1 << 20], cnt[1 << 20], st, num;
unordered_map <int, int> mp;

#define pii pair <int, int>
#define fi first
#define se second

pii sta[1 << 19];

int ans = -1e18, n, m;

void Init () {
  for (int i = 1; i <= num; i ++)
    v[i] = -1e18;
  num = 0; mp.clear ();
  for (int i = 1; i <= st; i ++) {
    int x = sta[i].fi, val = sta[i].se;
    if (mp[x] == 0) mp[x] = ++ num;
    v[mp[x]] = max (val, v[mp[x]]), cnt[mp[x]] = x;
  }
  st = 0; 
}

void Insert (int x, int val) {
  sta[++ st] = {x, val};
}

void DP () {
  
  cnt[1] = 0; num = 1; v[1] = 0;
  for (int i = 1; i <= n; i ++) {
    for (int j = 1; j <= num; j ++)
      cnt[j] <<= 2;
    for (int j = 1; j <= m; j ++) {
      for (int k = 1; k <= num; k ++) {
        int x = cnt[k], val = v[k];
        int u = ((x >> (2 * j)) & 3), r = ((x >> (2 * (j - 1))) & 3);
        if (!u and !r) {
          Insert (x, val);
          if (i != n and j != m) Insert (x + (1 << (2 * (j - 1))) + (2 << (2 * j)), val + a[i][j]);
        } 
        else if (!u and r) {
          if (i != n) Insert (x, val + a[i][j]);
          if (j != m) Insert (x - (r << (2 * (j - 1))) + (r << (2 * j)), val + a[i][j]);
        }
        else if (u and !r) {
          if (j != m) Insert (x, val + a[i][j]);
          if (i != n) Insert (x - (u << (2 * j)) + (u << (2 * (j - 1))), val + a[i][j]);
        }
        else if (u == 1 and r == 1) {
          int ur = 0, count = 0;
          for (int h = j; h <= m; h ++) {
            if (((x >> (2 * h)) & 3) == 1) count ++;
            else if (((x >> (2 * h)) & 3) == 2) count --;
            if (count == 0) { ur = h; break; }
          }
          Insert (x - (1 << (2 * j)) - (1 << (2 * (j - 1))) - (1 << (2 * ur)), val + a[i][j]); 
        }
        else if (u == 2 and r == 2) {
          int rl = 0, count = 0;
          for (int h = j - 1; h >= 0; h --) {
            if (((x >> (2 * h)) & 3) == 2) count ++;
            else if (((x >> (2 * h)) & 3) == 1) count --;
            if (count == 0) { rl = h; break; }
          }
          Insert (x - (2 << (2 * j)) - (2 << (2 * (j - 1))) + (1 << (2 * rl)), val + a[i][j]);
        }
        else if (u == 1 and r == 2) {
          Insert (x - (1 << (2 * j)) - (2 << (2 * (j - 1))), val + a[i][j]);
        }
        else if (x == ((2 << (2 * j)) + (1 << (2 * (j - 1)))))
          ans = max (ans, val + a[i][j]);
      }
      Init ();
    }
  }
}

signed main () {
  cin >> n >> m;
  for (int i = 1; i <= n; i ++) {
    for (int j = 1; j <= m; j ++)
      cin >> a[i][j];
  }
  
  DP ();
  cout << ans << "\n";
}
2023/9/8 16:25
加载中...