dp3-floatgrip / README.md
albilo's picture
Upload README.md with huggingface_hub
8b760fb verified
|
Raw History Blame Contribute Delete
5.51 kB
---
license: mit
tags:
- robotics
- imitation-learning
- diffusion-policy
- cloth-manipulation
---
# DP3 floating-gripper garment policies β€” Hang_Trousers
Five [3D Diffusion Policy](https://3d-diffusion-policy.github.io/) variants trained on
floating-gripper garment-hanging demonstrations in
[DexGarmentLab](https://huggingface.co/datasets/wayrise/DexGarmentLab)
(`Hang_Trousers_FloatGrip`, 195 demos, 12,516 transitions, Isaac Sim).
They differ **only** in the observation representation and the point encoder, which is
what makes them useful as an ablation: what does ground-truth cloth mesh state, and
true mesh connectivity, buy over a depth point cloud?
EMA weights only, ~0.98 GiB each. The source repo is currently private; everything
needed to *use* these checkpoints is below.
## The five variants
`EE MAE` is held-out end-effector position error against demonstration actions.
**It is not task success** β€” no policy here has been evaluated in closed loop yet.
| file | encoder | observation | edges | epoch | EE MAE |
|---|---|---|---|---|---|
| `trousers_pn_depth.ckpt` | PointNet | depth cloud 2048Γ—3 | β€” | 550 | **6.25 mm** |
| `trousers_gnn_depth.ckpt` | mesh-GNN | depth cloud | kNN | 200 | 6.85 mm |
| `trousers_pn_mesh.ckpt` | PointNet | decimated GT mesh + pothook | β€” | 900 | 6.94 mm |
| `trousers_gnn_mesh.ckpt` | mesh-GNN | GT mesh nodes | **true garment faces** | 300 | 7.24 mm |
| `trousers_gnn_knn.ckpt` | mesh-GNN | GT mesh nodes | kNN | 200 | 7.35 mm |
### How to read this table
- **Down a column** (same encoder) isolates the representation. **Across a row** isolates
the encoder. Never use a row to make a claim about representations.
- `gnn_mesh` vs `gnn_knn` is the sharpest comparison: identical nodes, identical encoder,
identical parameter count β€” the *only* difference is whether edges come from true mesh
faces or from k-nearest-neighbours. The gap is 0.11 mm, i.e. nothing.
- Checkpoints were selected by lowest EE MAE, **not** by validation loss. The two disagree
for every run, and selecting on `val_loss` would be up to 3.3 mm worse β€” it is diffusion
denoising MSE averaged over noise levels, not a measure of sampled action quality.
### Honest caveats
- All five are within ~1 mm of each other on a **4-episode** validation split. Treat them
as **near-indistinguishable**.
- Each figure is an argmin over 6–20 checkpoints on that same small split, so the numbers
are optimistic.
- They *are* learning: an observation-free "hold current position" baseline scores
48.9 mm and constant-velocity 28.3 mm, so the policies beat trivial baselines by 4.5–8Γ—.
- Encoder capacity, normalizer range and padding were each checked and ruled out as
explanations for the mesh variants underperforming. The likely reason is that the depth
cloud already contains ~75 % garment points plus the randomised target hook in one
view-consistent frame, and this task's difficulty is not cloth-state estimation.
## Usage
```bash
hf download albilo/dp3-floatgrip --local-dir checkpoints_pretrained
```
Needs `torch`, `diffusers`, `hydra-core`, `omegaconf`, `dill`, `einops`, `termcolor`, and
the `diffusion_policy_3d` package from DexGarmentLab's `IL_Baselines/Diffusion_Policy_3D`
on `PYTHONPATH` (the config in each checkpoint instantiates classes by dotted path).
```python
import torch, dill, hydra
from omegaconf import OmegaConf
OmegaConf.register_new_resolver("eval", eval, replace=True)
p = torch.load("trousers_pn_depth.ckpt", map_location="cpu", pickle_module=dill)
policy = hydra.utils.instantiate(p["cfg"].policy)
policy.load_state_dict(p["state_dicts"]["ema_model"])
policy.cuda().eval()
action = policy.predict_action(obs)["action"] # (B, 4, 8)
```
**Observation** β€” `n_obs_steps = 3`, so a history of 3 frames, oldest first:
| key | shape | notes |
|---|---|---|
| `point_cloud` | `(B, 3, 2048, 3)` float32 | xyz only, metres, ground removed, FPS-downsampled |
| `agent_pos` | `(B, 3, 8)` float32 | `[left_xyz, right_xyz, grip_L, grip_R]` |
`trousers_gnn_mesh` additionally requires `edges` `(B, 3, E, 2)`, `n_nodes` `(B, 3)` and
`n_edges` `(B, 3)`; without them it raises `KeyError: 'edges'`. The other four do not β€”
the kNN variants build their graph from coordinates at runtime.
**Action** β€” `(B, 4, 8)`: four future steps of
`[left_xyz, right_xyz, grip_L, grip_R]`, **absolute positions in metres**, not deltas.
There is no orientation (the grippers are attachment spheres; the demos never recorded
any). The normalizer travels inside the `state_dict`, so actions come back already in
metres β€” do not rescale. Gripper bits are regressed continuously; threshold at 0.5.
Execute all 4 actions before re-planning.
**Timing** β€” demonstrations were recorded every 10 physics steps with gripper motion
interpolated at 0.01 m/step, so one action β‰ˆ one recorded frame. Matching this matters:
a different velocity regime is out of distribution.
## Training
195 successful demonstrations, 12,516 transitions, one H100 per run. PointNet cells
1000 epochs @ batch 256 (~1.4 h); mesh-GNN cells 300 epochs @ batch 64 (~4.4 h, gradient
checkpointing β€” the dense padded edge activations peak at 77 GiB without it). The
mesh-GNN is a MeshGraphNets-style encoder β†’ processor over the garment graph; mesh nodes
are decimated to 1234–1535 with topology taken from the garment USDs and vertex
correspondence preserved per frame.
`manifest.json` records which training epoch each file came from.