kotmul commited on
Commit
7c424de
Β·
verified Β·
1 Parent(s): 053cb1e

Upload README.md with huggingface_hub

Browse files
Files changed (1) hide show
  1. README.md +114 -0
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 슀크립트λ₯Ό μ°Έκ³ .