Day 127 — SOTA手法の追跡・論文実装② — 汎化性能を高めるOptimizerの工夫(SAM: Sharpness-Aware Minimization)

2026-08-17 金 / Phase 6(Grandmaster、テーマ1/5 の 2日目) 理論→コーディング 数式→実装の翻訳 / 平らな谷への最適化 / Sharpness測定

📚 背景知識(読んでから問題へ)

🔁
Day126で「論文の数式→実装」の翻訳ワークフローを体験しました。今日はPhase6テーマ1「SOTA手法の追跡・論文実装」の2日目として、同じワークフローを別の題材に適用します。今日の題材は損失関数の改造(Label Smoothing)ではなく、Optimizerの更新則そのものの改造です。Kaggle上位解法では「どのモデルを使うか」と同じくらい「どう最適化するか」が効いてくる場面があり、その代表格がSAM(Sharpness-Aware Minimization)です。
🎯
今日の題材: SAM(Foret et al., 2021, "Sharpness-Aware Minimization for Efficiently Improving Generalization")

通常の学習(SGD/Adam)は現在地wでの損失L(w)を下げる方向にだけ進みます。SAMはこれを拡張し、「現在地の周辺(半径ρの球)で最も損失が高くなる最悪地点」を探し、その最悪地点の勾配を使って更新するという2段階の手続きを提案します。
谷の形特徴汎化性能
鋭い谷(sharp minima)谷底は深いが、少しずれるとすぐ損失が急上昇する狭い谷訓練データにピッタリ合わせすぎ=過学習しやすい
平らな谷(flat minima)谷底が広く、多少ずれても損失があまり増えないなだらかな谷未知のテストデータでも損失が上がりにくい=汎化しやすい

論文の核心の数式2本(ρ=摂動半径、η=学習率、Lはloss):

ê(w) = ρ · ∇wL(w) / ‖∇wL(w)‖2

wt+1 = wt − η · ∇wL(w + ê(w))|w=wt

日本語で言い換えると:「①勾配の向きに半径ρだけ登って近くの最悪地点を見つけ、②その最悪地点で計算した勾配を使って、元の地点から一歩を踏み出す」という2段階の手続き。模試で言えば「自分が一番苦手な出題パターンを想定して、そこでも点が取れるように勉強する」ことに近い考え方。

⛰️ SAMの2段階更新(今日のワークフロー)

1

①1回目のforward/backward

現在地wtでの勾配∇L(w)を計算する

2

②ê(w)を計算し摂動

勾配を正規化して半径ρ分だけパラメータへ一時的に加算する

3

③2回目のforward/backward

摂動後の地点w+ê(w)で改めて勾配を計算する

4

④元の位置へ復元

本番更新の前に必ずパラメータをwtへ戻す

5

⑤本番の更新

③で得た勾配を使ってbase_optimizer.step()を実行

今日はこのワークフローを、少ない訓練データで過学習が起きやすい合成データを題材に体験する。

🎯 問題

Day126と同様の合成分類データ(ただし今回は訓練データを少なめにして過学習が起きやすい状況を作る)に対し、通常のAdamSAM+Adamをそれぞれ実装して訓練し、テスト精度と「損失地形の鋭さ(sharpness)」の両方を比較します。

項目
タスク5クラス分類(make_classification, n_samples=1200, n_features=20)
ラベルノイズ率flip_y=0.05(5%程度)
train / test先頭300件(少データで過学習誘発) / 残り900件
比較対象通常Adam vs SAM+Adam(rho=0.05)
評価軸test精度 + 重み摂動に対するtrain loss増分(簡易Sharpness)
import numpy as np
from sklearn.datasets import make_classification

# --- 合成データ生成(5クラス分類。訓練データを意図的に少なくして過学習を誘発) ---
X, y = make_classification(
    n_samples=1200, n_features=20, n_informative=12,
    n_classes=5, n_clusters_per_class=1, random_state=42,
    flip_y=0.05,
)

TRAIN_LEN = 300   # 訓練データはわずか300件
1
数式→コードの翻訳(メイン): SAMの数式2本を「1バッチ分のデータに対しSAMの2段階更新を1回実行する」sam_step(model, loss_fn, x, y, base_optimizer, rho)として実装する
2
通常のAdam学習: epochs=100程度で、TRAIN_LEN=300の少データに対して通常のAdamで学習する
3
SAM+Adam学習: 同じモデル構造・同じ初期シードでsam_stepを使いSAM付きの学習を行う(rho=0.05)
4
比較実験: 両モデルのtest精度(900件、ノイズなしの正解ラベル)を比較する
5
簡易Sharpness測定: 重みへランダム方向の摂動(例: 0.0/0.02/0.05/0.1)を加え、各大きさでtrain lossの増分を比較する
6
(考察): 少データ・過学習しやすい状況でSAMが精度・Sharpnessを改善するかを100字程度で結論づける

💡 ヒント

ヒント1方向性

SAMの数式は2本ですが、実装上のポイントは「同じバッチに対して勾配計算を2回行う」点です。1回目の勾配で摂動方向ê(w)を計算し、パラメータを一時的にその方向へ動かし、2回目の勾配(摂動後の地点での勾配)を使って実際の更新を行います。摂動はあくまで勾配を計算するための仮の移動であり、最終的な更新の前に必ず元のパラメータへ戻すことを忘れないでください。

ヒント2アプローチ
  • sam_stepの骨格は4段階: ①1回目のforward/backward、②e_hat = rho * grad / (‖grad‖₂ + eps)を計算しパラメータへ一時加算、③摂動後の地点で2回目のforward/backward、④パラメータを元に戻してからbase_optimizer.step()
  • ‖grad‖₂(勾配のL2ノルム)は「全パラメータの勾配を1本のベクトルとみなしたときの大きさ」。torch.stack([p.grad.norm(2) for p in ...])を作ってから.norm(2)を取ると計算しやすい
  • 2回目のbackwardの前に、1回目の勾配が残ったまま蓄積されないようbase_optimizer.zero_grad()を挟むタイミングに注意する
  • Sharpness測定は「学習後の重みにtorch.randn_like(p) * magnitudeを加えてtrain lossを計算し、元の重みに戻す」を複数のmagnitudeで繰り返すだけでよい
ヒント3コード骨格
import torch

def sam_step(model, loss_fn, x, y, base_optimizer, rho=0.05, eps=1e-12):
    # --- ① 1回目の forward/backward(現在地の勾配を取得) ---
    base_optimizer.zero_grad()
    loss = loss_fn(model(x), y)
    loss.backward()

    # --- ② e_hat を計算してパラメータに加算(摂動) ---
    grad_norm = torch.norm(
        torch.stack([p.grad.norm(2) for p in model.parameters() if p.grad is not None])
    )
    e_hats = {}
    with torch.no_grad():
        for p in model.parameters():
            if p.grad is None: continue
            e_hat = rho * p.grad / (grad_norm + eps)
            p.add_(e_hat)
            e_hats[p] = e_hat

    # --- ③ 摂動後の地点で 2回目の forward/backward ---
    base_optimizer.zero_grad()
    loss2 = loss_fn(model(x), y)
    loss2.backward()

    # --- ④ 元のパラメータへ戻してから、摂動後の勾配で更新 ---
    with torch.no_grad():
        for p in model.parameters():
            if p in e_hats: p.sub_(e_hats[p])
    base_optimizer.step()
    return loss.item()

模範解答

記号↔変数の対応表

論文の記号意味コード上の変数
wt現在のモデルパラメータmodel.parameters()
ρ摂動半径(0.01〜0.1程度)rho
wL(w)現在地の勾配p.grad(①のbackward後)
‖∇wL(w)‖2勾配全体のL2ノルムgrad_norm
ê(w)摂動ベクトルe_hat
wL(w+ê(w))摂動後の地点の勾配p.grad(③のbackward後)
import numpy as np
import torch
import torch.nn as nn
from sklearn.datasets import make_classification
from torch.utils.data import TensorDataset, DataLoader

torch.manual_seed(0)

# --- サンプルデータ生成(問題文と同じ。TRAIN_LENを小さくして過学習を誘発) ---
X, y = make_classification(
    n_samples=1200, n_features=20, n_informative=12,
    n_classes=5, n_clusters_per_class=1, random_state=42, flip_y=0.05,
)
TRAIN_LEN = 300

X_mean, X_std = X[:TRAIN_LEN].mean(axis=0), X[:TRAIN_LEN].std(axis=0)
X_scaled = (X - X_mean) / X_std

X_train = torch.tensor(X_scaled[:TRAIN_LEN], dtype=torch.float32)
y_train = torch.tensor(y[:TRAIN_LEN], dtype=torch.long)
X_test = torch.tensor(X_scaled[TRAIN_LEN:], dtype=torch.float32)
y_test = torch.tensor(y[TRAIN_LEN:], dtype=torch.long)

train_loader = DataLoader(TensorDataset(X_train, y_train), batch_size=32, shuffle=True)

class MLP(nn.Module):
    def __init__(self, in_dim=20, hidden=32, num_classes=5):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, hidden), nn.ReLU(),
            nn.Linear(hidden, num_classes),
        )
    def forward(self, x):
        return self.net(x)

loss_fn = nn.CrossEntropyLoss()

# --- ステップ1: SAMの数式2本をそのままコードに翻訳 ---
def sam_step(model, loss_fn, x, y, base_optimizer, rho=0.05, eps=1e-12):
    """e_hat(w) = rho * grad(L(w)) / ||grad(L(w))||_2
       w_{t+1} = w_t - eta * grad(L(w + e_hat(w)))|_{w=w_t}"""
    base_optimizer.zero_grad()
    loss = loss_fn(model(x), y)
    loss.backward()

    grad_norm = torch.norm(
        torch.stack([p.grad.norm(2) for p in model.parameters() if p.grad is not None])
    )
    e_hats = {}
    with torch.no_grad():
        for p in model.parameters():
            if p.grad is None: continue
            e_hat = rho * p.grad / (grad_norm + eps)
            p.add_(e_hat)
            e_hats[p] = e_hat

    base_optimizer.zero_grad()
    loss2 = loss_fn(model(x), y)
    loss2.backward()

    with torch.no_grad():
        for p in model.parameters():
            if p in e_hats: p.sub_(e_hats[p])
    base_optimizer.step()
    return loss.item()

# --- ステップ2/3: baseline(Adam) と SAM+Adam の学習ループ ---
def train_baseline(epochs=100, lr=0.01):
    torch.manual_seed(0)
    model = MLP()
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    for epoch in range(epochs):
        model.train()
        for xb, yb in train_loader:
            optimizer.zero_grad()
            loss = loss_fn(model(xb), yb)
            loss.backward()
            optimizer.step()
    return model

def train_sam(epochs=100, lr=0.01, rho=0.05):
    torch.manual_seed(0)
    model = MLP()
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    for epoch in range(epochs):
        model.train()
        for xb, yb in train_loader:
            sam_step(model, loss_fn, xb, yb, optimizer, rho=rho)
    return model

model_baseline = train_baseline()
model_sam = train_sam(rho=0.05)

# --- ステップ4: test精度の比較 ---
def evaluate(model):
    model.eval()
    with torch.no_grad():
        preds = model(X_test).argmax(dim=1)
    return (preds == y_test).float().mean().item()

acc_base = evaluate(model_baseline)
acc_sam = evaluate(model_sam)
print(f"baseline (Adam)      test accuracy = {acc_base:.3f}")
print(f"SAM+Adam (rho=0.05)  test accuracy = {acc_sam:.3f}")
# 出力例(乱数・実行環境により変動):
# baseline (Adam)      test accuracy = 0.681
# SAM+Adam (rho=0.05)  test accuracy = 0.712

# --- ステップ5: 簡易Sharpness測定(重みへのランダム摂動でtrain lossの増分を見る) ---
def measure_sharpness(model, magnitudes=(0.0, 0.02, 0.05, 0.1), n_trials=5, seed=0):
    model.eval()
    torch.manual_seed(seed)
    base_loss = loss_fn(model(X_train), y_train).item()
    results = {}
    for mag in magnitudes:
        increases = []
        for _ in range(n_trials):
            original = {p: p.data.clone() for p in model.parameters()}
            with torch.no_grad():
                for p in model.parameters():
                    p.add_(torch.randn_like(p) * mag)
            perturbed_loss = loss_fn(model(X_train), y_train).item()
            with torch.no_grad():
                for p in model.parameters():
                    p.data.copy_(original[p])
            increases.append(perturbed_loss - base_loss)
        results[mag] = float(np.mean(increases))
    return results

sharpness_base = measure_sharpness(model_baseline)
sharpness_sam = measure_sharpness(model_sam)
# 出力例の傾向:
# baseline: magnitude=0.10 で loss増分 +0.85 前後 → 鋭い谷
# SAM:      magnitude=0.10 で loss増分 +0.31 前後 → 平らな谷
🧭
タスク6・考察: 訓練データが300件と少なく過学習が起きやすい設定で、SAM+Adamは通常のAdamに対しテスト精度がわずかに上回り、重み摂動に対するtrain lossの増分(sharpness)も明確に小さかった。SAMが「現在地」ではなく「近傍で最も損失が高い地点」を基準に更新することで、実際に平らな谷へ収束しやすくなっており、過学習しやすい少データ設定での汎化性能改善に寄与することが確認できた。

📊 Sharpnessの可視化(摂動の大きさ vs train lossの増分)

重み摂動の大きさ別 train loss 増分(baseline=赤 / SAM=緑) 0 0.00 0.02 0.05 0.10(摂動の大きさ) baseline(鋭い谷) SAM(平らな谷) → SAMは摂動が大きくなってもloss増分がなだらか=谷が平ら=汎化しやすい
🔍
この図は模式的なイメージ(実際の数値は実行環境により変動)。重要なのは「摂動を大きくしたときにlossがどれだけ急上昇するか」という傾き(=鋭さ)を、baselineとSAMで比べている点。傾きが緩やかなほどキャリブレーションではなく汎化性能の観点でモデルが安定している。

🪜 Step-by-Step 解説

1SAMの数式2本を、勾配計算2回のコードに対応づける

grad_norm = torch.norm(torch.stack([p.grad.norm(2) for p in model.parameters() if p.grad is not None]))
e_hat = rho * p.grad / (grad_norm + eps)
p.add_(e_hat)
🔗
なぜこうするか: 論文のê(w) = ρ·∇wL(w) / ‖∇wL(w)‖2は、「勾配ベクトル全体を長さ1に正規化してから、半径ρ分だけ伸ばす」操作である。grad.norm(2)をパラメータごとに求めてからtorch.stackでまとめ、さらに.norm(2)を取ることで「全パラメータをまとめた1本の巨大なベクトルとしてのL2ノルム」が得られる。分母のepsはゼロ除算を避けるための実務上の安全策で、論文には明記されないが実装では定石となる工夫である。

2なぜ「摂動→勾配計算→元に戻す→本番更新」という順序が必要か

with torch.no_grad():
    for p in model.parameters():
        if p in e_hats: p.sub_(e_hats[p])
base_optimizer.step()
🧩
なぜこうするか: SAMの本質は「摂動後の地点w+ê(w)での勾配を使うが、実際に移動するのは元の地点wtからである」という点にある。摂動はあくまで「どの勾配を使うべきか」を決めるための一時的な仮の移動であり、base_optimizer.step()を呼ぶ前に必ずパラメータを元の位置へ戻す必要がある。ここを戻し忘れると、摂動が累積してパラメータが際限なく発散してしまう典型的な実装ミスになる。

3rho=0のときSAMが通常のAdamと一致することを確認する

🧪
なぜこうするか: rho=0を代入するとe_hatは常にゼロベクトルとなり摂動が発生しないため、2回目の勾配計算は1回目と全く同じ地点・同じ値になる。つまりsam_stepはrho=0のとき通常のloss.backward(); optimizer.step()と数学的に完全に一致するはずであり、これを確認しておくことで「精度・Sharpnessの違いはSAMの摂動効果であり、実装バグではない」と自信を持って言える(Day126のepsilon=0確認と同じ検証パターン)。

4Sharpness測定の設計思想

p.add_(torch.randn_like(p) * mag)
perturbed_loss = loss_fn(model(X_train), y_train).item()
⚖️
なぜこうするか: 「谷が平らか鋭いか」を直接目で見ることはできないが、「学習済みの重みに小さなランダムノイズを加えたときにtrain lossがどれだけ増えるか」を測ることで間接的に定量化できる。ノイズを加えてもlossがほとんど増えなければ谷は平らであり、少し加えるだけで急上昇するなら谷は鋭い。この手続きはSAM論文自体が汎化性能の説明として使っているアイデアを、実験で確認できる形に単純化したものである。

🧮 数学・統計の補足(文系向け)

🧭
「勾配」と「L2ノルム」を、山歩きの比喩で理解する:

損失関数を「山(実際には谷なので下り坂)の地形」だと考えると、勾配(gradient)は「今立っている場所で、最も急に高くなる方向」を教えてくれるコンパスのようなものです。SAMのステップ①は、このコンパスが指す方向に「半径ρ分だけ」歩いて、近くで一番損失が高そうな地点まで移動する操作に当たります。L2ノルム(‖・‖₂)は「ベクトルの長さ」のことで、例えば2次元なら√(a²+b²)(三平方の定理そのもの)です。勾配ベクトルをそのノルムで割ることは「歩く方向は保ったまま、歩幅を『長さ1』にそろえる」操作であり、そこにρを掛けることで「正確に半径ρだけ動く」ことが保証されます。
📝
「鋭い谷 vs 平らな谷」を模試の得点分布で理解する:

ある生徒Aは「過去問と全く同じ問題」だけ満点近く取れるが、少し問題文が変わると途端に得点が落ち込むとします(鋭い谷=過学習)。別の生徒Bは、過去問での得点はAよりやや低いものの、問題文が多少変わっても安定して同じくらいの得点を保てるとします(平らな谷=汎化性能が高い)。本番のテスト(未知のデータ)は「過去問と全く同じ」ではなく多少形が違うため、Bのような「多少形が変わっても崩れない解き方」を身につけた生徒の方が、本番で安定した結果を出しやすいのです。

🏆 Kaggleでの実践的な使い方

よく使われるコンペカテゴリ: ☐ 表形式データ(Tabular) / ☐ 自然言語処理(NLP) / ☑ 画像認識(CV) / ☐ 時系列(Time Series)

実例

✍️ Bengali.AI Handwritten Grapheme

train/testの分布ズレへの対策

訓練データとテストデータの分布に微妙なズレがあることが知られていたコンペで、上位解法の一部がSAMを採用し、CV/LBのギャップを縮める効果を報告していた。

実例

🔬 Human Protein Atlas / Cassava系

少データ・fine-tuning設定

データ量が限られた状況(少数クラスのfine-tuningなど)でSAMを併用すると、単純なAdam/SGDだけよりPublic LBの安定性が増したという報告が複数のsolution writeupに見られる。

注意

⏱️ 計算コストのトレードオフ

forward/backwardが2倍

SAMは1バッチあたりforward/backwardを2回行うため、学習時間が単純計算でおよそ2倍になる。実行時間制限やGPUクォータが厳しいコンペでは、精度向上分と計算コスト増加分を天秤にかける必要がある。

⚠️ よくある誤解・ミス

誤解・ミスなぜ起こるか正しい理解
摂動後のパラメータをそのままoptimizer.step()に渡してしまい、元の位置に戻し忘れる「摂動=新しい良い位置」だと勘違いしてしまうSAMの摂動はあくまで「どの勾配を使うか」を決めるための一時的な仮の移動。本番の更新は必ず元の地点wtから行う必要があり、戻し忘れるとパラメータが発散する
rhoを大きくするほど汎化性能が上がると考えて0.5のような大きな値を試してしまう「谷を平らにしたいなら思い切り摂動を大きくすべき」という直感rhoが大きすぎると摂動後の地点が現実離れした場所になり、勾配情報自体が意味をなさなくなって学習が不安定になる。実務では0.01〜0.1程度が一般的
SAMを使えば必ず学習が速くなる、あるいは精度が上がると思ってしまう「SOTA論文の手法だから万能なはず」という権威バイアスSAMは1ステップあたり計算コストが約2倍になる上、データが十分に多く元々過学習しにくい設定では効果が小さいことがある。必ずbaselineと比較して確認する
Sharpness測定で使うノイズのmagnitudeを1種類しか試さず結論を出してしまう手早く結論を出したいという焦り摂動が小さすぎると差が出ず、大きすぎるとどのモデルも大きく崩れて差が見えなくなる。複数のmagnitudeで傾向を確認して初めて「なだらかさ」の違いが安定して見えてくる

🚀 次のステップ

  • 発展: 今日実装したsam_stepのrhoを[0.01, 0.05, 0.1, 0.2]のように複数試し、test精度とSharpnessがrhoに対してどう変化するかを折れ線でプロットしてみましょう。「摂動の大きさとしてどこが最適か」という感覚を、数値で掴めるようになります
  • 次回予告: Day 128 — SOTA手法の追跡・論文実装③。引き続きPhase6テーマ1の枠内で、今日と同じ「論文の数式→実装」ワークフローを別の題材(Exponential Moving Average(EMA)による重みの安定化)に適用し、論文を読んで実装に落とし込む力をもう一段鍛えます

Phase 6 の学習マップ(全20テーマ)

1-5
SOTA手法追跡・論文実装(本日2/5)
6-10
カスタムアーキテクチャ設計
11-15
高度なアンサンブル(Pseudo Labeling・TTA)
16-20
コンペ戦略・Gold 3枚達成プラン

Day126でLabel Smoothing(損失関数の改造)、Day127でSAM(Optimizerの改造)と、「論文→実装」ワークフローを異なる題材で2回反復。次はさらに別のSOTA手法へ。

📝 自己評価(解いた後に記入)

自分の回答・気づき・メモ: