diff --git a/src/lerobot/policies/pi0/modeling_pi0.py b/src/lerobot/policies/pi0/modeling_pi0.py index f6f4212fb..7a444f97f 100644 --- a/src/lerobot/policies/pi0/modeling_pi0.py +++ b/src/lerobot/policies/pi0/modeling_pi0.py @@ -16,7 +16,6 @@ import builtins import logging -import math from collections import deque from pathlib import Path from typing import TYPE_CHECKING, Literal, TypedDict, Unpack @@ -29,7 +28,6 @@ from lerobot.utils.import_utils import _transformers_available, require_package # Conditional import for type checking and lazy loading if TYPE_CHECKING or _transformers_available: - from transformers.cache_utils import DynamicCache from transformers.models.auto import CONFIG_MAPPING from transformers.models.gemma import modeling_gemma @@ -41,7 +39,6 @@ if TYPE_CHECKING or _transformers_available: ) else: CONFIG_MAPPING = None - DynamicCache = None modeling_gemma = None PiGemmaForCausalLM = None _gated_residual = None @@ -55,9 +52,17 @@ from lerobot.utils.constants import ( OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE, - OPENPI_ATTENTION_MASK_VALUE, ) +from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta +from ..common.vla_utils import ( + clone_past_key_values, + create_sinusoidal_pos_embedding, + make_att_2d_masks, + pad_vector, + prepare_attention_masks_4d, + resize_with_pad_torch, +) from ..pretrained import PreTrainedPolicy, T from ..rtc.modeling_rtc import RTCProcessor from .configuration_pi0 import DEFAULT_IMAGE_SIZE, PI0Config @@ -69,173 +74,6 @@ class ActionSelectKwargs(TypedDict, total=False): execution_horizon: int | None -def get_safe_dtype(target_dtype, device_type): - """Get a safe dtype for the given device type.""" - if device_type == "mps" and target_dtype == torch.float64: - return torch.float32 - if device_type == "cpu": - # CPU doesn't support bfloat16, use float32 instead - if target_dtype == torch.bfloat16: - return torch.float32 - if target_dtype == torch.float64: - return torch.float64 - return target_dtype - - -def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy) - time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu" -) -> Tensor: - """Computes sine-cosine positional embedding vectors for scalar positions.""" - if dimension % 2 != 0: - raise ValueError(f"dimension ({dimension}) must be divisible by 2") - - if time.ndim != 1: - raise ValueError("The time tensor is expected to be of shape `(batch_size, )`.") - - dtype = get_safe_dtype(torch.float64, device.type) - fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device) - period = min_period * (max_period / min_period) ** fraction - - # Compute the outer product - scaling_factor = 1.0 / period * 2 * math.pi - sin_input = scaling_factor[None, :] * time[:, None] - return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1) - - -def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy) - # Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU - alpha_t = torch.tensor(alpha, dtype=torch.float32) - beta_t = torch.tensor(beta, dtype=torch.float32) - dist = torch.distributions.Beta(alpha_t, beta_t) - return dist.sample((bsize,)).to(device) - - -def make_att_2d_masks(pad_masks, att_masks): # see openpi `make_att_2d_masks` (exact copy) - """Copied from big_vision. - - Tokens can attend to valid inputs tokens which have a cumulative mask_ar - smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to - setup several types of attention, for example: - - [[1 1 1 1 1 1]]: pure causal attention. - - [[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between - themselves and the last 3 tokens have a causal attention. The first - entry could also be a 1 without changing behaviour. - - [[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a - block can attend all previous blocks and all tokens on the same block. - - Args: - input_mask: bool[B, N] true if its part of the input, false if padding. - mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on - it and 0 where it shares the same attention mask as the previous token. - """ - if att_masks.ndim != 2: - raise ValueError(att_masks.ndim) - if pad_masks.ndim != 2: - raise ValueError(pad_masks.ndim) - - cumsum = torch.cumsum(att_masks, dim=1) - att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None] - pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None] - return att_2d_masks & pad_2d_masks - - -def clone_past_key_values(past_key_values): - """Clone the DynamicCache returned by prefix prefill for compiled denoising.""" - return DynamicCache( - tuple( - (keys.clone(), values.clone(), sliding_window) for keys, values, sliding_window in past_key_values - ) - ) - - -def pad_vector(vector, new_dim): - """Pad the last dimension of a vector to new_dim with zeros. - - Can be (batch_size x sequence_length x features_dimension) - or (batch_size x features_dimension) - """ - if vector.shape[-1] >= new_dim: - return vector - return F.pad(vector, (0, new_dim - vector.shape[-1])) - - -def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy) - images: torch.Tensor, - height: int, - width: int, - mode: str = "bilinear", -) -> torch.Tensor: - """PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion - by padding with black. If the image is float32, it must be in the range [-1, 1]. - - Args: - images: Tensor of shape [*b, h, w, c] or [*b, c, h, w] - height: Target height - width: Target width - mode: Interpolation mode ('bilinear', 'nearest', etc.) - - Returns: - Resized and padded tensor with same shape format as input - """ - # Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w] - if images.shape[-1] <= 4: # Assume channels-last format - channels_last = True - if images.dim() == 3: - images = images.unsqueeze(0) # Add batch dimension - images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w] - else: - channels_last = False - if images.dim() == 3: - images = images.unsqueeze(0) # Add batch dimension - - batch_size, channels, cur_height, cur_width = images.shape - - # Calculate resize ratio - ratio = max(cur_width / width, cur_height / height) - resized_height = int(cur_height / ratio) - resized_width = int(cur_width / ratio) - - # Resize - resized_images = F.interpolate( - images, - size=(resized_height, resized_width), - mode=mode, - align_corners=False if mode == "bilinear" else None, - ) - - # Handle dtype-specific clipping - if images.dtype == torch.uint8: - resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8) - elif images.dtype == torch.float32: - resized_images = resized_images.clamp(0.0, 1.0) - else: - raise ValueError(f"Unsupported image dtype: {images.dtype}") - - # Calculate padding - pad_h0, remainder_h = divmod(height - resized_height, 2) - pad_h1 = pad_h0 + remainder_h - pad_w0, remainder_w = divmod(width - resized_width, 2) - pad_w1 = pad_w0 + remainder_w - - # Pad - constant_value = 0 if images.dtype == torch.uint8 else 0.0 - padded_images = F.pad( - resized_images, - (pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom - mode="constant", - value=constant_value, - ) - - # Convert back to original format if needed - if channels_last: - padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c] - - return padded_images - - # Define the complete layer computation function for gradient checkpointing def compute_layer_complete(inputs_embeds, attention_mask, position_ids, adarms_cond, layers, rotary_emb): query_states = [] @@ -633,26 +471,18 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch` ) return func(*args, **kwargs) - def _prepare_attention_masks_4d(self, att_2d_masks): - """Helper method to prepare 4D attention masks for transformer.""" - att_2d_masks_4d = att_2d_masks[:, None, :, :] - return torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE) - def sample_noise(self, shape, device): - return torch.normal( - mean=0.0, - std=1.0, - size=shape, - dtype=torch.float32, - device=device, - ) + return sample_noise(shape, device) def sample_time(self, bsize, device): - time_beta = sample_beta( - self.config.time_sampling_beta_alpha, self.config.time_sampling_beta_beta, bsize, device + return sample_time_beta( + bsize, + device, + alpha=self.config.time_sampling_beta_alpha, + beta=self.config.time_sampling_beta_beta, + scale=self.config.time_sampling_scale, + offset=self.config.time_sampling_offset, ) - time = time_beta * self.config.time_sampling_scale + self.config.time_sampling_offset - return time.to(dtype=torch.float32, device=device) def embed_prefix( self, images, img_masks, lang_tokens, lang_masks @@ -783,7 +613,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch` att_2d_masks = make_att_2d_masks(pad_masks, att_masks) position_ids = torch.cumsum(pad_masks, dim=1) - 1 - att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks) + att_2d_masks_4d = prepare_attention_masks_4d(att_2d_masks) def forward_func(prefix_embs, suffix_embs, att_2d_masks_4d, position_ids, adarms_cond): (_, suffix_out), _ = self.paligemma_with_expert.forward( @@ -844,7 +674,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch` prefix_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks) prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1 - prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(prefix_att_2d_masks) + prefix_att_2d_masks_4d = prepare_attention_masks_4d(prefix_att_2d_masks) self.paligemma_with_expert.paligemma.model.language_model.config._attn_implementation = "eager" # noqa: SLF001 _, past_key_values = self.paligemma_with_expert.forward( @@ -855,44 +685,22 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch` use_cache=True, ) - dt = -1.0 / num_steps - - x_t = noise - for step in range(num_steps): - time = 1.0 + step * dt - time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize) - - def denoise_step_partial_call(input_x_t, current_timestep=time_tensor): - return self.denoise_step( - state=state, - prefix_pad_masks=prefix_pad_masks, - past_key_values=past_key_values, - x_t=input_x_t, - timestep=current_timestep, - ) - - if self._rtc_enabled(): - inference_delay = kwargs.get("inference_delay") - prev_chunk_left_over = kwargs.get("prev_chunk_left_over") - execution_horizon = kwargs.get("execution_horizon") - - v_t = self.rtc_processor.denoise_step( - x_t=x_t, - prev_chunk_left_over=prev_chunk_left_over, - inference_delay=inference_delay, - time=time, - original_denoise_step_partial=denoise_step_partial_call, - execution_horizon=execution_horizon, - ) - else: - v_t = denoise_step_partial_call(x_t) - - x_t = x_t + dt * v_t - - if self.rtc_processor is not None and self.rtc_processor.is_debug_enabled(): - self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t) - - return x_t + return euler_integrate( + lambda input_x_t, current_timestep: self.denoise_step( + state=state, + prefix_pad_masks=prefix_pad_masks, + past_key_values=past_key_values, + x_t=input_x_t, + timestep=current_timestep, + ), + noise, + num_steps, + rtc_processor=self.rtc_processor, + rtc_enabled=self._rtc_enabled(), + inference_delay=kwargs.get("inference_delay"), + prev_chunk_left_over=kwargs.get("prev_chunk_left_over"), + execution_horizon=kwargs.get("execution_horizon"), + ) def denoise_step( self, @@ -916,7 +724,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch` prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None] position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1 - full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks) + full_att_2d_masks_4d = prepare_attention_masks_4d(full_att_2d_masks) self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001 past_key_values = clone_past_key_values(past_key_values)