公式

B - ランプ列の分割スコア最大化 / Maximizing the Partition Score of a Lamp Sequence 解説 by MMNMM


この問題は、大きく \(2\) つの部分に分けることができます。

  • 操作を終えたあとの高橋君の装置の状態を求める
  • 最適な分割位置と、そのときの答えを求める

それぞれの部分に分けて説明します。


問題文中の操作は、最大でも \(\min\lbrace N,K\rbrace\) 回しか行われません。 最初の操作で右端に追加された点灯状態のランプが左端に到達したとき、必ず操作が終了するためです。

よって、操作を愚直に行っても十分高速になります。


操作後の高橋君の装置と青木君の列をあわせて、長さ \(N\) のランプ列が \(M+1\) 個あるとします。 分割する位置の候補は \(N-1\) 箇所あるので、これを全探索することで最適な分割位置とその答えを求めることができます。

時間計算量は \(O(N ^ 2M)\) となります。

時間計算量については、適切にビット演算を行うことで \(O(NM)\) としたり、左から \(i\) 番目の点灯しているランプの個数を集計しておくことで \(O(N ^ 2)\) に、加えて差分更新を適切に行うことで \(O(N)\) とできます。

実装例は以下のようになります。

#include <iostream>
#include <vector>
using namespace std;

int main() {
    int N, K, M;
    long X;
    cin >> N >> K >> M >> X;

    // 操作を行う
    for (int i = 0; i < K; ++i) {
        if (X & 1) break; // N 回以内に終了する
        X |= 1L << N; // 右端に追加して
        X /= 2; // 左端を削除する
    }

    // 「左から i 番目のランプのうち点灯しているものの個数」を集計しておく
    vector<long> bit_count(N);
    for (int i = 0; i < N; ++i) {
        bit_count[i] += (X >> i) & 1;
    }

    for (int i = 0; i < M; ++i) {
        long Y;
        cin >> Y;
        for (int j = 0; j < N; ++j) {
            bit_count[j] += (Y >> j) & 1;
        }
    }

    long ans = 0;
    for (int p = 1; p < N; ++p) { // 分割位置を全探索
        long tmp = 0;
        for (int i = 0; i < p; ++i) {
            tmp += bit_count[i] << i; // 左側
        }
        for (int i = p; i < N; ++i) {
            tmp += bit_count[i] << i - p; // 右側
        }
        ans = max(ans, tmp); // 最大値を求める
    }

    // 答えを出力
    cout << ans << endl;
    return 0;
}
N, K, M, X = map(int, input().split())

# 操作を行う
for i in range(K):
    if X % 2 == 1:
        break # N 回以内に終了する
    X |= 1 << N # 右端に追加して
    X //= 2 # 左端を削除する

# 「左から i 番目のランプのうち点灯しているものの個数」を集計しておく
bit_count = [0 for _ in range(N)]
for i in range(N):
    bit_count[i] += (X >> i) & 1

if M > 0:
    for Y in map(int, input().split()):
        for i in range(N):
            bit_count[i] += (Y >> i) & 1

ans = 0
for p in range(1, N): # 分割位置を全探索
    tmp = 0
    for i in range(p):
        tmp += bit_count[i] << i # 左側
    for i in range(p, N):
        tmp += bit_count[i] << i - p # 右側
    ans = max(ans, tmp) # 最大値を求める

# 答えを出力
print(ans)

投稿日時:
最終更新: