C++
 Computer >> コンピューター >  >> プログラミング >> C++

C++で「良い文字列」の総数を求めるアルゴリズムを解説

この記事では、動的計画法(DP)を使って「良い文字列(good strings)」の総数を効率的に求めるC++の手法を解説します。

問題の定義

長さnの2つの文字列s1とs2、そしてもう1つの文字列evilが与えられます。ここで「良い文字列」とは、以下の条件をすべて満たす文字列のことです。

  • 長さがnである
  • 辞書順でs1以上である
  • 辞書順でs2以下である
  • 部分文字列としてevilを含まない

答えは非常に大きな数になる可能性があるため、109 + 7で割った余りを返します。

入出力例

例として、n = 2、s1 = "bb"、s2 = "db"、evil = "a" の場合を考えてみましょう。このときの出力は51になります。内訳を見てみると、まずbで始まる良い文字列が25個("bb", "bc", "bd", ..., "bz")あり、次にcで始まる良い文字列が25個("cb", "cc", "cd", ..., "cz")あり、さらにdで始まるものとして"db"の1個が該当します。合計で51個となるわけです。

解法のアプローチ

この問題は、桁DP(数え上げDP)とオートマトンの遷移テーブルを組み合わせることで解けます。大まかな流れは以下の通りです。

  1. 定数N := 500、M := 50を定義する
  2. サイズ(N+1) × (M+1) × 2の配列dpを用意する
  3. サイズ(M+1) × 26の遷移テーブルtrを用意する
  4. モジュロ値m := 109 + 7を定義する

add関数(加算処理)

2つの値aとbを受け取り、モジュロ演算を適用しながら安全に足し合わせます。

((a mod m) + (b mod m)) mod m を返す

solve関数(コアロジック)

solve関数は引数としてn、s、eを受け取ります。主な処理ステップは以下の通りです。

  1. 文字列eを反転させる
  2. trとdpを0で初期化する
  3. eの各プレフィックスについて、26種類の文字を追加した際のKMP的な最長一致長を計算し、遷移テーブルtr[i][j]に格納する
  4. m := eの長さとする
  5. dp[n][0][1] := 1 で初期化する
  6. iをn-1から0まで減らしながら、各状態j・各文字k・境界フラグlについて遷移を行う
    • k > s[i] - 'a' の場合、nl := 0
    • k < s[i] - 'a' の場合、nl := 1
    • それ以外の場合、nl := l
    • dp[i][tr[j][k]][nl] に dp[i+1][j][l] を加算する
  7. ret := 0 とし、i = 0 から e.size() 未満まで dp[0][i][1] をretに加算して返す

メインメソッド(findGoodStrings)

  1. ok := 1 と初期化する
  2. s1がすべて'a'で構成されているかチェックする
  3. そうでない場合、s1を辞書順で1つ前の文字列にデクリメントする(繰り上がりの要領で'a'なら'z'に置き換える)
  4. left := (okが真なら0、そうでなければsolve(n, s1, evil))
  5. right := solve(n, s2, evil)
  6. (right - left + m) mod m を返す

このように「s2以下の良い文字列の数」から「s1未満の文字列の数」を引くことで、範囲[s1, s2]に含まれる良い文字列の総数を求めています。

C++による実装例

それでは、理解を深めるために実際の実装を見てみましょう。

#include <bits/stdc++.h>
using namespace std;
typedef long long int lli;
const int N = 500;
const int M = 50;
int dp[N + 1][M + 1][2];
int tr[M + 1][26];
const lli m = 1e9 + 7;
class Solution {
    public:
    int add(lli a, lli b){
        return ((a % m) + (b % m)) % m;
    }
    lli solve(int n, string s, string e){
        reverse(e.begin(), e.end());
        memset(tr, 0, sizeof(tr));
        memset(dp, 0, sizeof(dp));
        for (int i = 0; i < e.size(); i++) {
            string f = e.substr(0, i);
            for (int j = 0; j < 26; j++) {
                string ns = f + (char)(j + 'a');
                for (int k = i + 1;; k--) {
                    if (ns.substr(i + 1 - k) == e.substr(0, k)) {
                        tr[i][j] = k;
                        break;
                    }
                }
            }
        }
        int m = e.size();
        for (int i = 0; i <= n; i++) {
            for (int j = 0; j < m; j++) {
                dp[i][j][0] = dp[i][j][1] = 0;
            }
        }
        dp[n][0][1] = 1;
        for (int i = n - 1; i >= 0; i--) {
            for (int j = 0; j < e.size(); j++) {
                for (int k = 0; k < 26; k++) {
                    for (int l : { 0, 1 }) {
                        int nl;
                        if (k > s[i] - 'a') {
                            nl = 0;
                        }
                        else if (k < s[i] - 'a') {
                            nl = 1;
                        }
                        else
                           nl = l;
                        dp[i][tr[j][k]][nl] = add(dp[i][tr[j][k]]
                        [nl], dp[i + 1][j][l]);
                    }
                }
            }
        }
        lli ret = 0;
        for (int i = 0; i < e.size(); i++) {
            ret = add(ret, dp[0][i][1]);
        }
        return ret;
    }
    int findGoodStrings(int n, string s1, string s2, string evil) {
        bool ok = 1;
        for (int i = 0; i < s1.size() && ok; i++) {
            ok = s1[i] == 'a';
        }
        if (!ok) {
            for (int i = s1.size() - 1; i >= 0; i--) {
                if (s1[i] != 'a') {
                    s1[i]--;
                    break;
                }
                s1[i] = 'z';
            }
        }
        int left = ok ? 0 : solve(n, s1, evil);
        int right = solve(n, s2, evil);
        return (right - left + m) % m;
    }
};
main(){
    Solution ob;
    cout << (ob.findGoodStrings(2, "bb", "db", "a"));
}

入力

2, "bb", "db", "a"

出力

51

まとめ

この問題は、一見すると全列挙が必要に思えますが、桁DPの考え方とKMP法の失敗関数に似た遷移テーブルを組み合わせることで、O(n × |evil| × 26 × 2)程度の計算量で効率的に解くことができます。ポイントは以下の3つです。

  • 「s1未満」「s2以下」の数え上げを分けて行い、差分で範囲内の数を求める
  • evilとの一致状態をオートマトンの状態として持ち、evilが出現した経路は数えない
  • モジュロ演算を各加算時に適用し、オーバーフローを防ぐ

競技プログラミングや技術面接で頻出のパターンなので、ぜひマスターしておきましょう。

  1. C++で二分木内の重複するサブツリーをすべて検出する方法

    問題の概要二分木が与えられたとき、その中に重複するサブツリー(部分木)が存在するかどうかを判定する問題を考えてみましょう。例として、次のような二分木を取り上げます。この木には、サイズ2の同一のサブツリーが2つ存在します。さらに、それぞれのサブツリー内のDに注目すると、BDとBEもまた重複するサブツリーになっています。解決のアプローチ:木のシリアライズとハッシュこの問題は、木のシリアライズ(直列化)とハッシュテーブルを組み合わせることで効率的に解決できます。基本的な考え方は以下のとおりです。各サブツリーを間順走査(inorder traversal)で文字列としてシリアライズする空のノードには開

  2. 【C++入門】二次方程式のすべての解(根)を求めるプログラムの書き方

    二次方程式は一般に ax2 + bx + c = 0 の形で表されます。この方程式の解(根)は、以下に示す有名な「解の公式」によって求めることができます。判別式による3つの場合分け二次方程式の解の性質は、判別式 D = b2 − 4ac の値によって、次の3通りに分類されます。b2 < 4ac の場合:解は実数にならず、虚数を含む複素数になります。b2 = 4ac の場合:解は実数となり、両方の解が同じ値(重解)になります。b2 > 4ac の場合:解は実数となり、異なる2つの実数解を持ちます。それでは、これらすべての場合に対応した、二次方程式の解を求めるC++プログラムを見ていき