一文要約: 推論ではCheckpointから同じModelとTokenizerを復元し、PromptをToken IDsへEncodeし、1 Tokenずつ生成します。Scriptは短くても、モデル契約、乱数性、長さの境界は一つも捨てられません。
📦 2024年の元デモリポジトリ: github.com/waylandzhang/Transformer-from-scratch。当時のデモはData、Model、訓練、生成を一つの
model.pyに置いています。本章のinference.pyは、第18・19章で監査した独立ファイルと組み合わせます。
20.1 推論は「Updateしない訓練」だけではない
| 訓練 | 推論 | |
|---|---|---|
| Input | 連続Sequenceと1 Token shiftしたTargets | Promptと、これまでに生成したTokens |
| 各Forward | 多くの位置のLossを並列計算 | 最後の位置から1つのNext Tokenを選ぶ |
| Parameter Update | あり | なし |
| Dropout | train() でランダム | 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_config、tokenizer_name、model_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_configとtokenizer_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再計算を除きます。