Skip to content

vllm_omni.diffusion.models.pi0.modeling_pi0

Inference-only π0 (Pi-Zero) VLA math kernel for vllm-omni.

A self-contained inference kernel: only the math that turns a robot observation into an action chunk, with no serving/request glue — validated bit-for-bit against the LeRobot PI0Policy reference (max|Δ| < 1e-4).

π0 = PaliGemma (SigLIP vision + Gemma 2B LM) prefix + Gemma 300M action expert suffix + flow-matching head. Inference:

  1. Embed prefix (images + language) → bidirectional prefix tokens.
  2. Forward prefix through PaliGemma → a per-layer list[(k, v)] KV cache.
  3. For each Euler step t = 1.0, 1-dt, ..., 0: embed the suffix (state + noisy actions + timestep), run the action expert with cross attention over the cached prefix K/V, predict velocity v_t, and integrate x_t = x_t + dt * v_t.
  4. Return x_0 as the action chunk (batch, action_horizon, action_dim).

We walk Gemma decoder layers ourselves in _compute_layer_* so we never go through GemmaModel.forward (which has version-dependent mask/KV-cache behaviour). Targets modern transformers (PaliGemmaForConditionalGeneration with an inner PaliGemmaModel + DynamicCache).

Reference implementations
  • OpenPI: openpi/src/openpi/models_pytorch/pi0_pytorch.py, gemma_pytorch.py
  • LeRobot: lerobot/src/lerobot/policies/pi0/modeling_pi0.py

Weight source: https://huggingface.co/lerobot/pi0_base

DEFAULT_ACTION_DIM module-attribute

DEFAULT_ACTION_DIM = 32

DEFAULT_ACTION_HORIZON module-attribute

DEFAULT_ACTION_HORIZON = 50

DEFAULT_IMAGE_RESOLUTION module-attribute

DEFAULT_IMAGE_RESOLUTION = (224, 224)

DEFAULT_MAX_TOKEN_LEN module-attribute

DEFAULT_MAX_TOKEN_LEN = 48

DEFAULT_NUM_INFERENCE_STEPS module-attribute

DEFAULT_NUM_INFERENCE_STEPS = 10

OPENPI_ATTENTION_MASK_VALUE module-attribute

OPENPI_ATTENTION_MASK_VALUE = -2.3819763e+38

logger module-attribute

logger = logging.getLogger(__name__)

GemmaVariantConfig

depth instance-attribute

depth = depth

head_dim instance-attribute

head_dim = head_dim

mlp_dim instance-attribute

mlp_dim = mlp_dim

num_heads instance-attribute

num_heads = num_heads

num_kv_heads instance-attribute

num_kv_heads = num_kv_heads

width instance-attribute

width = width

PaliGemmaWithActionExpert

Bases: Module

Dual-backbone transformer: PaliGemma (Gemma 2B) + Action Expert (Gemma 300M).

forward has two inference modes: - prefix_only: inputs_embeds=[prefix, None] + use_cache=True → compute prefix hidden states + a layer-wise K/V cache. - suffix_only: inputs_embeds=[None, suffix] + the cache from the previous prefix pass → compute expert hidden states with cross-attention over the concatenated (prefix, suffix) K/V.

Both modes walk Gemma decoder layers manually through _compute_layer_* instead of GemmaModel.forward — that way we own the attention mask handling and the KV cache format (a plain list[(k, v)] one entry per layer).

Ref: lerobot/policies/pi0/modeling_pi0.py PaliGemmaWithExpertModel

gemma_expert instance-attribute

gemma_expert = GemmaForCausalLM(
    config=action_expert_config_hf
)

paligemma instance-attribute

paligemma = PaliGemmaForConditionalGeneration(
    config=vlm_config_hf
)

embed_image

embed_image(pixel_values: Tensor) -> Tensor

Encode images with SigLIP vision tower + PaliGemma projector.

We run the two steps explicitly instead of calling PaliGemmaModel.get_image_features because that convenience method historically has returned either the projected tensor (old) or the raw vision-tower output (new). Being explicit makes us independent of which transformers release we're on.

embed_language_tokens

embed_language_tokens(tokens: Tensor) -> Tensor

Embed language tokens with PaliGemma's embedding table, returning the * sqrt(hidden)-scaled embedding (PaliGemma/Gemma convention).

The scaling location moved across transformers releases
  • ≤ 5.3: embed_tokens is a plain nn.Embedding and the * sqrt(hidden) normalizer lives inside GemmaModel.forward (which we bypass) — so we apply it explicitly here.
  • ≥ 5.4: embed_tokens is a GemmaTextScaledWordEmbedding that self-applies embed_scale = hidden_size ** 0.5 — applying it again would double-scale (≈45×). We detect this and skip the manual scale.

This makes embed_language_tokens return the canonical scaled embedding on every transformers version.

forward

forward(
    attention_mask: Tensor | None = None,
    position_ids: LongTensor | None = None,
    past_key_values=None,
    inputs_embeds: list[Tensor | None] | None = None,
    use_cache: bool = False,
)

Dispatch to prefix_only / suffix_only and return ([prefix_out, suffix_out], past_key_values_or_None).

Pi0ForActionPrediction

Bases: Module

π0 VLA model for robot action prediction via flow matching.

Inference flow
  1. Embed prefix (images + language) → prefix tokens.
  2. Forward prefix through PaliGemma → layer-wise KV cache.
  3. For each denoising step t = 1.0, 1-dt, ..., 0: a. Embed suffix (state + x_t + timestep) → suffix tokens. b. Forward suffix through the action expert with the prefix cache. c. x_t = x_t + dt * v_t (Euler integration).
  4. Return x_0 as the predicted action chunk.

action_dim instance-attribute

action_dim = getattr(
    config, "max_action_dim", DEFAULT_ACTION_DIM
)

action_horizon instance-attribute

action_horizon = getattr(
    config, "chunk_size", DEFAULT_ACTION_HORIZON
)

action_in_proj instance-attribute

action_in_proj = nn.Linear(
    self.action_dim, self.expert_width
)

action_out_proj instance-attribute

action_out_proj = nn.Linear(
    self.expert_width, self.action_dim
)

action_time_mlp_in instance-attribute

action_time_mlp_in = nn.Linear(
    2 * self.expert_width, self.expert_width
)

action_time_mlp_out instance-attribute

action_time_mlp_out = nn.Linear(
    self.expert_width, self.expert_width
)

config instance-attribute

config = config

expert_width instance-attribute

expert_width = expert_config.width

max_state_dim instance-attribute

max_state_dim = getattr(
    config, "max_state_dim", self.action_dim
)

num_inference_steps instance-attribute

num_inference_steps = getattr(
    config,
    "num_inference_steps",
    DEFAULT_NUM_INFERENCE_STEPS,
)

paligemma_with_expert instance-attribute

paligemma_with_expert = PaliGemmaWithActionExpert(
    vlm_config, expert_config
)

state_proj instance-attribute

state_proj = nn.Linear(
    self.max_state_dim, self.expert_width
)

vlm_width instance-attribute

vlm_width = vlm_config.width

denoise_step

denoise_step(
    state: Tensor,
    prefix_pad_masks: Tensor,
    past_key_values,
    x_t: Tensor,
    timestep: Tensor,
) -> Tensor

Apply one flow-matching denoising step: predict v_t from x_t.

Uses the prefix KV cache from sample_actions and only runs the action expert. Ref: openpi PI0Pytorch.denoise_step.

embed_prefix

embed_prefix(
    images: list[Tensor],
    image_masks: list[Tensor],
    lang_tokens: Tensor,
    lang_masks: Tensor,
) -> tuple[Tensor, Tensor, Tensor]

Build the prefix embeddings, per-token padding mask, and AR mask.

Prefix tokens form a contiguous sequence [img_cam_0..., img_cam_1..., ..., lang_tokens...] with bidirectional attention; the returned att_masks are all zeros because each token is free to attend to every other prefix token (the suffix pass will put a causal boundary right before the state token).

Ref: openpi PI0Pytorch.embed_prefix

embed_suffix

embed_suffix(
    state: Tensor, noisy_actions: Tensor, timestep: Tensor
) -> tuple[Tensor, Tensor, Tensor]

Build the suffix embeddings + masks: [state_token, action_tokens×H].

AR mask layout: [1, 1, 0, 0, ..., 0] — the state token is a causal boundary (no later token attends backwards through it onto prefix by mistake), the first action token starts a new causal block, and the rest of the action tokens attend to each other bidirectionally.

Ref: openpi PI0Pytorch.embed_suffix

load_weights

load_weights(weights: Iterable[tuple[str, Tensor]])

Load weights from an lerobot/pi0_base safetensors checkpoint.

Checkpoint keys (after stripping the leading model. prefix) look like paligemma_with_expert.paligemma.<sub>.*. Modern transformers nests those sub-modules one level deeper (paligemma.model.<sub>), and PaliGemma ties lm_head.weight with embed_tokens.weight at post_init time. Two remap rules make the load lossless:

  1. paligemma.{vision_tower,multi_modal_projector,language_model}.* → paligemma.model.{...}.*
  2. paligemma.lm_head.weight → paligemma.model.language_model.embed_tokens.weight (the checkpoint does not store embed_tokens.weight on its own — only the tied lm_head.weight copy — and PaliGemmaForConditionalGeneration does not register lm_head.weight as a Parameter when tied, so without this rule the language embedding would silently remain at its random init.)

On any remaining mismatch the loader logs a warning listing a sample of the unmatched checkpoint keys and model params; downstream parity tests surface these as a hard failure.

sample_actions

sample_actions(
    images: list[Tensor],
    image_masks: list[Tensor],
    lang_tokens: Tensor,
    lang_masks: Tensor,
    state: Tensor,
    noise: Tensor | None = None,
    num_steps: int | None = None,
) -> Tensor

Generate an action chunk via iterative flow-matching denoising.

Convention: t=1 is noise, t=0 is the target — opposite of the published π0 paper but matches both OpenPI and LeRobot. Ref: openpi PI0Pytorch.sample_actions

create_sinusoidal_pos_embedding

create_sinusoidal_pos_embedding(
    time: Tensor,
    dimension: int,
    min_period: float = 0.004,
    max_period: float = 4.0,
    device: device = None,
) -> Tensor

Compute a sine/cosine positional embedding for scalar timesteps.

Ref: openpi/models_pytorch/pi0_pytorch.py create_sinusoidal_pos_embedding

get_gemma_config

get_gemma_config(variant: str) -> GemmaVariantConfig

make_att_2d_masks

make_att_2d_masks(
    pad_masks: Tensor, att_masks: Tensor
) -> Tensor

Build a 2D attention mask from a padding mask and an autoregressive mask.

Ref: openpi/models_pytorch/pi0_pytorch.py make_att_2d_masks

prepare_attention_masks_4d

prepare_attention_masks_4d(att_2d_masks: Tensor) -> Tensor

Convert (B, S, S) bool masks to (B, 1, S, S) float masks.

True → 0.0 (attend), False → OPENPI_ATTENTION_MASK_VALUE. Ref: openpi PI0Pytorch._prepare_attention_masks_4d