Official

D - Coefficient Stair Editorial by MMNMM


辞書順で全列挙を行う際には、深さ優先探索が便利です。

条件を満たす列のうち先頭 \(i\) 項が \(A _ 1,A _ 2,\ldots,A _ i\) であるものを辞書順にすべて列挙する関数として、次のような再帰関数を考えることができます。

  1. \(i=N\) のとき(つまり、先頭 \(N\) 項が決まっている場合)
    1. \(\displaystyle\sum _ {j=1} ^ Nj\times A _ j=K\) なら、\(A\) を出力する。
    2. そうでなければ、何も出力しない。
  2. そうでない場合、\(i+1\) 項目として \(0\) から \(\Biggl\lfloor\dfrac{K-\sum _ {j=1} ^ ij\times A _ j}{i+1}\Biggr\rfloor\) までの整数を順に試して再帰を行う。

Python では以下のような関数になります。

def dfs(i, A):
    s = sum((j + 1) * a for j, a in enumerate(A[:i]))
    if i == N:
        if s == K:
            print(*A)
        return
    for x in range(0, (K - s) // (i + 1) + 1):
        A[i] = x
        dfs(i + 1, A)

この関数を dfs(0, [0] * N) などで呼び出すことで正しい結果は得られますが、このまま実行時間制限に間に合わせるのは難しいです。

\(i=N\) であるような呼び出しは、非負整数からなる長さ \(N\) の列 \(A\) のうち \(\displaystyle\sum _ {j=1} ^ Nj\times A _ j\le K\) であるものごとに行われます。 非負整数からなる長さ \(N\) の列 \(A\) のうち \(\displaystyle\sum _ {j=1} ^ Nj\times A _ j=K\) であるものの個数を \(f(N,K)\) とすると、これは \(\displaystyle\sum _ {k=0} ^ Kf(N,k)\) であり、今回の制約のもとで \(10000100000\) まで大きくなってしまいます(\(N=2,K=2\times10 ^ 5\) のとき)。

ここで、再帰関数を以下のように変更してみます。

  1. \(i=N-1\) のとき(つまり、先頭 \(N-1\) 項が決まっている場合)
    1. \(A _ N\) をうまく決めることで \(\displaystyle\sum _ {j=1} ^ Nj\times A _ j=K\) とできるなら、\(A\) を出力する。
    2. そうでなければ、何も出力しない。
  2. そうでない場合、\(i+1\) 項目として \(0\) から \(\Biggl\lfloor\dfrac{K-\sum _ {j=1} ^ ij\times A _ j}{i+1}\Biggr\rfloor\) までの整数を順に試して再帰を行う。

Python では以下のような関数になります。

def dfs(i, A):
    s = sum((j + 1) * a for j, a in enumerate(A[:i]))
    if i == N - 1:
        if (K - s) % N == 0:
            A[i] = (K - s) // N
            print(*A)
        return
    for x in range(0, (K - s) // (i + 1) + 1):
        A[i] = x
        dfs(i + 1, A)

すると、今度は \(i=N-1\) での呼び出し回数は \(O(Nf(N,K))\) 回となることが示せます。 よって、全体の呼び出し回数は \(O(N ^ 2f(N,K))\) 回となり、s の計算に毎回 \(O(N)\) 時間かけているため全体で \(O(N ^ 3f(N,K))\) 時間となります。 言語によってはこのままでも実行時間制限に間に合わせることができます。 s を再帰関数の引数に加えることで全体で \(O(N ^ 2f(N,K))\) 時間とでき、これで十分高速です。

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

#include <iostream>
#include <vector>

int main() {
    using namespace std;
    int N, K;
    cin >> N >> K;

    vector<int> path(N);
    [N, K](this auto self, int index, int sum, vector<int>& path) -> void {
        if (index + 1 == N) { // 最後の 1 項になったら
            if ((K - sum) % N == 0) { // 合計を K にできるときのみ
                path.back() = (K - sum) / N;
                for (int p : path) // 出力
                    cout << p << ' ';
                cout << '\n';
            }
            return;
        } 

        // i 項目を決めて再帰
        for (int i = 0; sum + i * (index + 1) <= K; ++i) {
            path[index] = i;
            self(index + 1, sum + i * (index + 1), path);
        }
    }(0, 0, path);
    
    return 0;
}
N, K = map(int, input().split())

def dfs(i, s, A):
    if i == N - 1: # 最後の 1 項になったら
        if (K - s) % N == 0: # 合計を K にできるときのみ
            A[i] = (K - s) // N
            print(' '.join(str(a) for a in A)) # 出力
        return

    # i 項目を決めて再帰
    for x in range(0, (K - s) // (i + 1) + 1):
        A[i] = x
        dfs(i + 1, s + x * (i + 1), A)

dfs(0, 0, [0] * N)

posted:
last update: