📚 背景知識(読んでから問題へ)
通常の学習(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回目のforward/backward
現在地wtでの勾配∇L(w)を計算する
②ê(w)を計算し摂動
勾配を正規化して半径ρ分だけパラメータへ一時的に加算する
③2回目のforward/backward
摂動後の地点w+ê(w)で改めて勾配を計算する
④元の位置へ復元
本番更新の前に必ずパラメータをwtへ戻す
⑤本番の更新
③で得た勾配を使ってbase_optimizer.step()を実行
今日はこのワークフローを、少ない訓練データで過学習が起きやすい合成データを題材に体験する。
🎯 問題
Day126と同様の合成分類データ(ただし今回は訓練データを少なめにして過学習が起きやすい状況を作る)に対し、通常のAdamとSAM+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件
sam_step(model, loss_fn, x, y, base_optimizer, rho)として実装するsam_stepを使いSAM付きの学習を行う(rho=0.05)💡 ヒント
SAMの数式は2本ですが、実装上のポイントは「同じバッチに対して勾配計算を2回行う」点です。1回目の勾配で摂動方向ê(w)を計算し、パラメータを一時的にその方向へ動かし、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で繰り返すだけでよい
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 前後 → 平らな谷
🪜 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)
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()
base_optimizer.step()を呼ぶ前に必ずパラメータを元の位置へ戻す必要がある。ここを戻し忘れると、摂動が累積してパラメータが際限なく発散してしまう典型的な実装ミスになる。3rho=0のときSAMが通常のAdamと一致することを確認する
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()
🧮 数学・統計の補足(文系向け)
損失関数を「山(実際には谷なので下り坂)の地形」だと考えると、勾配(gradient)は「今立っている場所で、最も急に高くなる方向」を教えてくれるコンパスのようなものです。SAMのステップ①は、このコンパスが指す方向に「半径ρ分だけ」歩いて、近くで一番損失が高そうな地点まで移動する操作に当たります。L2ノルム(‖・‖₂)は「ベクトルの長さ」のことで、例えば2次元なら√(a²+b²)(三平方の定理そのもの)です。勾配ベクトルをそのノルムで割ることは「歩く方向は保ったまま、歩幅を『長さ1』にそろえる」操作であり、そこにρを掛けることで「正確に半径ρだけ動く」ことが保証されます。
ある生徒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テーマ)
Day126でLabel Smoothing(損失関数の改造)、Day127でSAM(Optimizerの改造)と、「論文→実装」ワークフローを異なる題材で2回反復。次はさらに別のSOTA手法へ。
📝 自己評価(解いた後に記入)
自分の回答・気づき・メモ: