コース進捗 コース目次 34レッスン中 34件を公開中
第0部―まず地図を広げる
第1部―簡単な目的で、なぜ動画が分かるのか
第2部―一本の動画を一台のエンコーダーへ
第3部―実験を読めてこそ、論文を読んだと言える
第4部―公式リポジトリから自分の実験へ
第5部―表現を世界モデル構想へ戻す
付録―必要なときに開く技術リュック
長距離運転の前に五点検
- 形箱の寸法
- 有限値NaN/Inf なら止める
- 勾配本当に動くか
- 未来漏洩過去区間を守る
- 資源記憶量・処理量
おもちゃの列車を二周だけ
新しい列車をいきなり全国へ走らせず、まず二周。車両数は合うか、車輪は回るか、赤信号で止まるか。成功しても長距離性能は分かりません。それでも高価な学習前に配線ミスを捕まえられます。
五つの関門
一つ目は配列の形です。データ読み込み器の既定出力は、大域視点が [B,16,3,224,224]、局所視点が [B,10,16,3,96,96]。main.py がエンコーダー用の [B,C,T,H,W] へ軸を並べ替えます。学習時にパッチの5%を残すと、大域視点は3,136個のパッチから157個と [CLS] を合わせた158トークン、各局所視点は576個から29個と [CLS] を合わせた30トークンになります。
二つ目と三つ目は数値と勾配です。pred_loss、sigreg_loss、総損失がいずれも有限値であり、backward() の後にエンコーダーと投影器の両方へ有限かつ非ゼロの勾配が流れている必要があります。2バッチだけを見て、損失の単調減少までは求めません。無作為な切り抜きと射影を使う小さなバッチでは値が揺れるからです。小バッチへの過学習を試すなら、同じバッチと乱数種を固定し、表現の分散も記録します。MSEだけが下がる場合は、表現崩壊かもしれません。
データ型の受け渡しも調べます。既定の normalize_on_gpu=true では、読み込みワーカーが8ビット符号なし整数を送り、ホスト側とプロセス間通信の記憶量をおよそ四分の一に抑えます。main.to_float_normalized() が255で割り、ImageNet の統計量で正規化します。データ読み込み器がすでに浮動小数点数へ変換した入力をもう一度正規化すると、値を気づかないまま壊します。バッチについて、データ型、最小値と最大値、配列の形を表示してください。8ビット符号なし整数なら値域は [0,255]、エンコーダーへ入る直前は有限な浮動小数点数でなければなりません。大域切り抜きには空間切り抜きと正規化を施し、局所切り抜きにはさらに色ゆらぎ、グレースケール化、反転を加えます。ただし、両者が参照する時間添字は完全に一致させます。
四つ目はブロック因果マスクの直接検査です。データも事前学習済み重みも要りません。後半二フレームだけを変えたとき、前半二フレームのパッチトークンが変わらないことを確かめます。動画窓全体を読む [CLS] は比較対象にしません。
uv run python - <<'PY'
import torch
from module import vit_tiny
torch.manual_seed(0)
m = vit_tiny(img_size=32, patch_size=16, num_frames=4, tubelet_size=1,
use_rope=True, token_drop_rate=0, attn_mode="block_causal").eval()
x = torch.randn(1, 3, 4, 32, 32)
y = x.clone(); y[:, :, 2:] = torch.randn_like(y[:, :, 2:])
with torch.inference_mode():
a, b = m(x), m(y)
d = (a[:, 1:9] - b[:, 1:9]).abs().max().item() # first 2 frames × 4 patches
print(f"causal prefix max diff: {d:.3g}")
assert d < 1e-5
PY
五つ目は、GPU記憶使用量の最大値、一ステップの所要時間、実際のトークン数の記録です。記憶不足になったら batch_size や局所切り抜き数を下げるか vit_tiny にしますが、それは診断用の別実験です。読み込みワーカーが停止する場合は loader.num_workers=0 にします。NaN が出たら、入力のデータ型と正規化、学習率、混合精度、SIGReg に渡すバッチサイズ、多GPU間の同期という順で調べます。training_workshop.ipynb は、CUDA、MPS、CPUの各環境で、視点生成、モデル、二つの損失、分布図までを通す小規模な検査です。
合格条件は検査式として明文化します。配列の形が一致し、損失と勾配がすべて有限値であること。固定したバッチを50〜100ステップ学習したとき、損失の移動平均が初期値より低く、各次元の標準偏差がそろってゼロへ向かわないこと。未来側だけを変えたときの過去区間の差が許容誤差未満であり、同じ乱数種なら最初の1ステップの損失を再現できることです。
SIGReg は毎ステップ無作為方向を取り直すため、基準ログと小数点以下まですべて一致するわけではありません。問題なのは、値が持続的に発散すること、NaN が出ること、埋め込みの固有値スペクトルがほぼゼロになることです。
点検表の出典
車庫に貼る三行
- データ読み込み器は
[B,T,C,H,W]、エンコーダーは[B,C,T,H,W]を使います。 - 因果性の検査では、パッチトークンの過去区間を比べます。
[CLS]は動画窓全体を見る設計です。 - 2バッチを通せても、短い動作確認に合格しただけであり、精度を再現したことにはなりません。
整備士の三問
- 既定の大域視点は学習中に約何トークンを出しますか。
- MSE と一緒に分散とSIGRegを見る理由は何ですか。
- 未来側のフレームを変えたとき、
[CLS]も不変であるべきですか。
答え
- 158です。残したパッチ157個に
[CLS]を加えます。 - 全入力を一点へ写しても、不変性を測るMSEは下がってしまうからです。
- いいえ。
[CLS]は動画窓全体を読む要約です。未来側の変更から守られるのは、過去側のパッチトークンです。