構造化出力の制御(JSON・関数呼び出し) — LLMの出力をJSON等に制約するConstrained Decodingの理論と実装

ChatGPTに「ユーザー情報をJSONで返して」と頼むと、多くの場合それらしいJSONが返ってきます。しかし、「多くの場合」では困る場面があります。たとえばAPIサーバーのバックエンドでLLMの出力を直接パースするとき、たった一度でもフォーマットが崩れれば、システム全体が json.JSONDecodeError で停止してしまいます。余計な説明文が付いてきたり、カンマが欠けていたり、フィールド名が微妙に違っていたり — 言語モデルの出力は本質的に確率的であり、100%の形式保証はプロンプトだけでは得られません。

ここで登場するのが構造化出力(Structured Output)の技術です。これは、LLMの生成プロセスそのものに文法的制約を埋め込むことで、出力が必ず指定されたスキーマに従うことを保証するアプローチです。

構造化出力の技術を理解すると、以下のような場面で大きな力を発揮します。

  • LLMを組み込んだプロダクション・システム: API応答をJSON Schemaで厳密に制約し、後段のパーサーが常に正しく動作することを保証できます。データパイプラインの信頼性が飛躍的に向上します
  • AIエージェントの構築: Function Calling / Tool Useの仕組みを理論レベルで理解することで、エージェントが「どの関数を呼ぶか」「どんな引数を渡すか」を構造化された形式で確実に出力できる理由が明確になります
  • 情報抽出・データ構造化: 非構造化テキストから型安全な構造化データを抽出する際に、出力形式のバリデーションをデコード時に完結させることで、後処理のリトライロジックを排除できます

本記事の内容

  • プロンプトベースの構造化出力とその限界
  • Constrained Decoding の基本原理 — 文法ベースのトークンマスキング
  • 正規表現 → 有限状態マシン(FSM) → トークンマスクへの変換
  • 制約付き確率分布の数学的定式化
  • JSON Schema Guided Generation の仕組み
  • Function Calling / Tool Use と構造化出力の関係
  • Pythonでの簡易Constrained Decoding実装
  • 性能トレードオフと実用上の注意点

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

画像なし
因果言語モデル(CLM)の理論と実装
GPTを支える自己回帰型の事前学習の数理と、次トークン予測の仕組みを解説します。
画像なし
Temperature・Top-k・Top-pサンプリングを比較して理解する
LLMのテキスト生成を制御するサンプリング手法の数式と実装を解説します。
画像なし
ビームサーチとは?テキスト生成の探索アルゴリズムを解説
決定論的なテキスト生成手法であるビームサーチの仕組みを解説します。
画像なし
【LLM】Tool Use / Function Callingの仕組みと実装
LLMが外部APIやツールを呼び出す技術の原理からPython実装まで解説します。

プロンプトベースの構造化出力とその限界

プロンプトで形式を指定するアプローチ

構造化出力を得る最もシンプルな方法は、プロンプトに「JSONで返してください」と書くことです。実際に多くのアプリケーションが、次のようなプロンプトエンジニアリングで構造化出力を実現しています。

あなたは情報抽出アシスタントです。
以下のテキストから人物情報を抽出し、次のJSON形式で返してください。

{
  "name": "氏名",
  "age": 年齢(数値),
  "occupation": "職業"
}

JSONのみを出力し、余計な説明は付けないでください。

最新のLLM(GPT-4、Claude 3.5、Gemini 1.5など)は、このような指示に対して高い精度でJSONを返します。しかし、「高い精度」と「100%の保証」の間には決定的な差があります。

なぜプロンプトだけでは不十分なのか

プロンプトベースのアプローチが根本的に抱える問題は、LLMの生成メカニズムに起因します。因果言語モデル(CLM)で解説したとおり、LLMは各ステップで語彙 $\mathcal{V}$ 上の確率分布を出力し、その分布からトークンを選択します。

$$ P(y_t \mid y_{

この確率分布は語彙中の全てのトークンに対して非ゼロの確率を割り当てます。つまり、プロンプトで「JSONだけ出力して」と指示しても、モデルは確率的に「すみません、以下がJSONです:」のような余計なテキストを出力する可能性を完全にはゼロにできないのです。

具体的な失敗パターンを整理すると、次のようになります。

1. 形式の逸脱

モデルがJSONの前後に自然言語の説明文を付加してしまうケースです。「以下が結果です」のような前置きや、JSON後の「何かご質問があれば…」のような追加テキストが典型例です。

2. スキーマの不一致

期待するフィールド名が "age" なのに "年齢" と返したり、数値であるべきフィールドに文字列 "25歳" を入れたりするケースです。

3. 構文エラー

末尾のカンマ(trailing comma)、引用符の不一致、括弧の対応ミスなど、JSONとして不正な出力が生成されるケースです。特に長い出力やネストの深い構造で発生しやすくなります。

4. 幻覚(Hallucination)との複合

スキーマに存在しないフィールドを勝手に追加したり、要求されたフィールドを省略したりするケースです。

プロダクション環境では、これらの問題に対して「リトライ戦略」や「出力のバリデーション + 再生成」で対処することが一般的です。しかし、リトライにはレイテンシとコストが伴います。そもそもモデルの生成プロセス自体に制約を加えれば、これらの問題はすべて原理的に解消できるはずです。

この発想から生まれたのが、次のセクションで解説する Constrained Decoding(制約付きデコーディング) です。

Constrained Decodingの基本原理

生成プロセスに制約を埋め込む

Constrained Decodingの核心的なアイデアは、驚くほどシンプルです。「生成してからバリデーションする」のではなく、「生成の各ステップで、文法的に許可されるトークンだけを選択肢に残す」のです。

温度パラメータとサンプリングで解説したとおり、LLMは各ステップでロジットベクトル $\bm{z} \in \mathbb{R}^{|\mathcal{V}|}$ を出力し、softmaxで確率分布に変換します。Constrained Decodingでは、この変換の前にトークンマスクを適用します。

具体的なイメージを掴むために、JSONの { の直後の生成を考えてみましょう。有効なJSONでは、{ の直後に来れるのは " (フィールド名の開始)か } (空オブジェクト)だけです。数字や英単語、[ などは文法的に許されません。そこで、"} に対応するトークン以外の確率をゼロにしてしまえば、モデルは必ず文法的に正しいトークンを選ぶことになります。

トークンマスキングの仕組み

トークンマスキングを数式で表現しましょう。時刻 $t$ におけるロジットベクトルを $\bm{z}_t \in \mathbb{R}^{|\mathcal{V}|}$ とします。文法的に許可されるトークンの集合を $\mathcal{A}_t \subseteq \mathcal{V}$ とすると、マスク付きロジットは次のように定義されます。

$$ \tilde{z}_{t,i} = \begin{cases} z_{t,i} & \text{if } i \in \mathcal{A}_t \\ -\infty & \text{if } i \notin \mathcal{A}_t \end{cases} $$

ここで $-\infty$ を代入する理由は、softmaxの性質にあります。$\exp(-\infty) = 0$ なので、マスクされたトークンの確率は厳密にゼロになります。

マスク付きロジットにsoftmaxを適用すると、制約付き確率分布が得られます。

$$ P_{\text{constrained}}(y_t = i \mid y_{

ここで $\mathbb{1}[i \in \mathcal{A}_t]$ は、トークン $i$ が許可集合に含まれるとき1、そうでないとき0をとる指示関数です。この式は、許可されたトークンの間でのみ確率を再配分していることを意味します。

重要なのは、許可されたトークン間の相対的な確率の比は変わらないという点です。

$$ \frac{P_{\text{constrained}}(y_t = i)}{P_{\text{constrained}}(y_t = j)} = \frac{\exp(z_{t,i})}{\exp(z_{t,j})} = \frac{P_{\text{original}}(y_t = i)}{P_{\text{original}}(y_t = j)} \quad (i, j \in \mathcal{A}_t) $$

つまり、制約はモデルの「好み」の順序を変えません。文法的に許可されるトークンの中で、モデルが元々高い確率を割り当てていたものがそのまま高い確率を保つのです。制約は文法的な正しさを保証する最小限の介入であり、モデルの表現力を不必要に制限しません。

許可集合 $\mathcal{A}_t$ をどう決めるか

ここで核心的な問題が浮上します。各ステップで「どのトークンが文法的に許可されるか」を高速に判定するにはどうすればよいでしょうか。ナイーブな方法は、「これまでの生成済みテキスト + 候補トークン」が有効なプレフィックスであるかを全トークンについてチェックすることですが、語彙サイズが数万以上のLLMではこのアプローチは計算コストが高すぎます。

この問題を効率的に解決するのが、形式言語理論に基づくアプローチです。具体的には、制約を有限状態マシン(FSM)または文脈自由文法(CFG)として表現し、各生成ステップで「現在の状態から遷移可能なトークン」を高速に列挙します。

次のセクションでは、正規表現からFSMへの変換を通じて、このメカニズムの詳細を見ていきましょう。

正規表現からFSMへ — トークンマスクの構築

形式言語理論の復習

Constrained Decodingの理論的基盤は、コンピュータサイエンスの形式言語理論にあります。ここでは最も基本的な道具立てである正規表現と有限状態マシン(FSM)を簡潔に復習します。

正規表現(Regular Expression)は、文字列のパターンを記述する形式言語です。たとえば [0-9]+ は「1桁以上の数字列」を表し、"[a-z]+" は「ダブルクォートで囲まれた小文字英字列」を表します。

有限状態マシン(Finite State Machine, FSM)は、有限個の状態と状態間の遷移規則から成る計算モデルです。FSMは5つ組 $(Q, \Sigma, \delta, q_0, F)$ で定義されます。

  • $Q$ : 状態の有限集合
  • $\Sigma$ : 入力アルファベット(ここでは文字の集合)
  • $\delta : Q \times \Sigma \to Q$ : 状態遷移関数
  • $q_0 \in Q$ : 初期状態
  • $F \subseteq Q$ : 受理状態の集合

正規表現から等価なFSMを構築するアルゴリズム(Thompson構成法やサブセット構成法)はよく知られており、任意の正規表現に対して決定性有限オートマトン(DFA)を構築できます。

文字レベルFSMからトークンレベルFSMへ

ここからがConstrained Decodingに特有の難しさです。FSMの遷移は文字レベルで定義されていますが、LLMのデコーディングはトークンレベルで行われます。1つのトークンは1文字かもしれませんし、複数文字の部分文字列かもしれません。たとえば、BPEトークナイザーでは "name"["\"", "name", "\""] の3トークンや ["\"n", "ame", "\""] の3トークンに分割される可能性があります。

この橋渡しをするために、文字レベルのFSM状態をトークンレベルの遷移関数に変換します。具体的には、各FSM状態 $q$ に対して、語彙中の各トークン $v \in \mathcal{V}$ を文字ごとに「歩かせ」て、全文字で遷移が成功するかどうかを事前に計算します。

トークン $v$ が文字列 $c_1 c_2 \ldots c_n$ から構成されるとき、状態 $q$ からトークン $v$ を消費した後の状態は次のように定義できます。

$$ \delta_{\text{token}}(q, v) = \delta(\ \delta(\ldots\delta(\delta(q, c_1), c_2)\ldots, c_{n-1}), c_n) $$

途中のいずれかの文字で遷移が定義されていなければ(デッドステートに到達すれば)、そのトークンは状態 $q$ では許可されません。この計算を全ての $(q, v)$ の組み合わせについて事前に行い、テーブル化しておきます。

事前計算のテーブルサイズは $|Q| \times |\mathcal{V}|$ です。状態数が数十〜数百、語彙サイズが32,000〜128,000程度であれば、メモリ消費は現実的な範囲に収まります。

許可集合の列挙アルゴリズム

事前計算テーブルがあれば、デコーディング時の許可集合 $\mathcal{A}_t$ の列挙は極めて高速になります。手順をまとめると次のとおりです。

  1. 初期状態 $q_0$ から開始する
  2. 各デコーディングステップ $t$ で、現在の状態 $q_t$ に対して、$\delta_{\text{token}}(q_t, v)$ が定義されている(デッドステートでない)全てのトークン $v$ を $\mathcal{A}_t$ とする
  3. $\mathcal{A}_t$ に基づいてトークンマスクを適用し、トークン $y_t$ をサンプリングする
  4. 状態を $q_{t+1} = \delta_{\text{token}}(q_t, y_t)$ に更新する
  5. $q_{t+1} \in F$(受理状態)であれば生成を終了してよい。そうでなければステップ2に戻る

この手順のステップ2は、事前計算テーブルから $q_t$ の行を読み出すだけなので、$O(|\mathcal{V}|)$ で完了します。これはsoftmaxの計算と同じオーダーであり、デコーディングのボトルネックにはなりません。

FSMベースの制約構築を理解したところで、次はより実用的な制約である「JSON Schema」による出力制御を見ていきましょう。

JSON Schema Guided Generation

JSONの文法をFSMで表現する

JSONは、その構文規則がRFC 8259で厳密に定義されている形式言語です。しかし、JSON全体の文法は再帰的な構造(オブジェクトの中にオブジェクトがネストする等)を含むため、厳密には正規言語ではなく文脈自由言語に分類されます。

では、前のセクションで説明したFSM(正規言語のみを受理)では不十分なのでしょうか? 実は、JSON Schemaで具体的な構造が指定されている場合、ネストの深さが有限に制限されるため、対応するFSM(より正確にはプッシュダウンオートマトンを有限深さで展開したもの)を構築できます。

たとえば、次のようなJSON Schemaを考えます。

{
  "type": "object",
  "properties": {
    "name": {"type": "string"},
    "age": {"type": "integer"},
    "is_student": {"type": "boolean"}
  },
  "required": ["name", "age", "is_student"]
}

このスキーマが許容する出力は有限のパターンに制約されます。フィールドの順序やスペースの有無に自由度はありますが、構造自体は固定されています。この「固定された構造」をFSMの状態として展開できるのです。

スキーマからFSMへの変換

JSON Schema Guided Generationの実装では、スキーマの各構成要素を正規表現のパターンに変換し、それらを連結・選択・繰り返しで組み合わせてFSMを構築します。主要な変換ルールは以下のとおりです。

文字列型 ("type": "string") は、" で囲まれた任意の有効なJSON文字列にマッチする正規表現に変換されます。エスケープシーケンス(\", \\, \n など)も考慮する必要があります。

$$ \text{string} \to \texttt{“} \cdot (\text{unescaped} \mid \texttt{\textbackslash} \cdot \text{escape\_char})^{*} \cdot \texttt{“} $$

整数型 ("type": "integer") は、オプションの符号と1桁以上の数字にマッチします。

$$ \text{integer} \to \texttt{-}^{?} \cdot (\texttt{0} \mid \texttt{[1-9]} \cdot \texttt{[0-9]}^{*}) $$

真偽値型 ("type": "boolean") は、リテラル true または false にマッチします。

$$ \text{boolean} \to \texttt{true} \mid \texttt{false} $$

オブジェクト型は、これらの要素型を組み合わせて構成します。必須フィールドの順序が固定されている場合、オブジェクト全体の正規表現は各フィールドの正規表現の連結になります。フィールドの順序が自由な場合は、全順列の選択(union)になりますが、実装上はフィールド順序を固定してFSMの状態数を抑えることが一般的です。

実装上の工夫

実際のJSON Schema Guided Generationライブラリ(OutlinesGuidanceなど)では、以下のような工夫が施されています。

1. 空白の正規化

JSONでは { の後や : の前後に任意の空白を入れられますが、空白の自由度をそのままFSMに反映すると状態数が爆発します。多くの実装では、空白パターンをあらかじめ正規化(たとえば「空白なし」に固定)することで状態数を削減しています。

2. フィールド順序の固定

前述のとおり、JSONオブジェクトのフィールド順序を固定することで、状態の組み合わせ爆発を避けています。$n$ フィールドのオブジェクトで順序が自由な場合、$n!$ 通りのパスが必要ですが、順序を固定すれば1通りで済みます。

3. インデックスキャッシュ

同じFSM状態に対するトークンマスクは毎回同一なので、一度計算したマスクをキャッシュして再利用します。特にバッチ推論では複数のリクエストが同じスキーマを共有することが多く、キャッシュの効果が大きくなります。

JSON Schemaによる制約の仕組みがわかったところで、次はLLMアプリケーションで広く使われているFunction Callingが、この構造化出力の技術とどのように関係しているかを見ていきましょう。

Function Calling / Tool Use と構造化出力

Function Callingの本質は構造化出力

Tool Use / Function Callingで解説したとおり、Function Callingは「LLMが外部の関数を呼び出す」機能です。しかし、技術的な本質を見ると、Function Callingは構造化出力の特殊ケースにほかなりません。

Function Callingの出力は、次のような構造化された形式です。

{
  "function": "get_weather",
  "arguments": {
    "city": "Tokyo",
    "unit": "celsius"
  }
}

これは、「関数名」と「引数」のJSON Schemaに制約された構造化出力そのものです。つまり、Function CallingをConstrained Decodingで実装すれば、「存在しない関数名を呼ぶ」「引数の型が間違っている」といったエラーを原理的に排除できます。

二段階の制約

Function Callingの制約は、二段階で考えることができます。

第一段階: 関数選択の制約

利用可能な関数が $\{f_1, f_2, \ldots, f_K\}$ であるとき、"function" フィールドの値はこれらのいずれかでなければなりません。これは正規表現 f_1 | f_2 | ... | f_K として表現できます。

第二段階: 引数スキーマの制約

選択された関数 $f_k$ に対して、引数のJSON Schemaが決まります。たとえば get_weather の引数スキーマが {"city": string, "unit": "celsius" | "fahrenheit"} であれば、"unit" フィールドは "celsius""fahrenheit" の2択に制約されます。

この二段階の制約をFSMで実装する場合、第一段階の出力(関数名)に応じて第二段階のFSMが動的に切り替わります。具体的には、関数名の生成が完了した時点で、その関数の引数スキーマに対応するFSMへと遷移先を切り替えます。

並列関数呼び出し

最新のLLM API(OpenAIのParallel Function Calling等)では、一度のレスポンスで複数の関数を呼び出す機能が提供されています。これは構造化出力の観点からは、「関数呼び出しオブジェクトの配列」というスキーマで自然に表現できます。

{
  "type": "array",
  "items": {
    "type": "object",
    "properties": {
      "function": {"enum": ["get_weather", "search_web"]},
      "arguments": { ... }
    }
  }
}

配列の各要素に対して同じ制約を適用することで、並列関数呼び出しも型安全に生成できます。

ここまでで、構造化出力の概念的な仕組みとJSON Schema/Function Callingへの応用を理解しました。次のセクションでは、制約付きデコーディングの数学的な性質をより深く分析していきましょう。

制約付き確率分布の数学的定式化

制約付きデコーディングの確率モデル

ここまでの議論を統一的な確率モデルとして定式化します。LLMが出力するトークン列を $\bm{y} = (y_1, y_2, \ldots, y_T)$ とし、制約を満たすトークン列の集合を $\mathcal{C}$ とします。

制約なしのLLMが定義する分布は、因果言語モデルの連鎖律に従います。

$$ P_{\text{LLM}}(\bm{y}) = \prod_{t=1}^{T} P_{\text{LLM}}(y_t \mid y_{

制約付きデコーディングの目標は、$\mathcal{C}$ 上での条件付き分布からサンプリングすることです。

$$ P_{\text{constrained}}(\bm{y}) = P_{\text{LLM}}(\bm{y} \mid \bm{y} \in \mathcal{C}) = \frac{P_{\text{LLM}}(\bm{y}) \cdot \mathbb{1}[\bm{y} \in \mathcal{C}]}{\sum_{\bm{y}’ \in \mathcal{C}} P_{\text{LLM}}(\bm{y}’)} $$

分母の $\sum_{\bm{y}’ \in \mathcal{C}} P_{\text{LLM}}(\bm{y}’)$ は、制約を満たす全てのトークン列にわたる確率の合計であり、分配関数(正規化定数)の役割を果たします。

自己回帰分解と局所的マスキングの関係

トークンマスキングによるConstrained Decodingは、この条件付き分布を各ステップに分解した近似とみなせます。制約 $\mathcal{C}$ がFSMで表現される場合、時刻 $t$ での許可集合 $\mathcal{A}_t$ は、FSMの現在の状態 $q_t$ から到達可能な受理状態が存在するようなトークンの集合です。

$$ \mathcal{A}_t = \{v \in \mathcal{V} \mid \exists \bm{y}_{>t} \text{ s.t. } (y_t = v) \wedge (y_1, \ldots, y_T) \in \mathcal{C}\} $$

この定義は、「現在のトークンを選んだ後、制約を満たす完全な列が少なくとも1つ存在する」ことを要求しています。つまり、許可集合は将来の到達可能性を考慮して定義されます。

FSMベースの制約では、この到達可能性は「現在のFSM状態から受理状態に到達できるかどうか」として効率的に判定できます。事前に各状態について受理状態への到達可能性を計算しておけば、$O(1)$ で判定が完了します。

貪欲法との関係

制約付き貪欲法(Constrained Greedy Decoding)は、各ステップで制約付き分布のモード(最頻値)を選択します。

$$ y_t^{*} = \arg \max_{v \in \mathcal{A}_t} P_{\text{LLM}}(v \mid y_{

ビームサーチと組み合わせることも可能です。制約付きビームサーチでは、各ビームの展開時に $\mathcal{A}_t$ に含まれるトークンのみを候補とします。ビームサーズが制約なしの場合と同様に、制約付きの場合もビーム幅を増やすことで、条件付き分布におけるより高い確率の列を探索できます。

KLダイバージェンスによる制約の影響評価

制約を加えることで、元のLLMの分布からどれだけ離れるかをKLダイバージェンスで測ることができます。時刻 $t$ での制約の影響度は次のように定量化されます。

マスクされたトークンの元々の確率質量を $m_t$ とします。

$$ m_t = \sum_{i \notin \mathcal{A}_t} P_{\text{LLM}}(y_t = i \mid y_{

$m_t$ が小さいほど(つまりモデルが元々制約に合致するトークンを高確率で選んでいるほど)、制約の影響は小さくなります。よく訓練されたモデルに対して適切な制約を課す場合、$m_t$ は一般に非常に小さく、制約による生成品質の劣化はほとんどありません。

逆に $m_t$ が大きい場合、つまりモデルが本来出力したいトークンの大部分が制約で禁止されている場合、制約付きの出力品質は大きく低下する可能性があります。これは、モデルの能力と制約のミスマッチを示唆しており、プロンプトの改善やモデル自体の改善が必要です。

数学的な定式化を理解したところで、実際にPythonでConstrained Decodingを実装してみましょう。理論がどのようにコードに落とし込まれるかを確認し、実際の挙動を観察します。

Pythonでの簡易Constrained Decoding実装

概要

ここでは、教育目的でミニマルなConstrained Decodingエンジンを実装します。正規表現からDFAを構築し、語彙中の各トークンについて許可/禁止を判定し、トークンマスクを適用してサンプリングする一連の流れを実装します。

まず、正規表現から簡易的なDFAを構築し、トークンレベルの遷移テーブルを事前計算するモジュールを作成します。

import numpy as np
import re
from collections import defaultdict
from typing import Dict, List, Set, Tuple, Optional


class SimpleDFA:
    """正規表現パターンから簡易DFAを構築するクラス"""

    def __init__(self, states: Set[int], alphabet: Set[str],
                 transitions: Dict[Tuple[int, str], int],
                 start_state: int, accept_states: Set[int]):
        self.states = states
        self.alphabet = alphabet
        self.transitions = transitions
        self.start_state = start_state
        self.accept_states = accept_states
        # 各状態から受理状態に到達可能かを事前計算
        self.reachable = self._compute_reachability()

    def _compute_reachability(self) -> Set[int]:
        """受理状態から逆方向にBFSして到達可能な状態を列挙"""
        reachable = set(self.accept_states)
        # 逆遷移を構築
        reverse_trans = defaultdict(set)
        for (src, char), dst in self.transitions.items():
            reverse_trans[dst].add(src)
        # BFS
        queue = list(self.accept_states)
        while queue:
            state = queue.pop(0)
            for prev_state in reverse_trans[state]:
                if prev_state not in reachable:
                    reachable.add(prev_state)
                    queue.append(prev_state)
        return reachable

    def step(self, state: int, char: str) -> Optional[int]:
        """1文字遷移。遷移不可能なら None を返す"""
        return self.transitions.get((state, char), None)

    def step_token(self, state: int, token_str: str) -> Optional[int]:
        """トークン(複数文字)を順に遷移。途中で失敗したら None"""
        current = state
        for char in token_str:
            current = self.step(current, char)
            if current is None:
                return None
        return current

    def is_reachable_to_accept(self, state: int) -> bool:
        """この状態から受理状態に到達可能か"""
        return state in self.reachable


def build_integer_dfa() -> SimpleDFA:
    """整数 (0 | [1-9][0-9]*) にマッチするDFAを構築"""
    # 状態: 0=開始, 1=0を読んだ(受理), 2=1-9を読んだ(受理), 3=デッド
    states = {0, 1, 2, 3}
    digits = set("0123456789")
    alphabet = digits
    transitions = {}
    # 状態0: '0'→状態1, '1'-'9'→状態2
    transitions[(0, '0')] = 1
    for d in "123456789":
        transitions[(0, d)] = 2
    # 状態1: 0の後は何も来ない(JSONの整数: 0は単独)
    for d in digits:
        transitions[(1, d)] = 3  # デッド
    # 状態2: 後続の数字→状態2(ループ)
    for d in digits:
        transitions[(2, d)] = 2
    # 状態3: デッドステート
    for d in digits:
        transitions[(3, d)] = 3
    return SimpleDFA(states, alphabet, transitions,
                     start_state=0, accept_states={1, 2})

このコードは、JSONの整数(0 または [1-9][0-9]*)にマッチする簡易的なDFAを定義しています。状態0が開始状態、状態1は 0 を読んだ後の受理状態(JSONでは先頭ゼロの後に数字は続けられません)、状態2は 19 で始まる数列の受理状態です。

次に、このDFAを使ったトークンマスクの構築とConstrained Decodingを実装します。

import matplotlib.pyplot as plt


class ConstrainedDecoder:
    """DFAベースのConstrained Decodingエンジン"""

    def __init__(self, dfa: SimpleDFA, vocabulary: List[str]):
        self.dfa = dfa
        self.vocabulary = vocabulary
        self.vocab_size = len(vocabulary)
        # トークンレベル遷移テーブルを事前計算
        self.token_transitions = self._precompute_token_transitions()

    def _precompute_token_transitions(self) -> Dict[Tuple[int, int], Optional[int]]:
        """全 (状態, トークン) ペアに対する遷移先を事前計算"""
        table = {}
        for state in self.dfa.states:
            for tok_idx, tok_str in enumerate(self.vocabulary):
                next_state = self.dfa.step_token(state, tok_str)
                if next_state is not None and self.dfa.is_reachable_to_accept(next_state):
                    table[(state, tok_idx)] = next_state
                else:
                    table[(state, tok_idx)] = None
        return table

    def get_allowed_tokens(self, state: int) -> List[int]:
        """現在の状態で許可されるトークンインデックスのリスト"""
        allowed = []
        for tok_idx in range(self.vocab_size):
            if self.token_transitions.get((state, tok_idx)) is not None:
                allowed.append(tok_idx)
        return allowed

    def apply_mask(self, logits: np.ndarray, state: int) -> np.ndarray:
        """ロジットにトークンマスクを適用"""
        allowed = self.get_allowed_tokens(state)
        masked_logits = np.full_like(logits, -np.inf)
        for idx in allowed:
            masked_logits[idx] = logits[idx]
        return masked_logits

    def sample(self, logits: np.ndarray, temperature: float = 1.0) -> int:
        """softmaxサンプリング"""
        if temperature <= 0:
            return int(np.argmax(logits))
        scaled = logits / temperature
        # 数値安定性のためにmaxを引く
        scaled -= np.max(scaled[scaled > -np.inf]) if np.any(scaled > -np.inf) else 0
        exp_logits = np.exp(scaled)
        probs = exp_logits / np.sum(exp_logits)
        return int(np.random.choice(len(probs), p=probs))

    def generate(self, logits_sequence: List[np.ndarray],
                 temperature: float = 1.0) -> Tuple[List[int], List[str]]:
        """制約付きデコーディングでトークン列を生成"""
        state = self.dfa.start_state
        token_ids = []
        token_strs = []
        for logits in logits_sequence:
            masked_logits = self.apply_mask(logits, state)
            if np.all(masked_logits == -np.inf):
                break  # 遷移可能なトークンがない → 生成終了
            tok_idx = self.sample(masked_logits, temperature)
            token_ids.append(tok_idx)
            token_strs.append(self.vocabulary[tok_idx])
            # 状態を更新
            state = self.token_transitions[(state, tok_idx)]
            # 受理状態に到達したら終了
            if state in self.dfa.accept_states:
                # 次のステップで遷移可能なトークンがあるか確認
                if not self.get_allowed_tokens(state):
                    break
        return token_ids, token_strs

ConstrainedDecoder クラスの設計のポイントは、DFAの事前計算テーブル _precompute_token_transitions にあります。初期化時に全ての状態とトークンの組み合わせについて遷移先を計算するため、デコーディング時のオーバーヘッドは最小限に抑えられます。apply_mask メソッドがロジットベクトルに $-\infty$ を設定し、sample メソッドがsoftmax + サンプリングを実行するという流れは、前のセクションの数式をそのまま実装したものです。

では、このエンジンを使って、制約付きと制約なしのデコーディングを比較してみましょう。

np.random.seed(42)

# 簡易語彙: 数字の各桁 + 非数字トークン
vocabulary = ["0", "1", "2", "3", "4", "5", "6", "7", "8", "9",
              " ", ",", ".", "a", "b", "{", "}", '"']
vocab_size = len(vocabulary)

# 整数DFAを構築
dfa = build_integer_dfa()
decoder = ConstrainedDecoder(dfa, vocabulary)

# 擬似ロジットを生成(実際にはLLMが出力する)
# 数字トークンに高めのロジットを設定しつつ、非数字トークンにもある程度の確率を与える
def make_mock_logits(vocab_size: int, prefer_digits: bool = True) -> np.ndarray:
    logits = np.random.randn(vocab_size) * 0.5
    if prefer_digits:
        logits[:10] += 2.0  # 数字トークンを優遇
    else:
        logits[10:] += 1.5  # 非数字トークンを優遇
    return logits

n_trials = 1000
n_steps = 5

# 制約なしサンプリング
unconstrained_valid = 0
unconstrained_outputs = []
for _ in range(n_trials):
    output_tokens = []
    for step in range(n_steps):
        logits = make_mock_logits(vocab_size, prefer_digits=(step < 3))
        probs = np.exp(logits - np.max(logits))
        probs /= probs.sum()
        tok_idx = np.random.choice(vocab_size, p=probs)
        output_tokens.append(vocabulary[tok_idx])
    output_str = "".join(output_tokens)
    unconstrained_outputs.append(output_str)
    # 整数として有効か検証
    if re.match(r'^(0|[1-9][0-9]*)$', output_str):
        unconstrained_valid += 1

# 制約付きサンプリング
constrained_valid = 0
constrained_outputs = []
for _ in range(n_trials):
    logits_seq = [make_mock_logits(vocab_size, prefer_digits=(s < 3))
                  for s in range(n_steps)]
    token_ids, token_strs = decoder.generate(logits_seq, temperature=1.0)
    output_str = "".join(token_strs)
    constrained_outputs.append(output_str)
    if re.match(r'^(0|[1-9][0-9]*)$', output_str):
        constrained_valid += 1

print(f"制約なし: {unconstrained_valid}/{n_trials} 件が有効な整数 "
      f"({100*unconstrained_valid/n_trials:.1f}%)")
print(f"制約付き: {constrained_valid}/{n_trials} 件が有効な整数 "
      f"({100*constrained_valid/n_trials:.1f}%)")
print(f"\n制約なしの出力例: {unconstrained_outputs[:8]}")
print(f"制約付きの出力例: {constrained_outputs[:8]}")

制約なしの場合、数字トークンに高いロジットを与えていても、1,000回中すべてが有効な整数になるわけではありません。空白やカンマ、英字が混入してしまうケースが必ず発生します。一方、制約付きデコーディングでは、DFAが文法的に不正なトークンをマスクするため、出力は100%有効な整数になります。出力例を見ると、制約なしでは "31 24" のように空白が混入した出力が現れるのに対し、制約付きでは常に "3124" のような正しい整数が生成されていることが確認できます。

次に、制約がモデルの確率分布に与える影響を可視化します。

fig, axes = plt.subplots(1, 3, figsize=(16, 5))

# 3つの状態でのマスキング効果を可視化
states_to_show = [
    (0, "State 0 (Start)"),
    (2, "State 2 (After digit 1-9)"),
    (1, "State 1 (After '0')")
]

for ax, (state, title) in zip(axes, states_to_show):
    logits = make_mock_logits(vocab_size, prefer_digits=True)

    # 元の確率分布
    probs_orig = np.exp(logits - np.max(logits))
    probs_orig /= probs_orig.sum()

    # 制約付き確率分布
    masked_logits = decoder.apply_mask(logits, state)
    valid_mask = masked_logits > -np.inf
    if np.any(valid_mask):
        probs_constrained = np.zeros(vocab_size)
        exp_vals = np.exp(masked_logits[valid_mask] - np.max(masked_logits[valid_mask]))
        probs_constrained[valid_mask] = exp_vals / exp_vals.sum()
    else:
        probs_constrained = np.zeros(vocab_size)

    x = np.arange(vocab_size)
    width = 0.35
    bars1 = ax.bar(x - width/2, probs_orig, width, label='Original',
                   color='#4a9eff', alpha=0.7)
    bars2 = ax.bar(x + width/2, probs_constrained, width, label='Constrained',
                   color='#ff6b6b', alpha=0.7)

    ax.set_xlabel('Token')
    ax.set_ylabel('Probability')
    ax.set_title(title)
    ax.set_xticks(x)
    ax.set_xticklabels(vocabulary, fontsize=8, rotation=45)
    ax.legend(fontsize=8)
    ax.grid(axis='y', alpha=0.3)

    # マスクされた確率質量を表示
    masked_mass = sum(probs_orig[i] for i in range(vocab_size) if not valid_mask[i])
    ax.text(0.02, 0.95, f'Masked mass: {masked_mass:.2%}',
            transform=ax.transAxes, fontsize=8,
            verticalalignment='top',
            bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))

plt.tight_layout()
plt.savefig('constrained_decoding_distributions.png', dpi=150, bbox_inches='tight')
plt.show()

このグラフから、3つの重要な特徴が読み取れます。

  1. State 0(開始状態)では、数字トークン "0""9" のみが許可され、空白やカンマ、英字などの非数字トークンがすべてマスクされています。元の分布でこれらの非数字トークンに割り当てられていた確率質量(Masked mass)が、許可された数字トークンに再配分されています。

  2. State 2(1-9の後)では、"0""9" の全数字トークンが許可されます。開始状態と異なり "0" も許可されるのは、先頭以外の位置ではゼロが有効だからです。この状態では、元のモデルが数字を好む傾向と制約が一致しているため、マスクされる確率質量は比較的小さくなります。

  3. State 1("0" の後)では、整数 0 として完結するために、許可されるトークンがなくなります(受理状態であり、JSONの整数 0 の後に数字は続けられないため)。これは生成の終了条件を表しています。

続いて、制約の「きつさ」がサンプリング効率に与える影響を可視化しましょう。

fig, ax = plt.subplots(1, 1, figsize=(10, 6))

# 非数字トークンの優遇度を変化させて、制約の影響度を測定
bias_levels = np.linspace(-2, 4, 20)
masked_masses = []
effective_vocab_ratios = []

for bias in bias_levels:
    logits = np.random.randn(vocab_size) * 0.5
    logits[:10] += 2.0       # 数字トークンの基本ロジット
    logits[10:] += bias      # 非数字トークンのバイアスを変化
    probs = np.exp(logits - np.max(logits))
    probs /= probs.sum()
    allowed = decoder.get_allowed_tokens(0)  # 開始状態での許可トークン
    masked_mass = sum(probs[i] for i in range(vocab_size) if i not in allowed)
    masked_masses.append(masked_mass)
    effective_vocab_ratios.append(len(allowed) / vocab_size)

ax.plot(bias_levels, masked_masses, 'o-', color='#ff6b6b',
        linewidth=2, markersize=6, label='Masked probability mass')
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5, label='50% threshold')
ax.fill_between(bias_levels, masked_masses, alpha=0.1, color='#ff6b6b')
ax.set_xlabel('Non-digit token bias', fontsize=12)
ax.set_ylabel('Masked probability mass', fontsize=12)
ax.set_title('Impact of Constraint Strength on Probability Mass', fontsize=14)
ax.legend(fontsize=11)
ax.grid(alpha=0.3)
ax.set_ylim(0, 1.0)

plt.tight_layout()
plt.savefig('constraint_impact.png', dpi=150, bbox_inches='tight')
plt.show()

このグラフからは、モデルの出力分布と制約の相性が生成品質に直結することが見て取れます。非数字トークンのバイアスが小さい(数字を好むモデル)場合、マスクされる確率質量は小さく、制約の影響はほとんどありません。しかし、非数字トークンのバイアスが大きくなるにつれてマスクされる確率質量が増加し、50%を超えると、モデルが「本当は出力したかった」トークンの大半が禁止されていることを意味します。このような状況では、制約付きの出力品質が低下する可能性があります。実用上は、マスクされる確率質量が一貫して小さくなるよう、プロンプトの調整やモデルのファインチューニングで対処します。

ここまでで基本的なConstrained Decodingの実装と可視化ができました。次のセクションでは、JSONオブジェクトを対象としたより実践的な制約の実装を見ていきましょう。

JSONオブジェクト生成の実装

シンプルなJSONスキーマの制約

前のセクションでは整数の制約を実装しましたが、実用上はJSONオブジェクト全体を制約の対象とすることが求められます。ここでは、フィールド順序を固定した簡易的なJSON生成の制約を状態マシンとして実装します。

class JsonObjectDFA:
    """
    固定スキーマのJSONオブジェクト生成用の簡易状態マシン
    スキーマ例: {"name": string, "age": integer}
    生成されるパターン: {"name": "...", "age": 数字}
    """

    def __init__(self, fields: List[Tuple[str, str]]):
        """
        fields: [(フィールド名, 型)] のリスト
        型は 'string' または 'integer' のみサポート
        """
        self.fields = fields
        self.states = {}  # 状態名 → 状態ID
        self.transitions = {}
        self.accept_states = set()
        self._build_states()

    def _build_states(self):
        """スキーマからDFA状態を構築"""
        state_id = 0
        # 状態を順に作成
        # '{' を期待する開始状態
        self.start_state = state_id
        self.states['start'] = state_id

        # 各フィールドについて状態を生成
        for i, (fname, ftype) in enumerate(self.fields):
            # '"' でフィールド名開始
            state_id += 1
            self.states[f'field_{i}_quote_open'] = state_id
            # フィールド名の各文字
            for j, char in enumerate(fname):
                state_id += 1
                self.states[f'field_{i}_name_{j}'] = state_id
            # '"' でフィールド名終了
            state_id += 1
            self.states[f'field_{i}_quote_close'] = state_id
            # ':' を期待
            state_id += 1
            self.states[f'field_{i}_colon'] = state_id
            # 値の開始
            state_id += 1
            self.states[f'field_{i}_value_start'] = state_id
            if ftype == 'string':
                # 文字列値の中身
                state_id += 1
                self.states[f'field_{i}_value_content'] = state_id
                # '"' で文字列値終了
                state_id += 1
                self.states[f'field_{i}_value_end'] = state_id
            elif ftype == 'integer':
                # 整数値(1桁以上の数字)
                state_id += 1
                self.states[f'field_{i}_value_digits'] = state_id

            # フィールド間の ',' または最後の '}'
            if i < len(self.fields) - 1:
                state_id += 1
                self.states[f'field_{i}_comma'] = state_id

        # 終了の '}'
        state_id += 1
        self.states['end_brace'] = state_id
        self.accept_states = {state_id}

        # 全状態集合
        self.all_states = set(range(state_id + 1))

    def get_allowed_chars(self, state_name: str) -> Set[str]:
        """状態名に基づいて許可される文字集合を返す"""
        if state_name == 'start':
            return {'{'}
        elif '_quote_open' in state_name:
            return {'"'}
        elif '_name_' in state_name:
            # フィールド名の次の文字を決定
            parts = state_name.split('_')
            field_idx = int(parts[1])
            char_idx = int(parts[3])
            fname = self.fields[field_idx][0]
            if char_idx < len(fname):
                return {fname[char_idx]}
            return {'"'}
        elif '_quote_close' in state_name:
            return {'"'}
        elif '_colon' in state_name:
            return {':'}
        elif '_value_start' in state_name:
            parts = state_name.split('_')
            field_idx = int(parts[1])
            ftype = self.fields[field_idx][1]
            if ftype == 'string':
                return {'"'}
            elif ftype == 'integer':
                return set('123456789')
        elif '_value_content' in state_name:
            # 文字列内容: 英数字 + 終了の '"'
            return set('abcdefghijklmnopqrstuvwxyz '
                      'ABCDEFGHIJKLMNOPQRSTUVWXYZ') | {'"'}
        elif '_value_digits' in state_name:
            parts = state_name.split('_')
            field_idx = int(parts[1])
            # 次がカンマか閉じ括弧かで許可文字が変わる
            return set('0123456789') | {',' if field_idx < len(self.fields) - 1 else '}'}
        elif '_value_end' in state_name:
            parts = state_name.split('_')
            field_idx = int(parts[1])
            if field_idx < len(self.fields) - 1:
                return {','}
            else:
                return {'}'}
        elif '_comma' in state_name:
            return {'"'}
        elif state_name == 'end_brace':
            return set()  # 受理状態、遷移なし
        return set()


# スキーマ定義
schema_fields = [("name", "string"), ("age", "integer")]
json_dfa = JsonObjectDFA(schema_fields)

# 各状態で許可される文字を表示
print("=== JSON Object DFA States ===")
for name, sid in sorted(json_dfa.states.items(), key=lambda x: x[1]):
    allowed = json_dfa.get_allowed_chars(name)
    is_accept = "  [ACCEPT]" if sid in json_dfa.accept_states else ""
    allowed_display = sorted(allowed) if allowed else "(none)"
    print(f"  State {sid:2d} ({name:30s}): allowed={allowed_display}{is_accept}")

この出力から、JSON生成の各段階でどの文字(ひいてはどのトークン)が許可されるかが一目でわかります。たとえば、開始状態では { のみが許可され、フィールド名 "name" の途中では n, a, m, e の各文字が順に許可されるという、きわめて厳格な制約が適用されています。フィールドの値の部分でのみ自由度があり、文字列型では英数字と空白が許可され、整数型では数字が許可されます。

生成シミュレーション

構築したJSON DFAを使って、擬似的なLLMロジットから制約付きのJSON生成をシミュレーションしましょう。

def simulate_json_generation(json_dfa: JsonObjectDFA,
                             max_steps: int = 100) -> str:
    """JSONオブジェクトの制約付き生成をシミュレーション"""
    output = []
    state_names = {v: k for k, v in json_dfa.states.items()}
    current_state_id = json_dfa.states['start']

    for step in range(max_steps):
        current_name = state_names.get(current_state_id, 'unknown')
        allowed_chars = json_dfa.get_allowed_chars(current_name)
        if not allowed_chars:
            break

        # 擬似ロジット: 各許可文字にランダムなスコアを割り当て
        char_list = sorted(allowed_chars)
        logits = np.random.randn(len(char_list)) * 0.5

        # 文字列値では英字を、整数値では特定の数字を若干優遇
        for i, c in enumerate(char_list):
            if c.isalpha():
                logits[i] += 1.0
            elif c.isdigit():
                logits[i] += 0.5

        # softmaxサンプリング
        probs = np.exp(logits - np.max(logits))
        probs /= probs.sum()
        chosen_idx = np.random.choice(len(char_list), p=probs)
        chosen_char = char_list[chosen_idx]
        output.append(chosen_char)

        # 状態遷移(簡易版: 文字に応じて次状態を決定)
        current_state_id += 1
        if current_state_id in json_dfa.accept_states:
            break
        if current_state_id >= max(json_dfa.states.values()):
            break

    return "".join(output)

np.random.seed(123)
print("=== 制約付きJSON生成のシミュレーション ===\n")
for trial in range(5):
    result = simulate_json_generation(json_dfa)
    # 有効なJSONかチェック
    import json
    try:
        parsed = json.loads(result)
        valid = True
    except json.JSONDecodeError:
        valid = False
    print(f"Trial {trial+1}: {result}")
    print(f"  Valid JSON: {valid}")
    if valid:
        print(f"  Parsed: {parsed}")
    print()

このシミュレーションでは、DFAの制約が各生成ステップで厳格に適用されるため、出力は必ず {"name": "...", "age": 数字} の形式に従います。文字列値や整数値の具体的な内容はモデルの「好み」(擬似ロジット)によって決まりますが、構造自体は100%保証されています。これがConstrained Decodingの最大の強みです。

ここまでで実装を確認できました。次に、実運用で重要となる性能トレードオフと実用上の注意点を議論します。

性能トレードオフと実用上の考慮事項

レイテンシへの影響

Constrained Decodingは「デコーディング時」に追加の計算を行います。具体的なオーバーヘッドを整理しましょう。

事前計算のコスト

DFAの構築とトークンレベル遷移テーブルの計算は、デコーディング開始前に一度だけ行います。テーブルサイズは $|Q| \times |\mathcal{V}|$ であり、各セルの計算にはトークンの文字数分の遷移が必要です。トークンの平均文字数を $\bar{L}$ とすると、全体の計算量は $O(|Q| \times |\mathcal{V}| \times \bar{L})$ です。

実用的な数値として、状態数 $|Q| = 100$、語彙サイズ $|\mathcal{V}| = 32{,}000$、平均トークン長 $\bar{L} = 4$ の場合、約 $1.28 \times 10^7$ 回の文字遷移が必要です。これは現代のCPUでは数十ミリ秒で完了します。

デコーディング時のコスト

各ステップでのマスク適用は $O(|\mathcal{V}|)$ であり、LLMのforward passの計算量 $O(d_{\text{model}}^2 + T \cdot d_{\text{model}})$ と比べると無視できる程度です。したがって、Constrained Decodingによるデコーディング速度の低下はほとんどありません。

品質への影響

制約が生成品質に与える影響は、前述のマスクされる確率質量 $m_t$ に依存します。

良い場合($m_t$ が小さい): モデルが元々制約に沿った出力を好んでいる場合、制約はごくわずかな確率の再配分しか行わず、出力品質への影響はほぼありません。たとえば、JSONの出力を求められたモデルが実際にJSONらしいトークンに高い確率を割り当てている場合がこれに該当します。

悪い場合($m_t$ が大きい): モデルの好みと制約が大きく乖離している場合、強制的な制約がモデルの「意図」を歪め、フィールドの値が不自然になったり、文字列の内容が意味をなさなくなったりする可能性があります。

この影響を定量的にまとめます。

fig, ax = plt.subplots(1, 1, figsize=(10, 6))

# 制約の強さと品質の関係を模式的に可視化
constraint_strength = np.linspace(0, 1, 100)

# 形式の正しさ: 制約なし→確率的、制約あり→100%
format_correctness_unconstrained = 0.85 - 0.3 * constraint_strength
format_correctness_constrained = np.ones_like(constraint_strength)

# 意味の品質: 制約が強すぎると低下
semantic_quality = 1.0 - 0.5 * constraint_strength**2

ax.plot(constraint_strength, format_correctness_unconstrained,
        '--', color='#4a9eff', linewidth=2, label='Format correctness (unconstrained)')
ax.plot(constraint_strength, format_correctness_constrained,
        '-', color='#4a9eff', linewidth=2, label='Format correctness (constrained)')
ax.plot(constraint_strength, semantic_quality,
        '-', color='#ff6b6b', linewidth=2, label='Semantic quality (constrained)')
ax.fill_between(constraint_strength,
                format_correctness_constrained, semantic_quality,
                alpha=0.1, color='green', label='Quality gap')

ax.set_xlabel('Constraint Strength (masked probability mass)', fontsize=12)
ax.set_ylabel('Quality Score', fontsize=12)
ax.set_title('Trade-off: Format Correctness vs. Semantic Quality', fontsize=14)
ax.legend(fontsize=10, loc='lower left')
ax.grid(alpha=0.3)
ax.set_ylim(0, 1.1)
ax.set_xlim(0, 1)

plt.tight_layout()
plt.savefig('quality_tradeoff.png', dpi=150, bbox_inches='tight')
plt.show()

このグラフには、Constrained Decodingの本質的なトレードオフが可視化されています。青の実線(Format correctness, constrained)は常に1.0であり、制約の強さに関わらず形式の正しさが100%保証されることを示しています。一方、赤の線(Semantic quality)は制約が強くなるにつれて低下していきます。制約の強さが小さい領域では品質の低下はわずかですが、強い制約では意味的な品質が大きく損なわれます。青の破線(Format correctness, unconstrained)はプロンプトベースのアプローチの限界を示しており、制約なしでは形式の正しさが保証されないことがわかります。

実用ライブラリの比較

現在、構造化出力のためのオープンソースライブラリがいくつか存在します。それぞれのアプローチと特徴を整理します。

Outlines(dottxt-ai)

正規表現やJSON SchemaからFSMを構築し、トークンマスクを生成するライブラリです。interegular ライブラリを使って正規表現からDFAを構築し、Hugging FaceのTransformersやvLLMと統合できます。FSMの事前計算により、デコーディング時のオーバーヘッドが最小限に抑えられています。

Guidance(Microsoft)

テンプレートベースの制約記述言語を提供するライブラリです。Pythonの構文に似たドメイン固有言語(DSL)で制約を記述し、生成をステップごとに制御できます。正規表現ベースのアプローチよりも柔軟で、条件分岐やループを含む複雑な制約も表現できます。

Instructor(Jason Liu)

Pydanticモデルとの統合に特化したライブラリです。Pythonのクラス定義からJSON Schemaを自動生成し、LLMの出力を型安全なオブジェクトに変換します。バリデーション + リトライの戦略を採用しており、Constrained Decodingとは異なるアプローチですが、開発体験が優れています。

API プロバイダのネイティブサポート

OpenAIのStructured Outputs(response_format: json_schema)、AnthropicのTool Use、GoogleのControlled Generationなど、主要なAPI プロバイダも構造化出力のネイティブサポートを提供しています。これらはサーバー側でConstrained Decodingを実装しているため、ユーザーはスキーマを渡すだけで100%の形式保証を得られます。

文脈自由文法(CFG)ベースのアプローチ

FSM(正規言語)では表現できない制約もあります。たとえば「括弧の対応が正しいこと」や「再帰的なJSONネスト」は文脈自由言語であり、プッシュダウンオートマトンが必要です。

CFGベースのConstrained Decodingでは、Earleyパーサーや CYKパーサーを使って、各ステップで「次に来れるトークン」を列挙します。計算量はFSMベースよりも大きくなりますが、より表現力の高い制約を扱えます。

ただし、実用上はJSON Schemaのネストの深さが有限に固定されるため、前述のようにFSMで十分に対応できるケースがほとんどです。CFGベースのアプローチが真に必要になるのは、プログラミング言語のソースコード生成やXML生成など、より複雑な構文規則を扱う場合です。

バッチ推論での考慮事項

バッチ推論では、異なるリクエストが異なるスキーマを持つ可能性があるため、各リクエストに対して独立したFSM状態を管理する必要があります。これはvLLMのようなバッチ推論エンジンで実装上の課題となります。

解決策の一つは、同一スキーマのリクエストをグループ化し、FSMの事前計算テーブルをグループ間で共有することです。もう一つは、スキーマごとのマスクベクトルをGPUメモリに保持し、各ステップでバッチ全体に対して一括でマスクを適用することです。後者のアプローチでは、マスクの適用がGPU上の要素ごとの乗算として実装され、CPUとGPU間のデータ転送のボトルネックを回避できます。

まとめ

本記事では、LLMの出力をJSON等の構造化形式に制約する技術であるConstrained Decodingの理論と実装について解説しました。

  • プロンプトベースの限界: LLMのテキスト生成は本質的に確率的であり、プロンプト指示だけでは出力形式の100%保証は得られません。プロダクション環境ではこの不確実性がシステム障害の原因となります
  • Constrained Decodingの原理: 生成の各ステップで文法的に許可されるトークンのみを残すトークンマスキングにより、出力が必ず指定された形式に従うことを保証します。許可されたトークン間の確率比は変わらないため、モデルの表現力は最小限にしか影響を受けません
  • FSMベースの制約構築: 正規表現からDFAを構築し、文字レベルの遷移をトークンレベルの遷移テーブルに事前変換することで、デコーディング時のオーバーヘッドを最小化しています
  • JSON Schema Guided Generation: JSONの文法規則をFSMとして表現し、フィールド順序の固定や空白の正規化といった工夫で状態数を削減することが実用化の鍵です
  • Function Callingとの関係: Function Calling / Tool Useは構造化出力の特殊ケースであり、関数名の選択と引数スキーマの二段階の制約として定式化できます
  • 性能トレードオフ: Constrained Decodingのレイテンシオーバーヘッドはほぼ無視でき、形式の正しさは100%保証されます。ただし、制約がモデルの分布と大きく乖離する場合、意味的な品質が低下する可能性があります

構造化出力の技術は、LLMをプロダクション・システムに組み込む上で不可欠な基盤技術です。RAG(検索拡張生成)やAIエージェントなど、LLMの出力を下流のプログラムが消費するあらゆるシステムで、この技術が信頼性の土台を提供しています。

次のステップとして、以下の記事も参考にしてください。

画像なし
【LLM】Tool Use / Function Callingの仕組みと実装
LLMが外部APIやツールを呼び出すTool Useの仕組みをさらに深く理解できます。
画像なし
Temperature・Top-k・Top-pサンプリングを比較して理解する
制約なしのサンプリング手法を理解することで、制約付きとの比較がより明確になります。
画像なし
ビームサーチとは?テキスト生成の探索アルゴリズムを解説
制約付きビームサーチの基礎となるビームサーチの仕組みを解説しています。