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.pyb1k/openpi-comet/src/openpi/policies/policy.py
Pi0.sample_actions
Added an optional keyword argument:
return_prefix_z: bool = False
Default behavior is unchanged. With return_prefix_z=False, the function still
returns only the sampled action chunk:
actions
With return_prefix_z=True, the function returns:
actions, prefix_z
prefix_z is computed from the same prefix forward pass already used to build
the KV cache for action generation:
- Run
embed_prefix(observation). - Run
PaliGemma.llm([prefix_tokens, None], ...). - Keep
prefix_out. - Apply
prefix_maskweighted mean pooling.
This matches the latent definition used by b1k/a2c2_create_dataset.py:
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:
[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:
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:
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:
self._sample_actions(..., return_prefix_z=True)
and unpacks:
actions, prefix_z
The returned dictionary contains the normal transformed policy output plus the untransformed COMET prefix latent:
{
"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:
{
"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.pyb1k/openpi-comet/scripts/serve_b1k_a2c2.pyb1k/openpi-comet/src/a2c2/model.pyb1k/openpi-comet/src/a2c2/dataset.py
A2C2B1KPolicyWrapper
Added a dedicated wrapper that keeps the original B1KPolicyWrapper baseline
untouched. The first version supports only:
control_mode = "receeding_horizon"
The wrapper now imports the A2C2 model directly from the bundled OpenPI-COMET source tree:
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:
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:
- Replans with
policy.infer_with_prefix_z(batch)when its queue is empty. - Stores base chunk context plus
prefix_zin the queue. - At every environment step, uses the current 256-d BEHAVIOR proprio state to predict the residual for the queued base action.
- Executes:
final_action = base_action + residual_scale * a2c2_delta
Optional clipping is available through delta_clip.
The queued context for each action is:
{
"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:
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.