--- license: mit tags: - world-model - jepa - dino-wm - robotics - pusht library_name: stable-worldmodel --- # 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 정규화 복원) | ## 설치 ```bash pip install stable-worldmodel # 또는 저장소에서 editable 설치 ``` ## 로드 (public repo → 내장 로더) ```python 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` 를 `/checkpoints/` 아래로 받아 `instantiate(config)` 후 가중치를 로드한다. ## 정규화 (중요) - **pixels**: ImageNet mean/std 정규화 후 224×224 - **proprio / action**: 아래 `norm_stats.json` 의 ZScore (학습과 반드시 동일해야 함) ```python 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]) ``` ## 추론 — 프레임 인코딩 & 다음 스텝 예측 ```python 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). 따라서 세 파일을 `/checkpoints//` 한 폴더에 두고 `policy` 를 그 `weights.pt` 로 지정하면 된다: ``` /checkpoints/dinowm-pusht-patch-prop/ weights.pt config.json norm_stats.json # eval 이 여기서 mean/std 복원 ``` MPC(CEM) planning 은 world model 을 imagination 으로 굴려 cost 를 최소화하고, 실제 env 에서 실행한다. 자세한 진입점은 리포지토리의 planning 스크립트를 참고.