一文要約: 推論ではCheckpointから同じModelとTokenizerを復元し、PromptをToken IDsへEncodeし、1 Tokenずつ生成します。Scriptは短くても、モデル契約、乱数性、長さの境界は一つも捨てられません。

信頼できるCheckpointからModelとTokenizerを復元し、PromptをEncode、Autoregressiveに生成してTextへDecodeする

📦 2024年の元デモリポジトリ: github.com/waylandzhang/Transformer-from-scratch。当時のデモはData、Model、訓練、生成を一つの model.py に置いています。本章の inference.py は、第18・19章で監査した独立ファイルと組み合わせます。


20.1 推論は「Updateしない訓練」だけではない

訓練推論
Input連続Sequenceと1 Token shiftしたTargetsPromptと、これまでに生成したTokens
各Forward多くの位置のLossを並列計算最後の位置から1つのNext Tokenを選ぶ
Parameter Updateありなし
Dropouttrain() でランダムeval() で無効
Autograd必要inference_mode() で無効

model.eval()torch.inference_mode() は別の仕事をします。前者はModule内のTrain/Eval behaviorを切り替え、後者はAutogradの記録を止めます。本書ModelにはDropoutはありますがBatchNormはありません。したがって、ここでの eval() の直接の役割はDropoutを止めることであり、「BatchNormを固定する」ことではありません。

eval() はSamplingも決定的にしません。temperature > 0 では依然として確率分布からSamplingします。temperature = 0 が第18章で定義したGreedy decodingです。


20.2 重みの前にモデル契約を復元する

第19章のCheckpointには model_configtokenizer_namemodel_state_dict が入っています。推論ScriptでLayer数、幅、Vocabulary sizeをもう一度推測するべきではありません。

def load_model(checkpoint_path, device):
    checkpoint = torch.load(
        checkpoint_path,
        map_location="cpu",
        weights_only=True,
    )
    required = {"model_state_dict", "model_config", "tokenizer_name"}
    if not isinstance(checkpoint, dict) or not required <= checkpoint.keys():
        raise ValueError("checkpoint is missing inference metadata")

    model_config = ModelConfig(**checkpoint["model_config"])
    tokenizer = tiktoken.get_encoding(checkpoint["tokenizer_name"])
    if tokenizer.n_vocab != model_config.vocab_size:
        raise ValueError("tokenizer vocabulary does not match the model")

    model = Model(model_config)
    model.load_state_dict(checkpoint["model_state_dict"], strict=True)
    model.to(device)
    model.eval()
    return model, tokenizer

まずCPUへLoadし、復元したModelを現在のDeviceへ移します。これなら、Checkpointを保存したときの特定GPUを今のMachineに求めません。strict=True は全ParameterのNameとShapeの一致を要求します。

weights_only=True はUnpicklerが構築できるObjectを制限しますが、どこかからDownloadしたFileを信頼できるDataに変えるわけではありません。自分で作ったか、出所を信頼できるCheckpointだけをLoadします。


20.3 Promptも同じTokenizer Contractを守る

日本語版では、蒔絵工房の道具目録をコーパスにした例を続けます。

prompt = "蒔絵筆:"
prompt_ids = tokenizer.encode(prompt, disallowed_special=())
if not prompt_ids:
    raise ValueError("prompt must contain at least one token")
x = torch.tensor(prompt_ids, dtype=torch.long, device=device)[None, :]

Token IDsのSampleを固定値で書きません。Tokenizer version、先頭Space、記号一つでも結果は変わります。Checkpointと一緒にLoadしたTokenizerの実際の結果を表示すべきです。

Promptが context_length を超える場合、第18章の教学実装は各Forwardで最後の context_length Tokensだけを使います。より古いPrompt tokensは返りTensorに残りますが、新しいTokenには影響しません。

本書では固定のSinusoidal position tableを使います。新しいVisible windowへCropすると、Window内のPosition indexも0から始め直します。これは教育用実装の明示的な方針であり、すべてのServing systemのDefaultではありません。


20.4 生成Parameterには境界があり、万能レシピはない

20.4.1 Temperature

第18章の定義は明確です。

  • temperature = 0: 最大Logitを直接選ぶGreedy decoding;
  • 0 < temperature < 1: 分布を鋭くする;
  • temperature = 1: Logitの相対Scaleを変えない;
  • temperature > 1: 分布を平らにする。

「事実質問は0.2、創作は0.9」という普遍的な処方箋ではありません。ModelのCalibration、Corpus、Taskで適切な値は変わります。

20.4.2 Top-K

top_k=K はLogitが大きいK Tokensを残し、その中でSoftmaxとSamplingを行います。本書CodeはKをVocabulary sizeまでに制限し、None はFilterなしを意味します。Top-KはLong tailを除きますが、正しさ、一貫性、非重複を保証しません。

20.4.3 Max New TokensとStop条件

この教育用ModelはEOS tokenもStop stringも定義していません。そのため max_new_tokens=80 なら新しい80 Tokensをちょうど生成します。「最大80で、Modelが自分で止まる」のではありません。


20.5 完成版 inference.py

import argparse
from pathlib import Path

import tiktoken
import torch

from model import Model, ModelConfig


def select_device():
    if torch.cuda.is_available():
        return torch.device("cuda")
    if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
        return torch.device("mps")
    return torch.device("cpu")


def load_model(checkpoint_path, device):
    checkpoint = torch.load(
        checkpoint_path,
        map_location="cpu",
        weights_only=True,
    )
    required = {"model_state_dict", "model_config", "tokenizer_name"}
    if not isinstance(checkpoint, dict) or not required <= checkpoint.keys():
        raise ValueError("checkpoint is missing inference metadata")

    model_config = ModelConfig(**checkpoint["model_config"])
    tokenizer = tiktoken.get_encoding(checkpoint["tokenizer_name"])
    if tokenizer.n_vocab != model_config.vocab_size:
        raise ValueError("tokenizer vocabulary does not match the model")

    model = Model(model_config)
    model.load_state_dict(checkpoint["model_state_dict"], strict=True)
    model.to(device)
    model.eval()
    return model, tokenizer


def encode_prompt(tokenizer, prompt, device):
    prompt_ids = tokenizer.encode(prompt, disallowed_special=())
    if not prompt_ids:
        raise ValueError("prompt must contain at least one token")
    return torch.tensor(
        prompt_ids, dtype=torch.long, device=device
    ).unsqueeze(0)


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--checkpoint",
        type=Path,
        default=Path("model/checkpoint.pt"),
    )
    parser.add_argument("--prompt", required=True)
    parser.add_argument("--max-new-tokens", type=int, default=80)
    parser.add_argument("--temperature", type=float, default=0.7)
    parser.add_argument("--top-k", type=int)
    parser.add_argument("--seed", type=int, default=1337)
    return parser.parse_args()


def main():
    args = parse_args()
    device = select_device()
    model, tokenizer = load_model(args.checkpoint, device)
    x = encode_prompt(tokenizer, args.prompt, device)

    torch.manual_seed(args.seed)

    with torch.inference_mode():
        y = model.generate(
            x,
            max_new_tokens=args.max_new_tokens,
            temperature=args.temperature,
            top_k=args.top_k,
        )

    parameter_count = sum(p.numel() for p in model.parameters())
    print(f"device={device} prompt_tokens={x.size(1)}")
    print(f"parameters={parameter_count:,}")
    print("---")
    print(tokenizer.decode(y[0].tolist()))


if __name__ == "__main__":
    main()

実際に訓練したCheckpointで実行します。

python inference.py --checkpoint model/checkpoint.pt --prompt "蒔絵筆:" --temperature 0.7 --top-k 50

OutputはそのCheckpointから実際に生成すべきです。旧版の整ったSample outputには再現可能なCheckpointがないため、本版では実験結果として掲載しません。


20.6 一つのStepを見るなら、実際の候補を表示する

もっともらしい確率表を手で作らず、現在のModelから取り出します。

@torch.inference_mode()
def show_next_candidates(model, tokenizer, x, k=5):
    x_crop = x[:, -model.config.context_length:]
    logits, _ = model(x_crop)
    probs = torch.softmax(logits[0, -1], dim=-1)
    top_probs, top_ids = torch.topk(probs, min(k, probs.numel()))

    for probability, token_id in zip(top_probs, top_ids):
        piece = tokenizer.decode_single_token_bytes(
            token_id.item()
        ).decode("utf-8", errors="backslashreplace")
        print(repr(piece), f"{probability.item():.4f}")

1 TokenはUTF-8 bytesの一部だけを持つことがあり、単独で完全な文字にDecodeできるとは限りません。そのため、この表示は「すべてのTokenは一語」と仮定せず decode_single_token_bytes() から始めます。


20.7 問題はまず境界から調べる

  • Parameter nameやShapeが合わない: strict=False で隠さず、model.py とCheckpointが同じContractか確認します。
  • Outputが重複する: Corpusの重複、Overfit、Calibration、Sampling設定のどれでも起こり得ます。Temperatureを上げるのは実験であり、保証ではありません。
  • Outputが支離滅裂: TemperatureとTop-Kの前に、実際のTrain/Valid loss、Tokenizer identity、Promptが訓練分布からどれほど外れるかを見ます。
  • 生成が遅い: この generate() は毎Step、現在Windowを再計算します。GPUは役立ち得ますが、第22章のKV CacheがStep間のK/V再計算を除きます。Flash Attentionは1回のAttention内のMemory trafficと実装を改善します。両者は同じ最適化ではありません。

20.8 本章のまとめ

  • Checkpointの model_configtokenizer_name からContractを復元し、推論Scriptで再推測しない;
  • TensorをまずCPUへLoadし、復元したModelを現在のDeviceへ移す;
  • eval() はModule behavior、inference_mode() はAutograd、Decoding ruleはSampling乱数性を決める;
  • Temperature、Top-K、生成長には正確な境界があり、普遍的な推奨範囲はない;
  • EOSの定義がないので、このModelはちょうど max_new_tokens 個を生成する;
  • Token IDs、候補確率、Parameter数、生成Textは、あらかじめ書いたOutputではなく現在のCodeとCheckpointから得る。

章末チェックリスト

  • Checkpointから完全に同じModelとTokenizerを復元できる
  • eval()inference_mode()、Greedy decoding、Samplingを区別できる
  • Prompt cropping、Temperature、Top-K、長さの境界を説明できる
  • 用意されたSample outputに頼らず、実Checkpointで inference.py を実行できる

Part 5 のまとめ

第18章でModelを定義し、第19章で訓練して完全な状態を保存し、第20章で推論に必要な部分だけをLoadしてTextを生成します。監査後の3ファイルはConfig check、Data境界、Recovery stateを含むため、旧版の「400行未満」より長くなりました。行数は目的ではありません。これは今も教育用の実装ですが、各簡略化の境界を明記します。

この連鎖を理解すればGPT-style decoderの主幹が見えます。しかしChatGPTやProduction serving systemを複製したことにはなりません。実際のSystemにはKV Cache、Batching、Custom kernel、Parallelism、Memory management、Request schedulingも必要です。


次章予告

現在の生成Loopは各StepでActive window全体のAttentionを再計算します。第21章は Flash Attention から始まります。Attentionの数学的な結果を変えず、計算とMemory accessを組み替え、1回のAttentionを効率化します。第22章ではKV Cacheを使い、Decoding step間のK/V再計算を除きます。

このページを引用する
Zhang, Wayland (2026). 第20章: 手書き inference.py - 推論ロジック. In Transformer アーキテクチャ:直感から実装まで. https://waylandz.com/llm-transformer-book-ja/chapter-20-inference-py/
@incollection{zhang2026transformer_ja_chapter-20-inference-py,
  author = {Zhang, Wayland},
  title = {第20章: 手書き inference.py - 推論ロジック},
  booktitle = {Transformer アーキテクチャ:直感から実装まで},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book-ja/chapter-20-inference-py/}
}