DINO-WM (PreJEPA) β€” PushT Β· patch tokens + proprio

ν”„λ ˆμž„μ„ μ–Όλ¦° DINOv2-small 백본으둜 latent μΈμ½”λ”©ν•˜κ³ , causal predictor 둜 λ‹€μŒ latent 을 μ˜ˆμΈ‘ν•˜λŠ” world model (DINO-WM 계열, JEPA loss). ν”½μ…€ μž¬κ΅¬μ„± μ—†μŒ.

  • backbone: dinov2_small (frozen), pixel_token=patch β†’ ν”„λ ˆμž„λ‹Ή 256 패치 Γ— 384-d
  • predictor: CausalPredictor, dim=404 (= pixel 384 + proprio_emb 10 + action_emb 10)
  • history_size=3, num_pred=1, frameskip=5
  • λΆ€κ°€μž…λ ₯: proprio(in_chans=4: agent pos+vel), action(in_chans=10 = raw 2 Γ— frameskip 5)
  • env: swm/PushT-v1

파일

파일 μ„€λͺ…
weights.pt λͺ¨λΈ κ°€μ€‘μΉ˜ (epoch 10)
config.json ꡬ쑰 (hydra instantiate 용)
norm_stats.json proprio/action ZScore meanΒ·std (eval μ •κ·œν™” 볡원)

μ„€μΉ˜

pip install stable-worldmodel   # λ˜λŠ” μ €μž₯μ†Œμ—μ„œ editable μ„€μΉ˜

λ‘œλ“œ (public repo β†’ λ‚΄μž₯ λ‘œλ”)

import stable_worldmodel as swm

model = swm.wm.utils.load_pretrained("kotmul/dinowm_patch_prop_pusht")
model = model.eval().requires_grad_(False)
model.interpolate_pos_encoding = True

load_pretrained λŠ” config.json + weights.pt λ₯Ό <cache>/checkpoints/ μ•„λž˜λ‘œ λ°›μ•„ instantiate(config) ν›„ κ°€μ€‘μΉ˜λ₯Ό λ‘œλ“œν•œλ‹€.

μ •κ·œν™” (μ€‘μš”)

  • pixels: ImageNet mean/std μ •κ·œν™” ν›„ 224Γ—224
  • proprio / action: μ•„λž˜ norm_stats.json 의 ZScore (ν•™μŠ΅κ³Ό λ°˜λ“œμ‹œ 동일해야 함)
import json, numpy as np
from huggingface_hub import hf_hub_download

norm = json.load(open(hf_hub_download("kotmul/dinowm_patch_prop_pusht", "norm_stats.json")))
p_mean, p_std = np.array(norm["proprio"]["mean"][0]), np.array(norm["proprio"]["std"][0])
a_mean, a_std = np.array(norm["action"]["mean"][0]),  np.array(norm["action"]["std"][0])

μΆ”λ‘  β€” ν”„λ ˆμž„ 인코딩 & λ‹€μŒ μŠ€ν… 예츑

import torch, numpy as np
import stable_pretraining as spt
from torchvision.transforms import v2 as T

tf = T.Compose([
    T.ToImage(), T.ToDtype(torch.float32, scale=True),
    T.Normalize(**spt.data.dataset_stats.ImageNet), T.Resize(224),
])
H, FS = model.history_size, 5   # 3 history steps, frameskip 5

# frames_uint8: (H, 224, 224, 3) uint8  β€” history_size 개의 연속 ν”„λ ˆμž„(frameskip 간격)
# proprio_raw : (H, 4)   각 μŠ€ν…μ˜ [agent_x, agent_y, agent_vx, agent_vy]
# action_raw  : (H, FS*2) 각 model-step 의 raw action FS개 묢음 ([-1,1]^2 Γ— FS)

pixels  = torch.stack([tf(im) for im in frames_uint8])[None]                 # (1,H,3,224,224)
proprio = torch.tensor(((proprio_raw - p_mean) / p_std)[None], dtype=torch.float32)  # (1,H,4)
action  = ((action_raw.reshape(H, FS, 2) - a_mean) / a_std).reshape(H, FS * 2)
action  = torch.tensor(action[None], dtype=torch.float32)                    # (1,H,10)

with torch.no_grad():
    # (a) 단일 ν”„λ ˆμž„ 인코딩 (patch latent)
    emb_img = model._encode_image(pixels[:, :1])   # (1, 1, 256, 384)

    # (b) λ‹€μŒ μŠ€ν… 예츑 (action/proprio 반영)
    info = {"pixels": pixels, "proprio": proprio, "action": action}
    info = model.encode(info, target="emb", is_video=False)
    pred = model.predict(info["emb"][:, :H])       # (1, H, 256, 404)
    next_latent = pred[:, -1]                       # μ˜ˆμΈ‘ν•œ λ‹€μŒ latent (1, 256, 404)
    # 404 = pixel(384) + proprio_emb(10) + action_emb(10).
    # planning cost λ“±μ—λŠ” 보톡 action ꡬ간(λ§ˆμ§€λ§‰ 10)을 μ œμ™Έν•œ actionless λΆ€λΆ„ μ‚¬μš©:
    actionless = next_latent[..., :394]

Planning / eval μ—μ„œ μ“°κΈ°

stable-worldmodel 의 planning(eval) μŠ€ν¬λ¦½νŠΈλŠ” 체크포인트 μ˜†μ˜ norm_stats.json 을 μžλ™μœΌλ‘œ μ°Ύμ•„ ν•™μŠ΅ λ•Œ μ •κ·œν™”λ₯Ό λ³΅μ›ν•œλ‹€(option B). λ”°λΌμ„œ μ„Έ νŒŒμΌμ„ <cache>/checkpoints/<run>/ ν•œ 폴더에 두고 policy λ₯Ό κ·Έ weights.pt 둜 μ§€μ •ν•˜λ©΄ λœλ‹€:

<cache>/checkpoints/dinowm-pusht-patch-prop/
    weights.pt
    config.json
    norm_stats.json      # eval 이 μ—¬κΈ°μ„œ mean/std 볡원

MPC(CEM) planning 은 world model 을 imagination 으둜 κ΅΄λ € cost λ₯Ό μ΅œμ†Œν™”ν•˜κ³ , μ‹€μ œ env μ—μ„œ μ‹€ν–‰ν•œλ‹€. μžμ„Έν•œ μ§„μž…μ μ€ λ¦¬ν¬μ§€ν† λ¦¬μ˜ planning 슀크립트λ₯Ό μ°Έκ³ .

Downloads last month
36
Video Preview
loading