4
int solve(int i, int j, int extra, vector<pair<int, int>> &groups) {7
if (dp[i][j][extra] != -1) return dp[i][j][extra];9
int ans = (groups[i].second + extra) * (groups[i].second + extra) + solve(i + 1, j, 0, groups);11
for (int g = i + 1; g <= j; g++)12
if (groups[g].first == groups[i].first)14
solve(i + 1, g - 1, 0, groups) + solve(g, j, extra + groups[i].second, groups));16
return dp[i][j][extra] = ans;18
int removeBoxes(vector<int> &boxes) {21
vector<pair<int, int>> groups;22
for (int i = 0; i < n; i++) {24
while (i + 1 < n and boxes[i + 1] == boxes[j]) i++;25
groups.push_back({boxes[j], i - j + 1});28
memset(dp, -1, sizeof(dp));29
return solve(0, groups.size() - 1, 0, groups);