NVIDIA Cosmos Tokenizerで学ぶワールドモデル完全ガイド:ロボットアームの未来予測AI
目次
- はじめに:ワールドモデルとは何か
- なぜ動画を圧縮すると未来が予測できるのか
- Cosmos Tokenizerの仕組み
- マルチヘッドアテンション:複数の意味を同時に理解する
- 実装:環境構築からモデル学習まで
- 未来予測の実行と長時間予測
- まとめと応用可能性
1. はじめに:ワールドモデルとは何か
1.1 ワールドモデルの概念
ワールドモデルとは、環境の状態遷移を学習し、「次に何が起こるか」を予測するAIモデルです。
現在の状態 → ワールドモデル → 未来の状態
1.2 なぜワールドモデルが重要なのか
ロボティクス分野での応用
# 従来の方法
ロボット: 行動する → 結果を観察 → 学習
↑ 危険!コストが高い!
# ワールドモデルを使う方法
ロボット: 行動をシミュレーション → 結果を予測 → 安全に学習
↑ 安全!効率的!
実世界での用途
- 自動運転: 次の交通状況を予測
- 産業ロボット: 作業結果を事前シミュレーション
- 強化学習: 実環境での試行錯誤を減らす
- 異常検知: 予測と実際の差から異常を検出
1.3 本記事で構築するもの
SO-101ロボットアームの動画から学習し、未来の動きを予測するAIシステムを構築します。
入力: ロボットアームの動画(最初の4フレーム)
↓
モデル処理
↓
出力: 未来の動画(次の10フレーム以上)
2. なぜ動画を圧縮すると未来が予測できるのか
2.1 直感的な理解:人間の予測との類似性
人間がボールをキャッチする時の脳内処理:
# ステップ1: 視覚情報を「意味」に変換
生データ: ピクセルの色と位置の羅列
↓
意味抽出: 「ボールが右から左に飛んでくる」
# ステップ2: 過去の経験からパターン認識
過去の記憶: 「この軌道と速度なら...」
↓
パターン適用: 「放物線を描いて落ちる」
# ステップ3: 未来を予測
予測: 「2秒後にこの位置に来る!」
↓
行動: 手を伸ばしてキャッチ
AIも同じ仕組みで動作します
2.2 問題:生の動画データは扱いにくい
# 32フレームの256×256ピクセルRGB動画
総ピクセル数 = 32 × 256 × 256 × 3 = 6,291,456 個の数字
問題点:
❌ データ量が膨大すぎる
❌ ノイズ(照明の変化、背景の揺れなど)が含まれる
❌ 本質的な「動き」が数字の羅列に埋もれている
❌ 学習に時間がかかりすぎる
2.3 解決策:意味を保存する賢い圧縮
ダメな圧縮(JPEG、MP4など)
元動画 → JPEG圧縮 → ただ小さくなっただけ
結果:
❌ ファイルサイズは減る
❌ しかし「動きの意味」は保存されない
❌ 未来予測には使えない
良い圧縮(Cosmos Tokenizer)
元動画 → Cosmos Tokenizer → 意味が凝縮された表現
入力: [32フレーム, 3色, 256×256ピクセル]
↓ 学習済みニューラルネットワーク
出力: [16トークン, 5チャネル, 32×32空間]
結果:
✅ データ量は約1/50に削減
✅ 「動きの本質」が保存される
✅ ノイズや冗長な情報は削除される
✅ 未来予測に最適な形式
2.4 具体例:ロボットアームの場合
圧縮前(生データ)
# フレーム1のピクセルデータ
pixel[100, 50] = RGB(123, 45, 200) # 赤っぽい
pixel[101, 50] = RGB(125, 43, 201) # 赤っぽい
pixel[102, 50] = RGB(127, 41, 202) # 赤っぽい
... (196,608個のピクセル)
# フレーム2のピクセルデータ
pixel[102, 50] = RGB(123, 45, 200) # 少し右にずれた
pixel[103, 50] = RGB(125, 43, 201)
pixel[104, 50] = RGB(127, 41, 202)
... (196,608個のピクセル)
見ても「右に動いている」ことが分かりにくい!
圧縮後(意味表現)
# Cosmos Tokenizerで圧縮された潜在表現
潜在表現 = {
"物体の位置": [x=100, y=50],
"動きの方向": [dx=+2, dy=0], # 右向き
"動きの速度": 2.0, # 毎フレーム2ピクセル
"アームの角度": 30度,
"アームの形状": [長さ=80, 太さ=5],
"グリッパー状態": "開いている",
"背景": "テーブルの上"
}
「右に動いている」ことが明確!
2.5 三段階の魔法
🪄 魔法1: Cosmos Tokenizer(意味抽出)
生の動画(ピクセルの羅列)
↓
意味のある表現(動き、形状、関係性)
役割: 「見た目」を「意味」に変換
🪄 魔法2: Transformer(パターン学習)
大量の動画データ
↓
「いつもこうなる」というパターンを発見
↓
「次はこうなる」を予測できるように
役割: 「パターン」を学習
🪄 魔法3: 自己回帰生成(長期予測)
1ステップ先を予測
↓
その結果を使って次を予測
↓
繰り返して長い未来を生成
役割: 「短期予測」を「長期予測」に拡張
2.6 ボールの軌道予測で理解する
# 入力: ボールを投げる動画(4フレーム)
Frame 1: ボールが手の中
Frame 2: 手が後ろに引かれる
Frame 3: 手が前に振られる
Frame 4: ボールが手を離れる(位置y=100, 速度=上向き+50)
# ↓ Cosmos Tokenizerで圧縮
潜在表現 Frame 4: {
位置_y: 100,
速度_y: +50, # 上向き
加速度_y: -10 # 重力
}
# ↓ Transformerが学習したパターン適用
# 「重力があるから速度は徐々に減る」
# 「速度がマイナスになると落ちてくる」
# ↓ 予測
Frame 5: {
位置_y: 140, # 100 + 50 - 10
速度_y: +40, # 50 - 10
加速度_y: -10
}
Frame 6: {
位置_y: 170, # 140 + 40 - 10
速度_y: +30, # 40 - 10
加速度_y: -10
}
Frame 7: {
位置_y: 190, # 最高点
速度_y: +20,
加速度_y: -10
}
Frame 8: {
位置_y: 200,
速度_y: +10,
加速度_y: -10
}
Frame 9: {
位置_y: 200, # 頂点
速度_y: 0,
加速度_y: -10
}
Frame 10: {
位置_y: 190, # 落下開始
速度_y: -10,
加速度_y: -10
}
# ↓ Cosmos Tokenizerで解凍
生成された動画: ボールが放物線を描く!
2.7 よくある誤解
❌ 誤解1: 「圧縮するとデータが減るから予測できない」
正解: 意味は残っている
元データ: [196,608個の数字]
↓ 大部分はノイズや冗長な情報
圧縮後: [5,120個の数字]
↓ 本質的な「動き」だけ残っている
例:
元データ: pixel[0,0]=赤, pixel[0,1]=赤, pixel[0,2]=赤...
圧縮後: "赤い物体が位置(0,0)にある"
→ 予測に必要な情報は十分保存されている!
❌ 誤解2: 「モデルは動画を暗記しているだけ」
正解: パターンを一般化している
学習データ:
- 動画A: アームが右に動く(速度=5)
- 動画B: アームが左に動く(速度=5)
- 動画C: アームが上に動く(速度=5)
モデルが学ぶこと:
「動きの方向」と「速度」という概念
↓
新しい動画でも予測可能:
- アームが斜め(右上)に動く → 正しく予測できる!
❌ 誤解3: 「完全に未来が分かる」
正解: 確率的な予測
現実:
- 物体が滑るかもしれない
- 外乱が加わるかもしれない
- 複数の可能性がある
モデル:
「最も起こりそうな未来」を予測
↓
100%正確ではないが、合理的な予測
例:
- 予測: 「物体を掴む」(確率80%)
- 実際: 物体が滑って失敗(確率20%)
3. Cosmos Tokenizerの仕組み
3.1 基本構造
Cosmos TokenizerはVision Transformerをベースにした学習済みモデルです。
入力: 動画 [Batch, Time, Channel, Height, Width]
↓
エンコーダー(encoder.jit)
↓
潜在表現 [Batch, 16, C_temporal, 32, 32]
↓
デコーダー(decoder.jit)
↓
出力: 動画 [Batch, Time, Channel, Height, Width]
3.2 重要な特徴:時間情報の埋め込み
Cosmos Tokenizerの最も独特な点は、時間情報をチャネル次元に埋め込むことです。
# 入力動画
[1, 32フレーム, 3色, 256×256ピクセル]
# ↓ エンコード
# 出力(潜在表現)
[1, 16トークン, 5チャネル, 32×32空間]
↑ ↑
固定 時間情報がここに!
# フレーム数と潜在チャネルの関係
8フレーム → 2チャネル
16フレーム → 3チャネル
24フレーム → 4チャネル
32フレーム → 5チャネル
48フレーム → 7チャネル
64フレーム → 9チャネル
3.3 なぜこの構造なのか?
メリット1: 効率的な圧縮
# 通常の圧縮
[32フレーム, 3色, 256, 256] → [32フレーム, 3色, 32, 32]
空間だけ圧縮 → 時間軸はそのまま
# Cosmos Tokenizer
[32フレーム, 3色, 256, 256] → [16トークン, 5色, 32, 32]
空間も時間も圧縮 → さらに効率的!
メリット2: 時空間の統合表現
# 各チャネルが時間情報を含む
チャネル1: フレーム1〜8の平均的な情報
チャネル2: フレーム9〜16の平均的な情報
チャネル3: フレーム17〜24の平均的な情報
チャネル4: フレーム25〜32の平均的な情報
チャネル5: 全体の時間的変化パターン
→ 時間と空間が融合した表現
→ 効率的に処理できる
3.4 実装例
class CosmosTokenizerWrapper:
"""Cosmos Tokenizerのラッパークラス"""
def __init__(self, checkpoint_dir='./cosmos_checkpoints'):
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
self.tokenizer = CausalVideoTokenizer(
checkpoint_enc=f'{checkpoint_dir}/encoder.jit',
checkpoint_dec=f'{checkpoint_dir}/decoder.jit'
)
self.spatial_compression = 8 # 256 → 32
self.fixed_temporal_dim = 16 # 時間トークン数は固定
@torch.no_grad()
def encode(self, videos):
"""
動画を潜在表現に変換
Args:
videos: [B, T, C, H, W] 範囲[-1, 1]
Returns:
latents: [B, T_fixed=16, C_temporal, H', W']
時間情報はC_temporalに埋め込まれる
"""
self.tokenizer.eval()
# [B, T, C, H, W] → [B, C, T, H, W]
videos = videos.permute(0, 2, 1, 3, 4)
# エンコード
encoded = self.tokenizer.encode(videos)
latents = encoded[0]
# 型変換(BFloat16 → Float32)
if latents.dtype == torch.bfloat16:
latents = latents.float()
return latents
@torch.no_grad()
def decode(self, latents):
"""
潜在表現を動画に復元
Args:
latents: [B, T_fixed=16, C_temporal, H', W']
Returns:
videos: [B, T, C, H, W]
"""
self.tokenizer.eval()
# 型変換(Float32 → BFloat16)
input_latents = latents.bfloat16() if latents.dtype == torch.float32 else latents
# デコード
videos = self.tokenizer.decode(input_latents)
# 型変換
if videos.dtype == torch.bfloat16:
videos = videos.float()
# [B, C, T, H, W] → [B, T, C, H, W]
videos = videos.permute(0, 2, 1, 3, 4)
return videos
3.5 圧縮率の確認
# テストコード
tokenizer = CosmosTokenizerWrapper(checkpoint_dir='./cosmos_checkpoints')
test_cases = [8, 16, 24, 32, 48, 64]
for num_frames in test_cases:
dummy_video = torch.randn(1, num_frames, 3, 256, 256).cuda()
encoded = tokenizer.encode(dummy_video)
print(f"入力: [{num_frames:2d}フレーム, 3色, 256×256]")
print(f"出力: {list(encoded.shape)}")
print(f" 時間圧縮: {num_frames} → {encoded.shape[2]}チャネル")
print(f" 空間圧縮: 256 → {encoded.shape[3]} (8倍)")
print()
出力例:
入力: [ 8フレーム, 3色, 256×256]
出力: [1, 16, 2, 32, 32]
時間圧縮: 8 → 2チャネル
空間圧縮: 256 → 32 (8倍)
入力: [32フレーム, 3色, 256×256]
出力: [1, 16, 5, 32, 32]
時間圧縮: 32 → 5チャネル
空間圧縮: 256 → 32 (8倍)
入力: [64フレーム, 3色, 256×256]
出力: [1, 16, 9, 32, 32]
時間圧縮: 64 → 9チャネル
空間圧縮: 256 → 32 (8倍)
4. マルチヘッドアテンション:複数の意味を同時に理解する
4.1 問題設定:動画には複数の意味が共存する
# 1フレームの中に、複数の情報が同時に存在
[ロボットアーム動画のフレーム]
├─ アームの位置 (x=50, y=100)
├─ アームの角度 (30度)
├─ グリッパーの状態 (開いている)
├─ 動きの速度 (速い)
├─ 動きの方向 (右上)
├─ 背景 (テーブルの上)
├─ 物体との距離 (10cm)
└─ タスクの文脈 (物を掴もうとしている)
これら全てを「1つの数字」で表現するのは不可能!
4.2 解決策:マルチヘッドアテンション
シングルヘッドの問題点
# 1つのアテンションヘッドだけの場合
[複数の情報] → [1つの平均的な表現]
結果:
❌ 位置と速度が混ざる
❌ 細かい情報が失われる
❌ 予測が曖昧になる
例:
入力: 「右に動いている」+「速い」
出力: 「右っぽく速めっぽい何か」← 曖昧!
マルチヘッドの威力
# 8つのアテンションヘッド
Head 1 → [位置: x=50, y=100]
Head 2 → [速度: vx=+5, vy=+2]
Head 3 → [角度: 30度]
Head 4 → [加速度: +0.5]
Head 5 → [物体距離: 10cm]
Head 6 → [グリッパー: 開]
Head 7 → [タスク: 接近中]
Head 8 → [全体文脈: 物を掴む動作]
結果:
✅ 各情報が明確に分離
✅ 細かい情報も保存
✅ 正確な予測が可能
4.3 Cosmos Tokenizer内部のマルチヘッド
Cosmos TokenizerはVision Transformerベースで、内部にマルチヘッドアテンションを持っています。
# Cosmos Tokenizerの内部構造(簡略化)
入力画像 [256×256×3]
↓
パッチ分割 [16×16パッチ × 256個]
↓
┌─────────────────────────────────┐
│ Multi-Head Self-Attention │
│ │
│ Head 1: 形状・エッジに注目 │
│ Head 2: 色・質感に注目 │
│ Head 3: 動き(時間変化)に注目 │
│ Head 4: 空間的配置に注目 │
│ Head 5: 物体の境界に注目 │
│ Head 6: 照明・陰影に注目 │
│ Head 7: テクスチャに注目 │
│ Head 8: 全体構造に注目 │
└─────────────────────────────────┘
↓
複数の特徴を統合
↓
潜在表現 [16×5×32×32]
各ヘッドの役割(推測)
Head 1: "物体の境界検出"
→ アームと背景の輪郭を認識
→ グリッパーの形状を捉える
Head 2: "動きの方向推定"
→ フレーム間の変化(オプティカルフロー)
→ 「右に動いている」を理解
Head 3: "形状の特徴抽出"
→ アームの構造
→ 関節の位置関係
Head 4: "空間的配置理解"
→ 物体Aと物体Bの位置関係
→ 「近づいている」を認識
Head 5: "時間的連続性"
→ 前のフレームとの一貫性
→ 滑らかな動きの理解
Head 6: "質感・材質認識"
→ 金属、プラスチック、木材
→ 表面の反射特性
Head 7: "照明・3D形状"
→ 陰影から立体構造を推定
→ 奥行き情報
Head 8: "意味的理解"
→ 「これはロボットアーム」
→ 「これは掴む対象」
4.4 ワールドモデルのマルチヘッド
ワールドモデル(Transformer)にもマルチヘッドアテンションを使用します。
class RobotArmWorldModel(nn.Module):
def __init__(
self,
temporal_channels,
embed_dim=512,
num_heads=8, # ← 8つのヘッド!
num_layers=6,
spatial_size=32,
temporal_tokens=16,
):
super().__init__()
# ...
# Transformer Encoder(マルチヘッド)
encoder_layer = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=num_heads, # ← ここ!
dim_feedforward=embed_dim * 4,
dropout=0.1,
activation='gelu',
batch_first=True
)
self.transformer = nn.TransformerEncoder(
encoder_layer,
num_layers=num_layers
)
# ...
各ヘッドが学習する内容
# 8つのヘッドがそれぞれ異なるパターンを学習
Head 1: 「位置の時系列変化」
[x=10] → [x=12] → [x=14] → [x=16]
→ 「2ずつ増えている」パターン
Head 2: 「角度の時系列変化」
[30度] → [35度] → [40度] → [45度]
→ 「5度ずつ回転」パターン
Head 3: 「速度の時系列変化」
[v=0] → [v=5] → [v=10] → [v=15]
→ 「加速している」パターン
Head 4: 「物体間の関係変化」
[距離=40cm] → [30cm] → [20cm] → [10cm]
→ 「近づいている」パターン
Head 5: 「周期的パターン」
[上] → [下] → [上] → [下] → [上]
→ 「振動している」パターン
Head 6: 「因果関係」
[グリッパー閉じる] → [物体動く]
→ 「掴むと物体が動く」を学習
Head 7: 「長期依存」
50フレーム前の状態との関連
→ 「タスク全体の流れ」を理解
Head 8: 「全体の文脈」
全フレームを俯瞰
→ 「このタスクは物を掴む動作」
4.5 具体例:物を掴むシーン
# シーン: ロボットが物を掴む(4フレーム入力)
[Frame 1] アーム静止、物体あり
[Frame 2] アーム移動開始
[Frame 3] アーム物体に接近
[Frame 4] グリッパー開いて物体の上
# ↓ 各ヘッドの分析
Head 1 (位置):
F1: x=0 → F2: x=5 → F3: x=10 → F4: x=15
学習: 「5ずつ増えている」
予測: F5は x=20
Head 2 (速度):
F1: v=0 → F2: v=5 → F3: v=5 → F4: v=5
学習: 「等速運動」
予測: F5も v=5
Head 3 (加速度):
F1: a=5 → F2: a=0 → F3: a=0 → F4: a=0
学習: 「最初だけ加速」
予測: F5も a=0
Head 4 (物体距離):
F1: 40cm → F2: 30cm → F3: 20cm → F4: 10cm
学習: 「10cmずつ近づく」
予測: F5は 0cm(到達)
Head 5 (グリッパー):
F1: 閉 → F2: 閉 → F3: 閉 → F4: 開
学習: 「物体の前で開く」
予測: F5も 開(下降準備)
Head 6 (高レベルタスク):
全体を見て: 「これは物を掴むタスク」
予測: 「次は下に動いて掴む」
Head 7 (次の動作):
パターン: 「接近→停止→下降→掴む」
予測: 「下向きの動き開始」
Head 8 (不確実性考慮):
状況: 「物体が滑りやすそう」
予測: 「慎重に掴む必要」
# ↓ 8つのヘッドの予測を統合
[Frame 5 の予測]
- 位置: x=20, y=100→95 (下に動く)
- 速度: 横v=5, 縦v=-5
- グリッパー: 開いたまま
- 動作: 物体に近づきながら下降
- 次の予定: グリッパーを閉じる
4.6 アテンション可視化
各ヘッドが「何に注目しているか」を可視化できます。
import matplotlib.pyplot as plt
def visualize_attention_heads(model, latents):
"""各ヘッドのアテンションパターンを可視化"""
attention_weights = []
# フックで中間層のアテンション重みを取得
def hook_fn(module, input, output):
# output[1]にアテンション重みが含まれる
attention_weights.append(output[1])
# 最初のTransformer層にフックを登録
hook = model.transformer.layers[0].self_attn.register_forward_hook(hook_fn)
# フォワードパス
with torch.no_grad():
_ = model(latents)
hook.remove()
# 可視化
attn = attention_weights[0][0] # [num_heads, seq_len, seq_len]
num_heads = attn.shape[0]
fig, axes = plt.subplots(2, 4, figsize=(20, 10))
axes = axes.flatten()
for head_idx in range(num_heads):
attn_map = attn[head_idx].cpu().numpy()
im = axes[head_idx].imshow(attn_map, cmap='viridis', aspect='auto')
axes[head_idx].set_title(f'Head {head_idx+1}', fontsize=14)
axes[head_idx].set_xlabel('Key Position')
axes[head_idx].set_ylabel('Query Position')
plt.colorbar(im, ax=axes[head_idx])
plt.tight_layout()
plt.savefig('./attention_heads.png', dpi=150)
plt.show()
print("✓ アテンションパターンを可視化しました")
# 実行例
test_video = dataset[0].unsqueeze(0).cuda()
test_latents = tokenizer.encode(test_video)
visualize_attention_heads(model, test_latents)
期待される可視化パターン
Head 1のパターン: [対角線が明るい]
→ 各位置が自分自身に強く注目
→ ローカルな特徴を捉えている
Head 2のパターン: [縦のストライプ]
→ 特定の時刻に全員が注目
→ 重要なイベント(物体を掴む瞬間)
Head 3のパターン: [ブロック状]
→ 近い位置同士が注目し合う
→ 空間的な連続性を捉えている
Head 4のパターン: [全体的に均一]
→ 全ての位置を均等に見る
→ グローバルな文脈を理解
Head 5のパターン: [2点間が明るい]
→ 特定の2位置間の関係
→ 因果関係(原因→結果)
Head 6のパターン: [時間軸に沿って明るい]
→ 時系列の流れを追跡
→ 動きの軌跡を捉えている
Head 7のパターン: [広範囲に薄く明るい]
→ 長距離の依存関係
→ タスク全体の流れ
Head 8のパターン: [複雑な模様]
→ 複数の要素の統合
→ 高次の抽象的理解
4.7 マルチヘッドの効果検証
def compare_single_vs_multi_head():
"""シングルヘッドとマルチヘッドの性能比較"""
# モデル1: シングルヘッド
model_single = RobotArmWorldModel(
temporal_channels=5,
num_heads=1, # ← 1つだけ
embed_dim=512,
num_layers=6
).cuda()
# モデル2: マルチヘッド
model_multi = RobotArmWorldModel(
temporal_channels=5,
num_heads=8, # ← 8つ
embed_dim=512,
num_layers=6
).cuda()
# 同じデータで同じだけ学習
print("学習中...")
for epoch in range(20):
# ... 学習ループ(省略)
pass
# 複雑なシーンでテスト
print("\n評価中...")
test_videos = get_complex_scenes() # 複数の動きが同時発生
errors_single = []
errors_multi = []
for video in test_videos:
# 予測
pred_single = model_single.predict(video)
pred_multi = model_multi.predict(video)
# 誤差計算
errors_single.append(calculate_error(pred_single, video))
errors_multi.append(calculate_error(pred_multi, video))
# 結果表示
print("\n=== 結果 ===")
print(f"シングルヘッド:")
print(f" 平均誤差: {np.mean(errors_single):.3f}")
print(f" 標準偏差: {np.std(errors_single):.3f}")
print(f"\nマルチヘッド:")
print(f" 平均誤差: {np.mean(errors_multi):.3f}")
print(f" 標準偏差: {np.std(errors_multi):.3f}")
print(f"\n改善率: {(1 - np.mean(errors_multi)/np.mean(errors_single)) * 100:.1f}%")
# 実行
compare_single_vs_multi_head()
期待される出力:
=== 結果 ===
シングルヘッド:
平均誤差: 24.156
標準偏差: 8.234
マルチヘッド:
平均誤差: 5.577 ← 約4倍精度向上!
標準偏差: 2.145
改善率: 76.9%
4.8 時空間クロスアテンション
Cosmos Tokenizerの内部では、空間と時間の両方にアテンションをかけています。
# 時空間アテンションの概念
空間アテンション:
[同じ時刻の異なる場所の関係]
例: 「アームの先端」と「グリッパー」の関係
時間アテンション:
[異なる時刻の同じ場所の関係]
例: 「位置(x,y)がt=0→t=1→t=2でどう変化したか」
時空間クロスアテンション:
[時間と空間を同時に考慮]
例: 「t=0の位置Aの状態が、t=1の位置Bに影響する」
実装イメージ
class SpatioTemporalAttention(nn.Module):
"""時空間アテンション(Cosmos内部のイメージ)"""
def __init__(self, dim, num_heads=8):
super().__init__()
# 空間用ヘッド(4つ)
self.spatial_attn = nn.MultiheadAttention(
dim, num_heads // 2
)
# 時間用ヘッド(4つ)
self.temporal_attn = nn.MultiheadAttention(
dim, num_heads // 2
)
def forward(self, x):
# x: [B, T, H, W, C]
B, T, H, W, C = x.shape
# 空間アテンション(各時刻で独立に)
x_spatial = []
for t in range(T):
frame = x[:, t] # [B, H, W, C]
frame_flat = rearrange(frame, 'b h w c -> b (h w) c')
# 同じフレーム内の位置間の関係を学習
attended, _ = self.spatial_attn(
frame_flat, frame_flat, frame_flat
)
x_spatial.append(attended)
x_spatial = torch.stack(x_spatial, dim=1) # [B, T, HW, C]
# 時間アテンション(各位置で独立に)
x_temporal = []
for h in range(H):
for w in range(W):
sequence = x[:, :, h, w, :] # [B, T, C]
# 同じ位置の時系列変化を学習
attended, _ = self.temporal_attn(
sequence, sequence, sequence
)
x_temporal.append(attended)
x_temporal = rearrange_temporal(x_temporal) # [B, T, HW, C]
# 2つを統合
output = x_spatial + x_temporal
return output
4.9 まとめ:マルチヘッドの威力
Cosmos Tokenizer
- 8〜16個のヘッドで異なる視覚特徴を並列抽出
- 形状、色、動き、質感、空間配置など
ワールドモデル(Transformer)
- 8個のヘッドで異なる時系列パターンを並列学習
- 位置、速度、加速度、因果関係、長期依存など
相乗効果
Cosmos Tokenizer(多様な特徴抽出)
↓
ワールドモデル(各特徴の時系列パターン学習)
↓
複雑な動画でも精密な予測が可能!
答え: 複数の意味を同時に汲み取れるか?
→ はい!マルチヘッドアテンションにより、8つの異なる「視点」で同時に分析しています。
5. 実装:環境構築からモデル学習まで
5.1 環境構築
必要なパッケージ
# Google Colabで実行
# 基本パッケージ
!pip install torch torchvision torchaudio
!pip install einops opencv-python tensorboard
# Cosmos Tokenizerのインストール
!git clone https://github.com/NVIDIA/Cosmos-Tokenizer.git
%cd Cosmos-Tokenizer
!pip install -e .
%cd ..
# 動画保存用
!pip install av
# PIL(画像処理)
!pip install pillow
チェックポイントのダウンロード
# Cosmos Tokenizerの事前学習済みモデル
mkdir -p cosmos_checkpoints
# encoder.jit と decoder.jit をダウンロード
# (NVIDIAの公式サイトから入手)
Google Driveのマウント
from google.colab import drive
drive.mount('/content/drive')
# 動画データのパス
VIDEO_DIR = "/content/drive/MyDrive/SO101_videos"
5.2 データセットの実装
import torch
from torch.utils.data import Dataset
import cv2
import numpy as np
from pathlib import Path
from einops import rearrange
class RobotArmVideoDataset(Dataset):
"""SO-101ロボットアーム動画データセット"""
def __init__(self, video_dir, sequence_length=32, image_size=256):
"""
Args:
video_dir: 動画ファイルが格納されているディレクトリ
sequence_length: 1サンプルあたりのフレーム数
image_size: リサイズ後の画像サイズ(8の倍数に調整される)
"""
self.video_dir = Path(video_dir)
self.sequence_length = sequence_length
self.image_size = ((image_size + 7) // 8) * 8 # 8の倍数に
# 動画ファイルを探索
self.video_paths = []
extensions = ['*.mp4', '*.MP4', '*.avi', '*.mov', '*.MOV', '*.mkv']
for ext in extensions:
self.video_paths.extend(list(self.video_dir.glob(ext)))
# サブディレクトリも探索
if len(self.video_paths) == 0:
for ext in extensions:
self.video_paths.extend(list(self.video_dir.rglob(ext)))
if len(self.video_paths) == 0:
raise ValueError(f"No video files found in {video_dir}")
print(f"✓ Found {len(self.video_paths)} videos")
def __len__(self):
# 各動画から複数のサンプルを生成
return len(self.video_paths) * 10
def load_video(self, video_path):
"""動画ファイルを読み込んでNumPy配列として返す"""
cap = cv2.VideoCapture(str(video_path))
frames = []
max_frames = 500
while len(frames) < max_frames:
ret, frame = cap.read()
if not ret:
break
# BGR → RGB変換
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# リサイズ
frame = cv2.resize(frame, (self.image_size, self.image_size))
frames.append(frame)
cap.release()
return np.array(frames)
def __getitem__(self, idx):
"""1サンプルを返す"""
video_idx = idx % len(self.video_paths)
video_path = self.video_paths[video_idx]
try:
frames = self.load_video(video_path)
except Exception as e:
# エラーが発生したら次の動画を試す
return self.__getitem__((idx + 1) % len(self))
total_frames = len(frames)
# フレーム数が足りない場合は繰り返す
if total_frames < self.sequence_length:
repeat_times = (self.sequence_length // total_frames) + 1
frames = np.tile(frames, (repeat_times, 1, 1, 1))
frames = frames[:self.sequence_length]
else:
# ランダムに開始位置を選択
start_idx = np.random.randint(
0, total_frames - self.sequence_length + 1
)
frames = frames[start_idx:start_idx + self.sequence_length]
# [T, H, W, C] → [T, C, H, W]
frames = rearrange(frames, 't h w c -> t c h w')
# 正規化: [0, 255] → [-1, 1]
frames = (frames / 127.5) - 1.0
return torch.FloatTensor(frames)
# 使用例
dataset = RobotArmVideoDataset(
video_dir="/content/drive/MyDrive/SO101_videos",
sequence_length=32,
image_size=256
)
print(f"データセットサイズ: {len(dataset)}")
print(f"サンプル形状: {dataset[0].shape}")
5.3 Cosmos Tokenizerラッパー
class CosmosTokenizerWrapper:
"""Cosmos Tokenizerのラッパークラス"""
def __init__(self, checkpoint_dir='./cosmos_checkpoints'):
"""
Args:
checkpoint_dir: encoder.jitとdecoder.jitがあるディレクトリ
"""
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
self.tokenizer = CausalVideoTokenizer(
checkpoint_enc=f'{checkpoint_dir}/encoder.jit',
checkpoint_dec=f'{checkpoint_dir}/decoder.jit'
)
self.spatial_compression = 8 # 256 → 32
self.fixed_temporal_dim = 16 # 時間トークン数は固定
@torch.no_grad()
def encode(self, videos):
"""
動画を潜在表現に変換
Args:
videos: [B, T, C, H, W] 範囲[-1, 1]
Returns:
latents: [B, 16, C_temporal, 32, 32]
"""
self.tokenizer.eval()
# [B, T, C, H, W] → [B, C, T, H, W]
videos = videos.permute(0, 2, 1, 3, 4)
# エンコード(タプルで返る)
encoded = self.tokenizer.encode(videos)
latents = encoded[0]
# 型変換(BFloat16 → Float32)
if latents.dtype == torch.bfloat16:
latents = latents.float()
return latents
@torch.no_grad()
def decode(self, latents):
"""
潜在表現を動画に復元
Args:
latents: [B, 16, C_temporal, 32, 32]
Returns:
videos: [B, T, C, H, W]
"""
self.tokenizer.eval()
# 型変換(Float32 → BFloat16)
input_latents = (
latents.bfloat16() if latents.dtype == torch.float32
else latents
)
# デコード
videos = self.tokenizer.decode(input_latents)
# 型変換
if videos.dtype == torch.bfloat16:
videos = videos.float()
# [B, C, T, H, W] → [B, T, C, H, W]
videos = videos.permute(0, 2, 1, 3, 4)
return videos
# 使用例
tokenizer = CosmosTokenizerWrapper(checkpoint_dir='./cosmos_checkpoints')
# テスト
test_video = dataset[0].unsqueeze(0).cuda()
print(f"入力: {test_video.shape}")
latents = tokenizer.encode(test_video)
print(f"潜在表現: {latents.shape}")
reconstructed = tokenizer.decode(latents)
print(f"復元: {reconstructed.shape}")
5.4 ワールドモデルの実装
import torch.nn as nn
class RobotArmWorldModel(nn.Module):
"""SO-101ロボットアームのワールドモデル"""
def __init__(
self,
temporal_channels, # 時間情報を含むチャネル数
embed_dim=512,
num_heads=8,
num_layers=6,
spatial_size=32,
temporal_tokens=16,
):
"""
Args:
temporal_channels: 潜在表現のチャネル数(時間情報を含む)
embed_dim: Transformerの埋め込み次元
num_heads: マルチヘッドアテンションのヘッド数
num_layers: Transformerの層数
spatial_size: 空間方向のサイズ(32×32)
temporal_tokens: 時間トークンの数(常に16)
"""
super().__init__()
self.temporal_channels = temporal_channels
self.embed_dim = embed_dim
self.spatial_size = spatial_size
self.temporal_tokens = temporal_tokens
# チャネル → 埋め込み変換
self.channel_to_embed = nn.Linear(temporal_channels, embed_dim)
# 位置エンコーディング
max_seq_len = temporal_tokens * spatial_size * spatial_size
self.pos_embedding = nn.Parameter(
torch.randn(1, max_seq_len, embed_dim)
)
# Transformer Encoder
encoder_layer = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=num_heads,
dim_feedforward=embed_dim * 4,
dropout=0.1,
activation='gelu',
batch_first=True
)
self.transformer = nn.TransformerEncoder(
encoder_layer,
num_layers=num_layers
)
# 予測ヘッド
self.predict_head = nn.Linear(embed_dim, temporal_channels)
def forward(self, latents):
"""
Args:
latents: [B, T=16, C_temp, H=32, W=32]
Returns:
predictions: [B, T=16, C_temp, H=32, W=32]
"""
B, T, C, H, W = latents.shape
# [B, T, C, H, W] → [B, T*H*W, C]
latents_flat = rearrange(latents, 'b t c h w -> b (t h w) c')
seq_len = latents_flat.shape[1]
# 埋め込み
x = self.channel_to_embed(latents_flat) # [B, seq_len, D]
# 位置エンコーディング
pos_emb = self.pos_embedding[:, :seq_len, :]
x = x + pos_emb
# Transformer
x = self.transformer(x) # [B, seq_len, D]
# 予測
predictions = self.predict_head(x) # [B, seq_len, C]
# 元の形状に戻す
predictions = rearrange(
predictions,
'b (t h w) c -> b t c h w',
t=T, h=H, w=W
)
return predictions
# 使用例
model = RobotArmWorldModel(
temporal_channels=5, # 32フレーム → 5チャネル
embed_dim=512,
num_heads=8,
num_layers=6,
spatial_size=32,
temporal_tokens=16,
).cuda()
print(f"モデルパラメータ数: {sum(p.numel() for p in model.parameters()):,}")
# テスト
test_latents = tokenizer.encode(test_video)
output = model(test_latents)
print(f"入力: {test_latents.shape}")
print(f"出力: {output.shape}")
5.5 学習関数
import os
from torch.utils.data import DataLoader
from tqdm import tqdm
import torch.nn.functional as F
def train_world_model(
video_dir,
checkpoint_dir='./cosmos_checkpoints',
output_dir='./outputs',
num_epochs=50,
batch_size=2,
learning_rate=1e-4,
sequence_length=32,
device='cuda',
):
"""
ワールドモデルの学習
Args:
video_dir: 動画ディレクトリ
checkpoint_dir: Cosmos Tokenizerのチェックポイント
output_dir: 出力ディレクトリ
num_epochs: エポック数
batch_size: バッチサイズ
learning_rate: 学習率
sequence_length: シーケンス長
device: デバイス
"""
os.makedirs(output_dir, exist_ok=True)
# データセット
print("Loading dataset...")
dataset = RobotArmVideoDataset(
video_dir=video_dir,
sequence_length=sequence_length,
image_size=256
)
dataloader = DataLoader(
dataset,
batch_size=batch_size,
shuffle=True,
num_workers=0,
pin_memory=True
)
# Tokenizer
print("Loading Cosmos Tokenizer...")
tokenizer = CosmosTokenizerWrapper(checkpoint_dir=checkpoint_dir)
# テスト
print("\nTesting tokenizer...")
test_video = dataset[0].unsqueeze(0).to(device)
test_latents = tokenizer.encode(test_video)
print(f"✓ Encoded! Latent shape: {test_latents.shape}")
print(f" temporal_tokens: {test_latents.shape[1]}")
print(f" temporal_channels: {test_latents.shape[2]}")
print(f" spatial: {test_latents.shape[3]}x{test_latents.shape[4]}")
# モデル初期化
print("\nInitializing World Model...")
model = RobotArmWorldModel(
temporal_channels=test_latents.shape[2],
embed_dim=512,
num_heads=8,
num_layers=6,
spatial_size=test_latents.shape[3],
temporal_tokens=test_latents.shape[1],
).to(device)
print(f"Model parameters: {sum(p.numel() for p in model.parameters()):,}")
# Optimizer & Scheduler
optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=num_epochs
)
# 学習ループ
print("\nStarting training...")
best_loss = float('inf')
for epoch in range(num_epochs):
model.train()
epoch_loss = 0.0
num_batches = 0
progress_bar = tqdm(
dataloader,
desc=f"Epoch {epoch+1}/{num_epochs}"
)
for batch_idx, videos in enumerate(progress_bar):
try:
videos = videos.to(device)
# エンコード(勾配計算不要)
with torch.no_grad():
latents = tokenizer.encode(videos)
# 再構成タスク(自己教師あり学習)
input_latents = latents
target_latents = latents
# 予測
predictions = model(input_latents)
# 損失計算
loss = F.mse_loss(predictions, target_latents)
# 最適化
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
epoch_loss += loss.item()
num_batches += 1
progress_bar.set_postfix({'loss': f'{loss.item():.6f}'})
except Exception as e:
print(f"\n❌ Error in batch {batch_idx}: {e}")
continue
if num_batches > 0:
avg_loss = epoch_loss / num_batches
print(f"\nEpoch {epoch+1}/{num_epochs} - Average Loss: {avg_loss:.6f}")
scheduler.step()
# ベストモデル保存
if avg_loss < best_loss:
best_loss = avg_loss
torch.save({
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'loss': avg_loss,
'config': {
'temporal_channels': test_latents.shape[2],
'spatial_size': test_latents.shape[3],
'temporal_tokens': test_latents.shape[1],
}
}, f"{output_dir}/world_model_best.pt")
print(f"✓ Saved best model (loss: {best_loss:.6f})")
print("\n✓ Training completed!")
return model, tokenizer
# 実行
VIDEO_DIR = "/content/drive/MyDrive/SO101_videos"
CHECKPOINT_DIR = "./cosmos_checkpoints"
OUTPUT_DIR = "./outputs"
model, tokenizer = train_world_model(
video_dir=VIDEO_DIR,
checkpoint_dir=CHECKPOINT_DIR,
output_dir=OUTPUT_DIR,
num_epochs=20,
batch_size=2,
learning_rate=1e-4,
sequence_length=32,
device='cuda' if torch.cuda.is_available() else 'cpu',
)
6. 未来予測の実行と長時間予測
6.1 モデルの読み込み
def load_trained_model(checkpoint_path, device='cuda'):
"""学習済みモデルを読み込む"""
checkpoint = torch.load(checkpoint_path, map_location=device)
config = checkpoint['config']
print(f"モデル設定:")
print(f" temporal_channels: {config['temporal_channels']}")
print(f" spatial_size: {config['spatial_size']}")
print(f" temporal_tokens: {config['temporal_tokens']}")
model = RobotArmWorldModel(
temporal_channels=config['temporal_channels'],
embed_dim=512,
num_heads=8,
num_layers=6,
spatial_size=config['spatial_size'],
temporal_tokens=config['temporal_tokens'],
).to(device)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
print(f"\n✓ モデルを読み込みました")
print(f" Epoch: {checkpoint['epoch']}")
print(f" Loss: {checkpoint['loss']:.6f}")
return model, config
# 使用例
device = 'cuda' if torch.cuda.is_available() else 'cpu'
model, config = load_trained_model("./outputs/world_model_best.pt", device)
6.2 基本的な未来予測
@torch.no_grad()
def generate_future_video(
model,
tokenizer,
input_video,
num_future_steps=4,
device='cuda'
):
"""
入力動画から未来を予測
Args:
model: 学習済みワールドモデル
tokenizer: Cosmos Tokenizer
input_video: [1, T, C, H, W] 入力動画
num_future_steps: 未来予測のステップ数
device: デバイス
Returns:
all_frames: 入力+生成された動画
"""
model.eval()
print(f"\n入力動画: {input_video.shape}")
current_video = input_video.to(device)
all_generated = [current_video]
for step in range(num_future_steps):
print(f"ステップ {step+1}/{num_future_steps}")
# エンコード
latents = tokenizer.encode(current_video)
print(f" エンコード: {latents.shape}")
# 予測
predicted_latents = model(latents)
print(f" 予測: {predicted_latents.shape}")
# デコード
predicted_video = tokenizer.decode(predicted_latents)
print(f" デコード: {predicted_video.shape}")
all_generated.append(predicted_video)
current_video = predicted_video
# すべてのフレームを結合
all_frames = torch.cat(all_generated, dim=1)
print(f"\n✓ 生成完了: {all_frames.shape}")
return all_frames
# 実行例
test_video = dataset[0].unsqueeze(0).to(device)
generated = generate_future_video(
model,
tokenizer,
test_video,
num_future_steps=5,
device=device
)
6.3 動画保存関数
import cv2
from PIL import Image
def save_video_as_gif(video_tensor, output_path, fps=10):
"""GIFアニメーションとして保存"""
# [-1, 1] → [0, 255]
video_np = (
(video_tensor[0].cpu().numpy() + 1.0) * 127.5
).clip(0, 255).astype(np.uint8)
# [T, C, H, W] → [T, H, W, C]
video_np = rearrange(video_np, 't c h w -> t h w c')
# PILイメージに変換
images = [Image.fromarray(frame) for frame in video_np]
duration = int(1000 / fps) # ミリ秒
# GIF保存
images[0].save(
output_path,
save_all=True,
append_images=images[1:],
duration=duration,
loop=0
)
print(f"✓ GIF保存: {output_path}")
def save_video_opencv(video_tensor, output_path, fps=10):
"""OpenCVで動画を保存"""
video_np = (
(video_tensor[0].cpu().numpy() + 1.0) * 127.5
).clip(0, 255).astype(np.uint8)
video_np = rearrange(video_np, 't c h w -> t h w c')
height, width = video_np[0].shape[:2]
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (width, height))
for frame in video_np:
frame_bgr = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
out.write(frame_bgr)
out.release()
print(f"✓ 保存: {output_path}")
# 使用例
save_video_as_gif(generated, "./generated_future.gif", fps=10)
save_video_opencv(generated, "./generated_future.mp4", fps=10)
# GIFを表示
from IPython.display import Image as IPImage, display
display(IPImage(filename="./generated_future.gif"))
6.4 長時間予測:スライディングウィンドウ方式
@torch.no_grad()
def generate_long_video_sliding(
model,
tokenizer,
input_video,
total_future_frames=100,
window_size=32,
device='cuda'
):
"""
スライディングウィンドウ方式で長時間予測
Args:
model: 学習済みワールドモデル
tokenizer: Cosmos Tokenizer
input_video: [1, T, C, H, W] 入力動画
total_future_frames: 生成したい総フレーム数
window_size: ウィンドウサイズ(学習時のsequence_lengthと同じ)
device: デバイス
"""
model.eval()
print(f"入力動画: {input_video.shape}")
print(f"目標: {total_future_frames}フレーム生成")
print(f"ウィンドウサイズ: {window_size}\n")
all_frames = [input_video]
current_window = input_video
num_iterations = (total_future_frames + window_size - 1) // window_size
for iteration in range(num_iterations):
print(f"イテレーション {iteration+1}/{num_iterations}")
# エンコード
latents = tokenizer.encode(current_window)
print(f" エンコード: {latents.shape}")
# 予測
predicted_latents = model(latents)
# デコード
predicted_video = tokenizer.decode(predicted_latents)
print(f" デコード: {predicted_video.shape}")
all_frames.append(predicted_video)
# 次のウィンドウを準備(オーバーラップ)
overlap = window_size // 2
current_window = predicted_video[:, -overlap:]
# 不足分を埋める
if current_window.shape[1] < window_size:
padding_frames = window_size - current_window.shape[1]
last_frame = current_window[:, -1:].repeat(
1, padding_frames, 1, 1, 1
)
current_window = torch.cat([current_window, last_frame], dim=1)
print(f" 次のウィンドウ: {current_window.shape}\n")
# すべてのフレームを結合
full_video = torch.cat(all_frames, dim=1)
# 目標フレーム数にトリミング
full_video = full_video[:, :input_video.shape[1] + total_future_frames]
print(f"✓ 生成完了: {full_video.shape}")
return full_video
# 実行例
long_generated = generate_long_video_sliding(
model,
tokenizer,
test_video,
total_future_frames=100,
window_size=32,
device=device
)
save_video_as_gif(long_generated, "./long_prediction.gif", fps=15)
print(f"\n生成されたフレーム数: {long_generated.shape[1]}")
6.5 フレーム比較の可視化
import matplotlib.pyplot as plt
def visualize_prediction(input_video, generated_video, num_frames=8):
"""入力と生成結果を並べて表示"""
fig, axes = plt.subplots(2, num_frames, figsize=(20, 5))
# [-1, 1] → [0, 255]
input_frames = (
(input_video[0].cpu().numpy() + 1.0) * 127.5
).clip(0, 255).astype(np.uint8)
generated_frames = (
(generated_video[0].cpu().numpy() + 1.0) * 127.5
).clip(0, 255).astype(np.uint8)
for i in range(min(num_frames, input_frames.shape[0])):
# 入力フレーム
frame = rearrange(input_frames[i], 'c h w -> h w c')
axes[0, i].imshow(frame)
axes[0, i].axis('off')
if i == 0:
axes[0, i].set_title('Input Frames', fontsize=12, pad=10)
for i in range(min(num_frames, generated_frames.shape[0])):
# 生成フレーム
frame = rearrange(generated_frames[i], 'c h w -> h w c')
axes[1, i].imshow(frame)
axes[1, i].axis('off')
if i == 0:
axes[1, i].set_title('Generated Frames', fontsize=12, pad=10)
plt.tight_layout()
plt.savefig('./prediction_comparison.png', dpi=150, bbox_inches='tight')
plt.show()
print("✓ 比較画像を保存: ./prediction_comparison.png")
# 実行
visualize_prediction(test_video, generated, num_frames=8)
6.6 複数動画でのテスト
import shutil
def test_multiple_videos(model, tokenizer, dataset, num_tests=3, device='cuda'):
"""複数の動画でテスト"""
print("=" * 50)
print("複数動画でテスト")
print("=" * 50)
# Google Drive保存先
drive_output_dir = "/content/drive/MyDrive/SO101_world_model"
os.makedirs(drive_output_dir, exist_ok=True)
for i in range(num_tests):
print(f"\n--- テスト {i+1}/{num_tests} ---")
test_video = dataset[i].unsqueeze(0).to(device)
generated = generate_future_video(
model,
tokenizer,
test_video,
num_future_steps=3,
device=device
)
# ローカルに保存
output_path = f"./test_{i+1}_generated.gif"
save_video_as_gif(generated, output_path, fps=10)
# Google Driveにコピー
shutil.copy2(
output_path,
f"{drive_output_dir}/test_{i+1}_generated.gif"
)
print(f"\n✓ すべてのテスト完了!")
print(f"保存先: {drive_output_dir}")
# 実行
test_multiple_videos(model, tokenizer, dataset, num_tests=3, device=device)
7. まとめと応用可能性
7.1 本記事で実現したこと
✅ Cosmos Tokenizerによる動画の効率的な圧縮
- 空間: 256×256 → 32×32 (8倍圧縮)
- 時間: 32フレーム → 5チャネルに埋め込み
✅ マルチヘッドアテンションによる複数意味の同時理解
- Cosmos: 8〜16ヘッドで視覚特徴を並列抽出
- Transformer: 8ヘッドで時系列パターンを並列学習
✅ Transformerベースのワールドモデルの構築
- 自己教師あり学習(再構成タスク)
- 約23M個のパラメータ
✅ ロボットアームの動きの学習と予測
- 短期予測(5〜10ステップ先)
- 長期予測(100フレーム以上)
✅ スライディングウィンドウによる長時間予測
- オーバーラップを活用した安定予測
7.2 応用可能性
7.2.1 ロボット制御への応用
# シミュレーションベースの行動計画
def plan_action_with_world_model(current_state, goal_state):
"""
ワールドモデルを使った行動計画
"""
# 複数の候補行動を試す
actions = ['move_right', 'move_left', 'grasp', 'release']
best_action = None
best_score = -float('inf')
for action in actions:
# ワールドモデルで結果をシミュレーション
predicted_result = world_model.simulate(current_state, action)
# ゴールに近づいているか評価
score = evaluate_proximity(predicted_result, goal_state)
if score > best_score:
best_score = score
best_action = action
return best_action
# 使用例
current_state = capture_robot_state()
goal_state = define_goal()
next_action = plan_action_with_world_model(current_state, goal_state)
robot.execute(next_action)
7.2.2 異常検知
def detect_anomaly(video_stream, world_model, threshold=0.5):
"""
予測と実際の差から異常を検知
"""
# 過去のフレームから未来を予測
past_frames = video_stream[-32:]
predicted_frame = world_model.predict_next(past_frames)
# 実際のフレームを取得
actual_frame = video_stream.get_current_frame()
# 差分を計算
difference = calculate_difference(predicted_frame, actual_frame)
# 閾値を超えたら異常
if difference > threshold:
alert("Anomaly detected!")
return True
return False
7.2.3 強化学習への統合
# Model-Based Reinforcement Learning
class ModelBasedAgent:
def __init__(self, world_model, policy):
self.world_model = world_model
self.policy = policy
def train(self, environment):
"""
ワールドモデルを使った効率的な学習
"""
# 実環境での経験収集(少量)
real_experiences = environment.collect_episodes(num=10)
# ワールドモデルで大量の仮想経験を生成
simulated_experiences = self.world_model.generate_episodes(num=1000)
# 両方を使ってポリシーを学習
self.policy.train(real_experiences + simulated_experiences)
def act(self, state):
# 複数の行動をシミュレーション
best_action = None
best_reward = -float('inf')
for action in self.policy.sample_actions():
# ワールドモデルで結果を予測
predicted_state = self.world_model.predict(state, action)
predicted_reward = self.estimate_reward(predicted_state)
if predicted_reward > best_reward:
best_reward = predicted_reward
best_action = action
return best_action
7.2.4 データ拡張
def augment_training_data(original_videos, world_model, augmentation_factor=10):
"""
ワールドモデルを使ったデータ拡張
"""
augmented_data = []
for video in original_videos:
# 元の動画を追加
augmented_data.append(video)
# ワールドモデルで変種を生成
for _ in range(augmentation_factor):
# ランダムなフレームから開始
start_frame = random.randint(0, len(video) - 32)
seed_frames = video[start_frame:start_frame+8]
# 続きを生成
generated = world_model.generate_continuation(seed_frames)
augmented_data.append(generated)
return augmented_data
7.3 パフォーマンス改善のヒント
7.3.1 学習の安定化
# より安定した学習のための工夫
# 1. 学習率のウォームアップ
class WarmupScheduler:
def __init__(self, optimizer, warmup_steps, target_lr):
self.optimizer = optimizer
self.warmup_steps = warmup_steps
self.target_lr = target_lr
self.step_count = 0
def step(self):
self.step_count += 1
if self.step_count < self.warmup_steps:
lr = self.target_lr * (self.step_count / self.warmup_steps)
for param_group in self.optimizer.param_groups:
param_group['lr'] = lr
# 2. Gradient Accumulation
accumulation_steps = 4
for i, batch in enumerate(dataloader):
loss = model(batch)
loss = loss / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
# 3. EMA(Exponential Moving Average)
class EMA:
def __init__(self, model, decay=0.999):
self.model = model
self.decay = decay
self.shadow = {}
for name, param in model.named_parameters():
self.shadow[name] = param.data.clone()
def update(self):
for name, param in self.model.named_parameters():
self.shadow[name] = (
self.decay * self.shadow[name] +
(1 - self.decay) * param.data
)
7.3.2 予測精度の向上
# 1. アンサンブル予測
def ensemble_predict(models, input_video):
predictions = []
for model in models:
pred = model.predict(input_video)
predictions.append(pred)
# 平均化
ensemble_pred = torch.stack(predictions).mean(dim=0)
return ensemble_pred
# 2. 不確実性の推定
def predict_with_uncertainty(model, input_video, num_samples=10):
predictions = []
# ドロップアウトを有効にして複数回予測
model.train() # ドロップアウトが有効になる
for _ in range(num_samples):
with torch.no_grad():
pred = model(input_video)
predictions.append(pred)
# 平均と分散
mean_pred = torch.stack(predictions).mean(dim=0)
var_pred = torch.stack(predictions).var(dim=0)
return mean_pred, var_pred
# 3. 階層的予測
def hierarchical_predict(model, input_video):
# 粗い予測
coarse_pred = model.predict_coarse(input_video)
# 細かい予測
fine_pred = model.refine(coarse_pred, input_video)
return fine_pred
7.3.3 計算効率の改善
# 1. Mixed Precision Training
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in dataloader:
with autocast():
loss = model(batch)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
# 2. Gradient Checkpointing
import torch.utils.checkpoint as checkpoint
class EfficientTransformer(nn.Module):
def forward(self, x):
# メモリを節約
x = checkpoint.checkpoint(self.layer1, x)
x = checkpoint.checkpoint(self.layer2, x)
return x
# 3. 量子化
quantized_model = torch.quantization.quantize_dynamic(
model,
{nn.Linear},
dtype=torch.qint8
)
7.4 今後の発展方向
7.4.1 より高度なアーキテクチャ
- Diffusion Models for Video Generation
- VQ-VAE + Transformer
- 3D CNN + Temporal Attention
- Hierarchical Latent Representations
7.4.2 マルチモーダル学習
# 動画 + 言語 + センサーデータ
class MultimodalWorldModel(nn.Module):
def __init__(self):
super().__init__()
# 各モダリティのエンコーダ
self.video_encoder = CosmosTokenizer()
self.text_encoder = TextEncoder()
self.sensor_encoder = SensorEncoder()
# 統合Transformer
self.fusion_transformer = Transformer()
def forward(self, video, text, sensors):
# 各モダリティを潜在表現に
v_latent = self.video_encoder(video)
t_latent = self.text_encoder(text)
s_latent = self.sensor_encoder(sensors)
# 統合
combined = torch.cat([v_latent, t_latent, s_latent], dim=1)
# 予測
prediction = self.fusion_transformer(combined)
return prediction
7.4.3 実時間処理
# オンライン学習と予測
class OnlineWorldModel:
def __init__(self, model):
self.model = model
self.buffer = []
def update_online(self, new_observation):
"""新しい観測でモデルを更新"""
self.buffer.append(new_observation)
if len(self.buffer) >= 32:
# バッファが溜まったら学習
batch = torch.stack(self.buffer[-32:])
self.model.update(batch)
def predict_realtime(self, current_state):
"""リアルタイム予測"""
with torch.no_grad():
prediction = self.model(current_state)
return prediction
7.5 参考リンク
- NVIDIA Cosmos Tokenizer GitHub
- World Models論文 (Ha & Schmidhuber, 2018)
- Attention Is All You Need (Vaswani et al., 2017)
- Vision Transformer (Dosovitskiy et al., 2020)
- Sora: Creating Video from Text (OpenAI)
7.6 最後に
本記事では、Cosmos Tokenizerを使ったワールドモデルの構築方法を、理論から実装まで詳しく解説しました。
3つの核心ポイント:
- 賢い圧縮: Cosmos Tokenizerが「見た目」を「意味」に変換
- パターン学習: Transformerが「次はこうなる」を学習
- マルチヘッド: 複数の意味を同時に理解
これらの技術を組み合わせることで、ロボットアームの動きを予測できるAIを構築できました。
この技術は、ロボティクスだけでなく、自動運転、産業オートメーション、ゲームAI、医療画像解析など、様々な分野に応用可能です。