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:

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:

  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:

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.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:

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:

  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:
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.