a2c2 / modification.md
dennis96's picture
Upload modification.md with huggingface_hub
5bac388 verified
|
Raw
History Blame Contribute Delete
5.34 kB
# OpenPI-COMET Backbone Modifications
This file records local changes made to `b1k/openpi-comet/` for integrating the
trained A2C2 correction head into BEHAVIOR online evaluation.
## 2026-05-29: expose COMET prefix latent from action sampling
### Files changed
- `b1k/openpi-comet/src/openpi/models/pi0.py`
- `b1k/openpi-comet/src/openpi/policies/policy.py`
### `Pi0.sample_actions`
Added an optional keyword argument:
```python
return_prefix_z: bool = False
```
Default behavior is unchanged. With `return_prefix_z=False`, the function still
returns only the sampled action chunk:
```python
actions
```
With `return_prefix_z=True`, the function returns:
```python
actions, prefix_z
```
`prefix_z` is computed from the same prefix forward pass already used to build
the KV cache for action generation:
1. Run `embed_prefix(observation)`.
2. Run `PaliGemma.llm([prefix_tokens, None], ...)`.
3. Keep `prefix_out`.
4. Apply `prefix_mask` weighted mean pooling.
This matches the latent definition used by `b1k/a2c2_create_dataset.py`:
```text
mask-pooled prefix_out from COMET/OpenPI PI0.5 PaliGemma prefix forward
over image, prompt, and discrete state tokens
```
Expected `prefix_z` shape for the current PI0.5 COMET checkpoint is:
```text
[batch, 2048]
```
### `Policy` JAX jit setup
Updated `openpi.policies.policy.Policy.__init__` so JAX policies detect whether
`model.sample_actions` supports `return_prefix_z`. If it does, the jitted wrapper
is created with:
```python
static_argnames=("return_prefix_z",)
```
This keeps the Python branch static under JAX JIT and allows future A2C2 wrapper
code to call `sample_actions(..., return_prefix_z=True)` safely.
### `Policy.infer_with_prefix_z`
Added a new method:
```python
Policy.infer_with_prefix_z(obs, *, noise=None) -> dict
```
This method follows the same preprocessing, input transforms, sampling, output
transforms, and timing pattern as `Policy.infer()`, but calls:
```python
self._sample_actions(..., return_prefix_z=True)
```
and unpacks:
```python
actions, prefix_z
```
The returned dictionary contains the normal transformed policy output plus the
untransformed COMET prefix latent:
```python
{
"actions": ..., # normal output-transformed action chunk
"prefix_z": ..., # float32-like numpy array, shape [2048]
"policy_timing": ...,
}
```
`prefix_z` is appended after output transforms so task-specific transforms such
as `B1kOutputs` cannot drop or reshape it.
This method currently supports the JAX / Orbax OpenPI-COMET path only. It raises
`NotImplementedError` for PyTorch policies or models whose `sample_actions`
does not expose `return_prefix_z`.
### Current behavior
`Policy.infer()` behavior was not changed. Existing OpenPI-COMET baseline
serving still returns the same output as before:
```python
{
"state": ...,
"actions": ...,
"policy_timing": ...,
}
```
The new latent return path is available through `Policy.infer_with_prefix_z()`
for a future A2C2-specific wrapper.
## 2026-05-30: add A2C2 online B1K wrapper
### Files changed
- `b1k/openpi-comet/src/openpi/shared/a2c2_b1k_wrapper.py`
- `b1k/openpi-comet/scripts/serve_b1k_a2c2.py`
- `b1k/openpi-comet/src/a2c2/model.py`
- `b1k/openpi-comet/src/a2c2/dataset.py`
### `A2C2B1KPolicyWrapper`
Added a dedicated wrapper that keeps the original `B1KPolicyWrapper` baseline
untouched. The first version supports only:
```text
control_mode = "receeding_horizon"
```
The wrapper now imports the A2C2 model directly from the bundled OpenPI-COMET
source tree:
```python
from a2c2.model import A2C2CorrectionHead, A2C2CorrectionHeadConfig
```
It no longer inserts the external `b1k/a2c2/src` directory into `sys.path`.
`a2c2_root` is kept as a deprecated CLI compatibility argument but is ignored.
This is intentional because the A2C2 training tuples are:
```text
o_{t+k}, base_chunk_t[k], base_chunk_t, z_t, k -> delta_{t+k}
```
So online execution should not correct the full future chunk at replan time.
Instead, the wrapper:
1. Replans with `policy.infer_with_prefix_z(batch)` when its queue is empty.
2. Stores base chunk context plus `prefix_z` in the queue.
3. At every environment step, uses the current 256-d BEHAVIOR proprio state to
predict the residual for the queued base action.
4. Executes:
```python
final_action = base_action + residual_scale * a2c2_delta
```
Optional clipping is available through `delta_clip`.
The queued context for each action is:
```python
{
"base_action": base_chunk[k],
"base_action_chunk": base_chunk,
"base_policy_z": prefix_z,
"valid_action_mask": valid_mask,
"chunk_index": k,
}
```
### `serve_b1k_a2c2.py`
Added a server entrypoint mirroring `scripts/serve_b1k.py`, but wrapping the
OpenPI-COMET policy with `A2C2B1KPolicyWrapper`.
Example:
```bash
uv run --no-sync scripts/serve_b1k_a2c2.py \
--task_name=tidying_bedroom \
--control_mode=receeding_horizon \
--max_len=32 \
--port=8001 \
--a2c2_checkpoint=/a2c2/results/latest.pt \
--a2c2_device=cuda \
policy:checkpoint \
--policy.config=pi05_b1k-base \
--policy.dir=/20TB_02/dennis_openpi/b1k/openpi-comet/checkpoints/pi05-b1kpt12-cs32
```
This server still speaks the same BEHAVIOR websocket protocol as the baseline
server. The BEHAVIOR evaluation command can keep using `policy=websocket`.