JEPA4Japan · チュートリアル

第8章 — 完全な順伝播:1バッチを最初から最後まで追う

1,132文字 4分で読めます #LeWorldModel#World Models#JEPA

観測の符号化、行動条件、次状態予測、予測対象表現の構成、2つの損失、勾配の流れを、最小構成と形状検査で結び付けます。

コース進捗 コース目次 48レッスン中 48件を公開中

第0部 読み方ガイド:私たちは何を学ぶのか

  1. 01 第0章 — はじめる前に 公開中

第1部 世界モデル:エージェントの頭の中にある実験場

  1. 02 第1章 — なぜエージェントには「未来を想像する」力が必要なのか 公開中
  2. 03 第2章 — なぜ次の画像をそのまま予測しないのか 公開中
  3. 04 第3章 — JEPAの発想:画面の複製ではなく意味を予測する 公開中
  4. 05 第4章 — 1枚の図でLeWMを理解する 公開中

第2部 画面を状態に変える:LeWMのモデル構造

  1. 06 第5章 — 軌跡データ:モデルにとって世界は画像集ではない 公開中
  2. 07 第6章 — 視覚エンコーダー:各フレームに「状態パスポート」を発行する 公開中
  3. 08 第7章 — 動力学予測器:頭の中で時間を前へ進める 公開中
  4. 09 第8章 — 完全な順伝播:1バッチを最初から最後まで追う 現在のレッスン

第3部 モデルの抜け道を防ぐ:予測損失とSIGReg

  1. 10 第9章 — 最も危険な近道:表現崩壊 公開中
  2. 11 第10章 — 予測損失:モデルはどのように次の一歩を学ぶのか 公開中
  3. 12 第11章 — SIGRegの直感:表現空間に「呼吸」をさせる 公開中
  4. 13 第12章 — 必要最小限の数学 公開中
  5. 14 第13章 — オリジナルLeWMのエンドツーエンド学習の仕組み 公開中
  6. 15 第14章 — すぐに表現崩壊しないモデルを訓練する 公開中

第4部 モデルを行動に使う:潜在空間での計画

  1. 16 第15章 — 目標条件付き計画:「今いる場所」から「行きたい場所」へ 公開中
  2. 17 第16章 — 潜在ユークリッド距離:便利だが、常に信頼できるとは限らない 公開中
  3. 18 第17章 — CEM:勝ち抜き方式で行動を探索する 公開中
  4. 19 第18章 — MPC:モデルを一度に長く信じすぎない 公開中
  5. 20 第19章 — 長期ロールアウト:小さな誤差が大事故へ育つまで 公開中
  6. 21 第20章 — 最小のLeWMプランナーをゼロから実装する 公開中

第5部 エンジニアリング再現:論文から動くシステムへ

  1. 22 第21章 — 公式リポジトリと実験環境 公開中
  2. 23 第22章 — 最初の実験:TwoRoomのスモークテスト 公開中
  3. 24 第23章 — 2つ目の実験:PushTを再現する 公開中
  4. 25 第24章 — 世界モデルを公平に評価する方法 公開中
  5. 26 第25章 — 失敗診断マニュアル 公開中

第6部 LeWMは何を学んだのか

  1. 27 第26章 — 線形プローブ:潜在状態にはどの物理量が含まれるのか 公開中
  2. 28 第27章 — 潜在空間を「健康診断」する 公開中
  3. 29 第28章 — 期待違反:モデルは「あり得ない出来事」に驚くのか 公開中
  4. 30 第29章 — 「世界を理解する」を厳密に語るには 公開中

第7部 「予測が正確」でも「計画がうまくいかない」のはなぜか

  1. 31 第30章 — 訓練目的と計画目的のあいだにある亀裂 公開中
  2. 32 第31章 — 大域的には表現崩壊していなくても、タスクに必要な動力学が保たれるとは限らない 公開中
  3. 33 第32章 — 等方ガウス事前分布はいつ強すぎるのか 公開中
  4. 34 第33章 — 長期計画:より遠くを予測するか、より賢く計画するか 公開中
  5. 35 第34章 — 位置の距離からタスクの進捗へ 公開中
  6. 36 第35章 — マルチタスク、実ロボット、視覚的外乱 公開中
  7. 37 第36章 — 理論的な境界:真の状態はいつ同定できるのか 公開中

第8部 再現者から研究者へ

  1. 38 第37章 — 信頼できるLeWM改良実験を設計する 公開中
  2. 39 第38章 — 実行可能な12の研究課題 公開中
  3. 40 第39章 — LeWM研究の未解決問題 公開中

付録 数学・実装・再現・査読のための参照資料

  1. 41 付録A — 最低限必要な数学ツールキット 公開中
  2. 42 付録B — PyTorch実装クイックリファレンス 公開中
  3. 43 付録C — テンソル形状の完全一覧 公開中
  4. 44 付録D — 実験設定カード 公開中
  5. 45 付録E — 論文タイムラインとエビデンスレベル 公開中
  6. 46 付録F — 用語集 公開中
  7. 47 付録G — 再現チェックリスト 公開中
  8. 48 付録H — 専門家査読チェックリスト 公開中

まず大きな絵

  1. 4つencodez0 z1 z2 z3
  2. 3つ揃えるa0 a1 a2
  3. 3つ予測p1 p2 p3
  4. targetとzipp1↔z1, p2↔z2, p3↔z3
shapeが正しいだけでは足りません。全time indexとgradient routeが正しい物語を語る必要があります。

4つのobserved latentから3つのshifted prediction–target pairを作ります。

小さなお話

工場は正しい大きさの箱に、間違ったlabelを入れて出荷できます。LeWM batchも同じです。prediction 3本とtarget 3本のshapeが揃っていても、predictionをnextではなくpresentと比べたり、最後のpairだけ学習したり、targetをdetachしたりできます。

本章の約束はprovenanceです。各slotで、どのobservation、action block、prediction、targetが対応するかを指差せるようにします。

技術バックパック

frozen defaultはforward pass全体を具体化します。batch B=128、history_size=3、num_preds=1、latent width D=192、frameskip=5です。

pixels                 [B,4,C,H,W]
flattened pixels       [B×4,C,H,W]
observation latents    [B,4,192]
raw action blocks      [B,4,5×A]
encoded actions        [B,4,192]
state/action context   [B,3,192]
connected targets      [B,3,192]
predictions            [B,3,192]
SIGReg input           [4,B,192]
losses                 scalar

paperはTwoRoomのhistoryを1、PushT/OGBench-Cubeを3と報告し、frozen global defaultは3です。これらはfrozen defaultのshapeであり、全LeWMの定数ではありません。

最短のfaithful forward cardは次です。

z = encode_each_frame(images)              # [B,4,D]
u = encode_each_action_block(actions)       # [B,4,D]
z_pred = causal_predict(z[:,:3], u[:,:3])  # [B,3,D]
z_next = z[:,1:]                            # [B,3,D], detachしない
L_pred = mean_square_gap(z_pred, z_next)
L_sig = sigreg(time_first(z))               # [4,B,D]
L_total = L_pred + lambda * L_sig

3つのshifted one-step pairすべてが寄与します。num_preds=1はone-position target offsetであり、最後の1出力ではありません。このdata configの1 model stepは5 environment actionsをまたぎます。

prediction branchはvisual encoder、encoder projector、action encoder、causal predictor、prediction projectorを直接学習します。z_nextはconnectedなので、predictionは同じvisual encoderのtarget useからも戻ります。SIGRegが直接学習するのはvisual encoderとそのprojectorだけです。1 optimizerがconnected learned systemを更新し、target-encoder optimizerやEMA updateはありません。

SIGRegは各relative time positionでbatch populationを見ます。B×Tを1 populationへflattenしてはいけません。temporal changeとacross-example diversityは別物です。frozen codeは1,024 random projectionsと0〜3の17 frequency knotsを使います。paper appendixのquadrature presentationは異なります。paper method weightは0.1、frozen YAMLは0.09です。source labelを残します。

だまされる仕掛け

わざと間違えたbranchを作ります。

wrong targets: [z0,z1,z2]
right targets: [z1,z2,z3]

どちらも[B,3,D]です。wrong branchはpresentのcopyを学び、よいlossを出せます。time checksumを見える形で付け、prediction_slot / last_visible_context / target_sourceをprintします。正しいrowは0/0/1、1/1/2、2/2/3です。

次にgraphを試します。disposable probeで正しいtarget sliceだけdetachし、target-side gradient routeが消え、context routeが残ることを確かめます。encoderに何らかのgradientがあるだけでは十分ではありません。

5つのcheap checkは、exact shape、shift provenance、causal future-mutation、expected moduleのfinite/non-null gradient、NaN/infinityなしです。通過が証明するのはmechanicsであり、有用なdynamicsやplanningではありません。

実験レシート

LeWorldModel v3、固定train.py、jepa.py、module.pyがfrozen alignmentとconnected graphを裏づけます。paperのdisplayed training pseudocodeにはmalformed F.mse_loss callがあるため、exact indexingはexecutable train.pyを典拠にします。

next targetをplanning goalと呼ぶ、JEPAの癖でdetachする、SIGRegでbatch/timeを混ぜる、matching shapeをcorrectnessとみなす、のはいずれも誤りです。

3問クイックチェック

  1. 4 observationsから作る3 prediction–target pairsは何ですか。
  2. prediction lossはvisual encoderのどの2つのuseから流れますか。
  3. wrong temporal targetが通常のshape assertionを全部通れるのはなぜですか。