seungkukim/dexgys_grouped_wan22ti2v5b_L18_g6_pc48_txtcache512_joint-60k
DexGYS grouped-action (arm A2) @ step 60,000
DiT4DiT joint checkpoint — a Wan2.2-TI2V-5B video backbone plus a DiT-B action head that predicts the 63-D Shadow-Hand keypoint grasp as 6 anatomically grouped action tokens, conditioned on 48 spatial PartField point-cloud tokens. Trained on `seungkukim/dexgys_video224_a0_lerobot` (37,158 grasps over 1,258 objects) with cached umT5 text embeddings and feature extraction at block 18.
Full training run name: dexgys_grouped_ditjoint_wan22_L18_g15-9-9-9-12-9_pc48_txtcache512_pd16_ga1_effgb64_local
What is in here
Inference files only — 12 files, 11.84 GB of weights plus a 124 MB PartField bank. The 1,318 tensors in the shards are:
Training-only artifacts are excluded: the DeepSpeed global_step60000/ optimizer state (~70 GB), rng_state_*.pth, scheduler.pt, trainer_state.json, training_args.bin, zero_to_fp32.py, latest, wandb_config.json.
The six anatomical action groups
action_group_widths = [15, 9, 9, 9, 12, 9] (sum = 63 = max_action_dim). Order is load-bearing — the head slices the 63-D target by these widths, so a permutation reorders the target with no error:
The PartField bank ships WITH this repo
DexGYS is stateless, so the state slot carries the object: the state group is one scalar join key (point_cloud_index) and the processor swaps it for 48 × 1024 PartField triplane tokens read from a frozen memmapped bank. That bank is data, not weights — it is absent from the shards, and the processor cannot build the state stream without it. Included here as:
ckpt/dexgys_pc_features/dexgys_partfield_k8.f16
ckpt/dexgys_pc_features/dexgys_partfield_k8.f16.jsonSidecar claims, all asserted at load: n_scenes=1258, num_tokens=48, feat_dim=1024, pool_kernel=8, split='train', frame = 'canonical object frame (render views=1, azimuth 0)'.
The split field matters: the eval bank RENUMBERS point_cloud_index densely over the 309 test scenes, so every test index is also a valid train index — a bank from the wrong split returns another object's features with no error anywhere.
External dependencies
Currently wan_model_path = /data/seungku/hf_cache/Wan2.2-TI2V-5B-Diffusers, pc_feature_path = /data/seungku/projects/wam/ckpt/dexgys_pc_features/dexgys_partfield_k8.f16.
config.json and processor_config.json still carry the local training paths (wan_model_path, pc_feature_path) and wan_local_files_only=true. Override both before loading this off-box, or cd into the snapshot and set pc_feature_path=ckpt/dexgys_pc_features/dexgys_partfield_k8.f16.
The base Wan2.2 snapshot is required either way — from_pretrained() builds the backbone from it and overlays these shards:
huggingface-cli download Wan-AI/Wan2.2-TI2V-5B-Diffusers --local-dir /path/to/Wan2.2-TI2V-5B-DiffusersText conditioning is not part of the checkpoint. This run trained against a precomputed umT5 cache (seqlen 512, `formalizelanguage=False) selected by the WAMTEXTEMBEDCACHE` environment variable, which `config.json` does not record. Either bake that cache (`scripts/wamdit4dit/precomputetextembedswan22.py`, ~87 GB at 512) or let umT5 encode live from the base snapshot. `formalizelanguage` must match what the cache was baked with.
Load
from gr00t.model.wam_dit4dit import WAMDiT4DiT
model = WAMDiT4DiT.from_pretrained("seungkukim/dexgys_grouped_wan22ti2v5b_L18_g6_pc48_txtcache512_joint-60k", torch_dtype="bfloat16")