Upload README.md with huggingface_hub
Browse files
README.md
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: mit
|
| 3 |
+
tags:
|
| 4 |
+
- world-model
|
| 5 |
+
- jepa
|
| 6 |
+
- dino-wm
|
| 7 |
+
- robotics
|
| 8 |
+
- pusht
|
| 9 |
+
library_name: stable-worldmodel
|
| 10 |
+
---
|
| 11 |
+
|
| 12 |
+
# DINO-WM (PreJEPA) β PushT Β· patch tokens + proprio
|
| 13 |
+
|
| 14 |
+
νλ μμ μΌλ¦° **DINOv2-small** λ°±λ³ΈμΌλ‘ latent μΈμ½λ©νκ³ , causal predictor λ‘
|
| 15 |
+
λ€μ latent μ μμΈ‘νλ world model (DINO-WM κ³μ΄, JEPA loss). ν½μ
μ¬κ΅¬μ± μμ.
|
| 16 |
+
|
| 17 |
+
- backbone: `dinov2_small` (frozen), `pixel_token=patch` β νλ μλΉ 256 ν¨μΉ Γ 384-d
|
| 18 |
+
- predictor: `CausalPredictor`, `dim=404` (= pixel 384 + proprio_emb 10 + action_emb 10)
|
| 19 |
+
- `history_size=3`, `num_pred=1`, `frameskip=5`
|
| 20 |
+
- λΆκ°μ
λ ₯: `proprio`(in_chans=4: agent pos+vel), `action`(in_chans=10 = raw 2 Γ frameskip 5)
|
| 21 |
+
- env: `swm/PushT-v1`
|
| 22 |
+
|
| 23 |
+
## νμΌ
|
| 24 |
+
|
| 25 |
+
| νμΌ | μ€λͺ
|
|
| 26 |
+
|---|---|
|
| 27 |
+
| `weights.pt` | λͺ¨λΈ κ°μ€μΉ (epoch 10) |
|
| 28 |
+
| `config.json` | ꡬ쑰 (hydra instantiate μ©) |
|
| 29 |
+
| `norm_stats.json` | proprio/action ZScore meanΒ·std (eval μ κ·ν 볡μ) |
|
| 30 |
+
|
| 31 |
+
## μ€μΉ
|
| 32 |
+
|
| 33 |
+
```bash
|
| 34 |
+
pip install stable-worldmodel # λλ μ μ₯μμμ editable μ€μΉ
|
| 35 |
+
```
|
| 36 |
+
|
| 37 |
+
## λ‘λ (public repo β λ΄μ₯ λ‘λ)
|
| 38 |
+
|
| 39 |
+
```python
|
| 40 |
+
import stable_worldmodel as swm
|
| 41 |
+
|
| 42 |
+
model = swm.wm.utils.load_pretrained("kotmul/dinowm_patch_prop_pusht")
|
| 43 |
+
model = model.eval().requires_grad_(False)
|
| 44 |
+
model.interpolate_pos_encoding = True
|
| 45 |
+
```
|
| 46 |
+
|
| 47 |
+
`load_pretrained` λ `config.json` + `weights.pt` λ₯Ό `<cache>/checkpoints/` μλλ‘ λ°μ
|
| 48 |
+
`instantiate(config)` ν κ°μ€μΉλ₯Ό λ‘λνλ€.
|
| 49 |
+
|
| 50 |
+
## μ κ·ν (μ€μ)
|
| 51 |
+
|
| 52 |
+
- **pixels**: ImageNet mean/std μ κ·ν ν 224Γ224
|
| 53 |
+
- **proprio / action**: μλ `norm_stats.json` μ ZScore (νμ΅κ³Ό λ°λμ λμΌν΄μΌ ν¨)
|
| 54 |
+
|
| 55 |
+
```python
|
| 56 |
+
import json, numpy as np
|
| 57 |
+
from huggingface_hub import hf_hub_download
|
| 58 |
+
|
| 59 |
+
norm = json.load(open(hf_hub_download("kotmul/dinowm_patch_prop_pusht", "norm_stats.json")))
|
| 60 |
+
p_mean, p_std = np.array(norm["proprio"]["mean"][0]), np.array(norm["proprio"]["std"][0])
|
| 61 |
+
a_mean, a_std = np.array(norm["action"]["mean"][0]), np.array(norm["action"]["std"][0])
|
| 62 |
+
```
|
| 63 |
+
|
| 64 |
+
## μΆλ‘ β νλ μ μΈμ½λ© & λ€μ μ€ν
μμΈ‘
|
| 65 |
+
|
| 66 |
+
```python
|
| 67 |
+
import torch, numpy as np
|
| 68 |
+
import stable_pretraining as spt
|
| 69 |
+
from torchvision.transforms import v2 as T
|
| 70 |
+
|
| 71 |
+
tf = T.Compose([
|
| 72 |
+
T.ToImage(), T.ToDtype(torch.float32, scale=True),
|
| 73 |
+
T.Normalize(**spt.data.dataset_stats.ImageNet), T.Resize(224),
|
| 74 |
+
])
|
| 75 |
+
H, FS = model.history_size, 5 # 3 history steps, frameskip 5
|
| 76 |
+
|
| 77 |
+
# frames_uint8: (H, 224, 224, 3) uint8 β history_size κ°μ μ°μ νλ μ(frameskip κ°κ²©)
|
| 78 |
+
# proprio_raw : (H, 4) κ° μ€ν
μ [agent_x, agent_y, agent_vx, agent_vy]
|
| 79 |
+
# action_raw : (H, FS*2) κ° model-step μ raw action FSκ° λ¬Άμ ([-1,1]^2 Γ FS)
|
| 80 |
+
|
| 81 |
+
pixels = torch.stack([tf(im) for im in frames_uint8])[None] # (1,H,3,224,224)
|
| 82 |
+
proprio = torch.tensor(((proprio_raw - p_mean) / p_std)[None], dtype=torch.float32) # (1,H,4)
|
| 83 |
+
action = ((action_raw.reshape(H, FS, 2) - a_mean) / a_std).reshape(H, FS * 2)
|
| 84 |
+
action = torch.tensor(action[None], dtype=torch.float32) # (1,H,10)
|
| 85 |
+
|
| 86 |
+
with torch.no_grad():
|
| 87 |
+
# (a) λ¨μΌ νλ μ μΈμ½λ© (patch latent)
|
| 88 |
+
emb_img = model._encode_image(pixels[:, :1]) # (1, 1, 256, 384)
|
| 89 |
+
|
| 90 |
+
# (b) λ€μ μ€ν
μμΈ‘ (action/proprio λ°μ)
|
| 91 |
+
info = {"pixels": pixels, "proprio": proprio, "action": action}
|
| 92 |
+
info = model.encode(info, target="emb", is_video=False)
|
| 93 |
+
pred = model.predict(info["emb"][:, :H]) # (1, H, 256, 404)
|
| 94 |
+
next_latent = pred[:, -1] # μμΈ‘ν λ€μ latent (1, 256, 404)
|
| 95 |
+
# 404 = pixel(384) + proprio_emb(10) + action_emb(10).
|
| 96 |
+
# planning cost λ±μλ λ³΄ν΅ action ꡬκ°(λ§μ§λ§ 10)μ μ μΈν actionless λΆλΆ μ¬μ©:
|
| 97 |
+
actionless = next_latent[..., :394]
|
| 98 |
+
```
|
| 99 |
+
|
| 100 |
+
## Planning / eval μμ μ°κΈ°
|
| 101 |
+
|
| 102 |
+
`stable-worldmodel` μ planning(eval) μ€ν¬λ¦½νΈλ 체ν¬ν¬μΈνΈ μμ `norm_stats.json` μ
|
| 103 |
+
μλμΌλ‘ μ°Ύμ νμ΅ λ μ κ·νλ₯Ό 볡μνλ€(option B). λ°λΌμ μΈ νμΌμ
|
| 104 |
+
`<cache>/checkpoints/<run>/` ν ν΄λμ λκ³ `policy` λ₯Ό κ·Έ `weights.pt` λ‘ μ§μ νλ©΄ λλ€:
|
| 105 |
+
|
| 106 |
+
```
|
| 107 |
+
<cache>/checkpoints/dinowm-pusht-patch-prop/
|
| 108 |
+
weights.pt
|
| 109 |
+
config.json
|
| 110 |
+
norm_stats.json # eval μ΄ μ¬κΈ°μ mean/std 볡μ
|
| 111 |
+
```
|
| 112 |
+
|
| 113 |
+
MPC(CEM) planning μ world model μ imagination μΌλ‘ κ΅΄λ € cost λ₯Ό μ΅μννκ³ ,
|
| 114 |
+
μ€μ env μμ μ€ννλ€. μμΈν μ§μ
μ μ 리ν¬μ§ν 리μ planning μ€ν¬λ¦½νΈλ₯Ό μ°Έκ³ .
|