Robotics
Transformers
Safetensors
openvla
feature-extraction
vla
openvla-oft
xarm
task-conditioned-gate
custom_code
Instructions to use AAyano/gate_setting2_chunksize25_batch32_from20000 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use AAyano/gate_setting2_chunksize25_batch32_from20000 with Transformers:
# Load model directly from transformers import AutoModelForVision2Seq model = AutoModelForVision2Seq.from_pretrained("AAyano/gate_setting2_chunksize25_batch32_from20000", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """ | |
| modeling_prismatic.py | |
| Core HuggingFace-style PrismaticPreTrainedModel and PrismaticForConditionalGeneration class definitions. | |
| Inherits from the default `transformers.PretrainedModel`. Meant to be standalone and self-contained, | |
| but exactly replicate the logic in `prismatic.models.vlms.prismatic.py`. | |
| """ | |
| import logging | |
| import math | |
| from dataclasses import dataclass | |
| from functools import partial | |
| from typing import Any, Callable, ClassVar, Dict, List, Optional, Tuple, Union | |
| import numpy as np | |
| import timm | |
| import tokenizers | |
| import torch | |
| import torch.nn as nn | |
| import transformers | |
| from timm.models.vision_transformer import LayerScale | |
| from transformers import AutoModelForCausalLM, PretrainedConfig, PreTrainedModel | |
| from transformers.modeling_outputs import ModelOutput | |
| from prismatic.training.train_utils import ( | |
| get_current_action_mask, | |
| get_next_actions_mask, | |
| ) | |
| from prismatic.vla.constants import ( | |
| ACTION_DIM, | |
| ACTION_PROPRIO_NORMALIZATION_TYPE, | |
| ACTION_TOKEN_BEGIN_IDX, | |
| IGNORE_INDEX, | |
| NUM_ACTIONS_CHUNK, | |
| STOP_INDEX, | |
| NormalizationType, | |
| ) | |
| from .configuration_prismatic import OpenVLAConfig, PrismaticConfig | |
| # Set up logger | |
| logger = logging.getLogger(__name__) | |
| # === Utility Functions for Monkey-Patching === | |
| def unpack_tuple(fn: Callable[[Any], Tuple[Any]]) -> Callable[[Any], Any]: | |
| def wrapper(*args: Any, **kwargs: Any) -> Any: | |
| result = fn(*args, **kwargs) | |
| return result[0] if isinstance(result, tuple) else result | |
| return wrapper | |
| # HF Transformers overwrites parameters with names containing `gamma`; we're going to patch VisionBackbone.LayerScale. | |
| # =>> TIMM :: https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py#L109 | |
| # =>> Transformers :: https://github.com/huggingface/transformers/blob/main/src/transformers/modeling_utils.py#L3960 | |
| def _ls_new_forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return x.mul_(self.scale_factor) if self.inplace else x * self.scale_factor | |
| def ls_apply_patch(ls_module: LayerScale): | |
| ls_module.scale_factor = nn.Parameter(ls_module.gamma.clone()) | |
| ls_module.forward = _ls_new_forward.__get__(ls_module, LayerScale) | |
| del ls_module.gamma | |
| # === Prismatic Vision Backbone (nn.Module) Definitions (w/ Fused Backbone Support) === | |
| class PrismaticVisionBackbone(nn.Module): | |
| """ | |
| Vision backbone for Prismatic models that handles image feature extraction. | |
| Supports both single backbone (e.g., SigLIP) and fused backbone (e.g., SigLIP + DINOv2) configurations. | |
| For fused backbones, features from both models are concatenated along the feature dimension. | |
| """ | |
| def __init__( | |
| self, | |
| use_fused_vision_backbone: bool, | |
| image_sizes: List[int], | |
| timm_model_ids: List[str], | |
| timm_override_act_layers: List[Optional[str]], | |
| ) -> None: | |
| """ | |
| Initialize the vision backbone. | |
| Args: | |
| use_fused_vision_backbone: Whether to use two backbones and fuse their features | |
| image_sizes: List of image sizes for each backbone | |
| timm_model_ids: List of TIMM model IDs to use for each backbone | |
| timm_override_act_layers: List of activation layer overrides for each backbone | |
| """ | |
| super().__init__() | |
| self.use_fused_vision_backbone = use_fused_vision_backbone | |
| self.num_images_in_input = 1 # Default value, can be overridden later | |
| # Validate number of (fused) vision backbones | |
| if len(timm_model_ids) > 2: | |
| raise ValueError("Prismatic models only support up to 2 (fused) vision backbones!") | |
| # Create primary featurizer | |
| self.featurizer = self._create_featurizer( | |
| model_id=timm_model_ids[0], img_size=image_sizes[0], act_layer=timm_override_act_layers[0] | |
| ) | |
| self.embed_dim = self.featurizer.embed_dim | |
| # Create secondary featurizer if using fused backbone | |
| if self.use_fused_vision_backbone: | |
| self.fused_featurizer = self._create_featurizer( | |
| model_id=timm_model_ids[1], img_size=image_sizes[1], act_layer=timm_override_act_layers[1] | |
| ) | |
| self.embed_dim += self.fused_featurizer.embed_dim | |
| # Patch LayerScale modules for HF compatibility | |
| self._patch_layer_scales() | |
| def _create_featurizer(self, model_id: str, img_size: int, act_layer: Optional[str]) -> nn.Module: | |
| """ | |
| Create a TIMM-based featurizer model with appropriate configurations. | |
| Args: | |
| model_id: The TIMM model ID to load | |
| img_size: Input image size for the model | |
| act_layer: Override for the activation layer type | |
| Returns: | |
| A configured featurizer model | |
| """ | |
| featurizer = timm.create_model( | |
| model_id, | |
| pretrained=False, | |
| num_classes=0, | |
| img_size=img_size, | |
| act_layer=act_layer, | |
| ) | |
| # Monkey-patch the forward function to extract the second-to-last layer features | |
| num_blocks = len(featurizer.blocks) | |
| featurizer.forward = unpack_tuple(partial(featurizer.get_intermediate_layers, n={num_blocks - 2})) | |
| return featurizer | |
| def _patch_layer_scales(self) -> None: | |
| """ | |
| Patch all LayerScale modules to be compatible with HF's parameter naming. | |
| HF Transformers overwrites parameters with names containing 'gamma', | |
| so we need to rename and modify the forward method. | |
| """ | |
| # Patch primary featurizer | |
| for module in self.featurizer.modules(): | |
| if isinstance(module, LayerScale): | |
| ls_apply_patch(module) | |
| # Patch secondary featurizer if it exists | |
| if self.use_fused_vision_backbone: | |
| for module in self.fused_featurizer.modules(): | |
| if isinstance(module, LayerScale): | |
| ls_apply_patch(module) | |
| def get_num_patches(self) -> int: | |
| """ | |
| Returns the number of vision patches output by the vision backbone. | |
| Returns: | |
| Number of patches per image | |
| """ | |
| return self.featurizer.patch_embed.num_patches | |
| def get_num_images_in_input(self) -> int: | |
| """ | |
| Returns the number of input images for the vision backbone. | |
| Returns: | |
| Number of images expected in the input | |
| """ | |
| return self.num_images_in_input | |
| def set_num_images_in_input(self, num_images_in_input: int) -> None: | |
| """ | |
| Sets the number of input images for the vision backbone. | |
| Args: | |
| num_images_in_input: Number of images to expect in the input | |
| """ | |
| self.num_images_in_input = num_images_in_input | |
| def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Implements the forward pass for the vision backbone. | |
| If `self.use_fused_vision_backbone == True`, uses both SigLIP and DINOv2 transformers to extract visual features | |
| (otherwise uses SigLIP only). Allows multi-image inputs (but only for fused vision backbone). | |
| Args: | |
| pixel_values (torch.Tensor): Pixels for input image(s), (B, C, H, W). | |
| """ | |
| if self.num_images_in_input == 1: | |
| if not self.use_fused_vision_backbone: | |
| return self.featurizer(pixel_values) | |
| # Split `pixel_values :: [bsz, 2 * 3, resolution, resolution]` =>> featurize =>> channel stack | |
| img, img_fused = torch.split(pixel_values, [3, 3], dim=1) | |
| patches, patches_fused = self.featurizer(img), self.fused_featurizer(img_fused) | |
| return torch.cat([patches, patches_fused], dim=2) | |
| else: | |
| assert self.use_fused_vision_backbone, "Multi-image inputs require using fused backbone!" | |
| # Split `pixel_values` into individual images (each with 6 channels: 3 for SigLIP + 3 for DINOv2) | |
| images = torch.split(pixel_values, [6] * self.num_images_in_input, dim=1) | |
| # Process each image and collect patches | |
| all_patches = [] | |
| for img in images: | |
| # Split each image further into two stacks of channels (each with 3 channels) | |
| img_regular, img_fused = torch.split(img, [3, 3], dim=1) | |
| # Get patches from both SigLIP and DINOv2 vision transformers | |
| patches = self.featurizer(img_regular) | |
| patches_fused = self.fused_featurizer(img_fused) | |
| # Concatenate SigLIP and DINOv2 patches along the hidden dimension | |
| combined_patches = torch.cat([patches, patches_fused], dim=2) | |
| all_patches.append(combined_patches) | |
| # Concatenate all patches along the patch dimension | |
| return torch.cat(all_patches, dim=1) | |
| # === Prismatic Projector (nn.Module) Definitions === | |
| class PrismaticProjector(nn.Module): | |
| def __init__(self, use_fused_vision_backbone: bool, vision_dim: int, llm_dim: int) -> None: | |
| super().__init__() | |
| self.use_fused_vision_backbone = use_fused_vision_backbone | |
| self.vision_dim, self.llm_dim = vision_dim, llm_dim | |
| # Switch on `use_fused_vision_backbone` =>> use slightly different MLPs and projection factors! | |
| if not self.use_fused_vision_backbone: | |
| self.fc1 = nn.Linear(self.vision_dim, self.llm_dim, bias=True) | |
| self.fc2 = nn.Linear(self.llm_dim, self.llm_dim, bias=True) | |
| self.act_fn1 = nn.GELU() | |
| else: | |
| initial_projection_dim = 4 * vision_dim | |
| self.fc1 = nn.Linear(self.vision_dim, initial_projection_dim, bias=True) | |
| self.fc2 = nn.Linear(initial_projection_dim, self.llm_dim, bias=True) | |
| self.fc3 = nn.Linear(self.llm_dim, self.llm_dim, bias=True) | |
| self.act_fn1 = nn.GELU() | |
| self.act_fn2 = nn.GELU() | |
| def forward(self, img_patches: torch.Tensor) -> torch.Tensor: | |
| if not self.use_fused_vision_backbone: | |
| projected_features = self.fc1(img_patches) | |
| projected_features = self.act_fn1(projected_features) | |
| projected_features = self.fc2(projected_features) | |
| else: | |
| projected_features = self.fc1(img_patches) | |
| projected_features = self.act_fn1(projected_features) | |
| projected_features = self.fc2(projected_features) | |
| projected_features = self.act_fn2(projected_features) | |
| projected_features = self.fc3(projected_features) | |
| return projected_features | |
| class GazingSelfAttentionBlock(nn.Module): | |
| def __init__(self, dim: int, num_heads: int, ffn_mult: int = 4): | |
| super().__init__() | |
| self.attn_norm = nn.LayerNorm(dim) | |
| self.attn = nn.MultiheadAttention(dim, num_heads, batch_first=True, dropout=0.0) | |
| self.ffn_norm = nn.LayerNorm(dim) | |
| self.ffn = nn.Sequential( | |
| nn.Linear(dim, ffn_mult * dim), | |
| nn.GELU(), | |
| nn.Linear(ffn_mult * dim, dim), | |
| ) | |
| def forward(self, tokens: torch.Tensor) -> torch.Tensor: | |
| attn_input = self.attn_norm(tokens) | |
| attn_output, _ = self.attn(attn_input, attn_input, attn_input, need_weights=False) | |
| tokens = tokens + attn_output | |
| tokens = tokens + self.ffn(self.ffn_norm(tokens)) | |
| return tokens | |
| class TextConditionedTokenGate(nn.Module): | |
| def __init__( | |
| self, | |
| llm_dim: int, | |
| use_text_summary: bool = True, | |
| use_vision_tokens: bool = True, | |
| hidden_dim: int = 512, | |
| text_pooling_mode: str = "mean", | |
| text_pool_hidden_dim: int = 128, | |
| cross_attention_dim: int = 256, | |
| cross_attention_heads: int = 1, | |
| gate_mlp_depth: int = 1, | |
| contrastive_visual_tau: float = 0.1, | |
| contrastive_text_tau: float = 1.0, | |
| gazing_mode: str = "mlp", | |
| self_attn_dim: int = 512, | |
| self_attn_heads: int = 8, | |
| self_attn_layers: int = 1, | |
| ): | |
| super().__init__() | |
| if text_pooling_mode not in {"mean", "mlp", "cross_attention", "contrastive_alignment_score"}: | |
| raise ValueError( | |
| "`text_pooling_mode` must be one of " | |
| "['mean', 'mlp', 'cross_attention', 'contrastive_alignment_score'], " | |
| f"got {text_pooling_mode!r}" | |
| ) | |
| if gate_mlp_depth <= 0: | |
| raise ValueError(f"`gate_mlp_depth` must be positive, got {gate_mlp_depth}") | |
| if contrastive_visual_tau <= 0: | |
| raise ValueError(f"`contrastive_visual_tau` must be positive, got {contrastive_visual_tau}") | |
| if contrastive_text_tau <= 0: | |
| raise ValueError(f"`contrastive_text_tau` must be positive, got {contrastive_text_tau}") | |
| if gazing_mode not in {"mlp", "self_attention"}: | |
| raise ValueError(f"`gazing_mode` must be one of ['mlp', 'self_attention'], got {gazing_mode!r}") | |
| if self_attn_dim <= 0: | |
| raise ValueError(f"`self_attn_dim` must be positive, got {self_attn_dim}") | |
| if self_attn_heads <= 0: | |
| raise ValueError(f"`self_attn_heads` must be positive, got {self_attn_heads}") | |
| if self_attn_dim % self_attn_heads != 0: | |
| raise ValueError( | |
| "`self_attn_dim` must be divisible by `self_attn_heads`; " | |
| f"got {self_attn_dim} and {self_attn_heads}" | |
| ) | |
| if self_attn_layers <= 0: | |
| raise ValueError(f"`self_attn_layers` must be positive, got {self_attn_layers}") | |
| self.llm_dim = llm_dim | |
| self.use_text_summary = use_text_summary | |
| self.use_vision_tokens = use_vision_tokens | |
| self.hidden_dim = hidden_dim | |
| self.text_pooling_mode = text_pooling_mode | |
| self.text_pool_hidden_dim = text_pool_hidden_dim | |
| self.cross_attention_dim = cross_attention_dim | |
| self.cross_attention_heads = cross_attention_heads | |
| self.cross_attention_head_dim = None | |
| self.gate_mlp_depth = gate_mlp_depth | |
| self.contrastive_visual_tau = contrastive_visual_tau | |
| self.contrastive_text_tau = contrastive_text_tau | |
| self.gazing_mode = gazing_mode | |
| self.self_attn_dim = self_attn_dim | |
| self.self_attn_heads = self_attn_heads | |
| self.self_attn_layers = self_attn_layers | |
| self.text_pool = None | |
| self.text_query = None | |
| self.text_key = None | |
| self.text_value = None | |
| self.text_output = None | |
| self.visual_attn_ln = None | |
| self.text_attn_ln = None | |
| self.gate = None | |
| self.context_proj = None | |
| self.gazing_blocks = nn.ModuleList() | |
| self.gate_out = None | |
| self.gazing_output_ln = None | |
| self.gate_visual_ln = nn.LayerNorm(llm_dim) | |
| self.gate_text_ln = nn.LayerNorm(llm_dim) | |
| self.last_text_pool_mask = None | |
| self.last_contrastive_scores = None | |
| gate_input_dim = 2 * llm_dim if self.use_text_summary and self.use_vision_tokens else llm_dim | |
| if self.text_pooling_mode == "mlp": | |
| self.text_pool = nn.Sequential( | |
| nn.Linear(llm_dim, text_pool_hidden_dim), | |
| nn.Tanh(), | |
| nn.Linear(text_pool_hidden_dim, 1), | |
| ) | |
| elif self.text_pooling_mode == "cross_attention": | |
| if self.cross_attention_heads <= 0: | |
| raise ValueError( | |
| "`cross_attention_heads` must be positive, " | |
| f"got {self.cross_attention_heads}" | |
| ) | |
| if self.cross_attention_dim % self.cross_attention_heads != 0: | |
| raise ValueError( | |
| "`cross_attention_dim` must be divisible by `cross_attention_heads`; got " | |
| f"{self.cross_attention_dim} and {self.cross_attention_heads}" | |
| ) | |
| self.cross_attention_head_dim = self.cross_attention_dim // self.cross_attention_heads | |
| self.visual_attn_ln = nn.LayerNorm(llm_dim) | |
| self.text_attn_ln = nn.LayerNorm(llm_dim) | |
| self.text_query = nn.Linear(llm_dim, cross_attention_dim, bias=False) | |
| self.text_key = nn.Linear(llm_dim, cross_attention_dim, bias=False) | |
| self.text_value = nn.Linear(llm_dim, cross_attention_dim, bias=False) | |
| self.text_output = nn.Linear(cross_attention_dim, llm_dim, bias=False) | |
| if self.gazing_mode == "mlp": | |
| gate_layers = [nn.Linear(gate_input_dim, hidden_dim), nn.GELU()] | |
| for _ in range(gate_mlp_depth - 1): | |
| gate_layers.extend([nn.Linear(hidden_dim, hidden_dim), nn.GELU()]) | |
| gate_layers.append(nn.Linear(hidden_dim, 1)) | |
| self.gate = nn.Sequential(*gate_layers) | |
| else: | |
| self.context_proj = nn.Linear(gate_input_dim, self_attn_dim) | |
| self.gazing_blocks = nn.ModuleList( | |
| [GazingSelfAttentionBlock(self_attn_dim, self_attn_heads) for _ in range(self_attn_layers)] | |
| ) | |
| self.gazing_output_ln = nn.LayerNorm(self_attn_dim) | |
| self.gate_out = nn.Linear(self_attn_dim, 1) | |
| def _pool_text_embeddings( | |
| language_embeddings: torch.Tensor, language_attention_mask: Optional[torch.Tensor] = None | |
| ) -> torch.Tensor: | |
| if language_attention_mask is None: | |
| return language_embeddings.mean(dim=1) | |
| text_mask = language_attention_mask.to(dtype=language_embeddings.dtype).unsqueeze(-1) # [B, T, 1] | |
| valid_token_counts = text_mask.sum(dim=1).clamp(min=1.0) # [B, 1] | |
| masked_language_embeddings = language_embeddings * text_mask | |
| return masked_language_embeddings.sum(dim=1) / valid_token_counts | |
| def forward( | |
| self, | |
| visual_tokens: torch.Tensor, | |
| language_embeddings: torch.Tensor, | |
| language_attention_mask: Optional[torch.Tensor] = None, | |
| language_text_pool_mask: Optional[torch.Tensor] = None, | |
| patches_per_image: Optional[int] = None, | |
| ): | |
| # visual_tokens: [B, N, D] | |
| # language_embeddings: [B, T, D] | |
| self.last_text_pool_mask = None | |
| self.last_contrastive_scores = None | |
| text_pool_mask = language_text_pool_mask | |
| if text_pool_mask is not None: | |
| text_pool_mask = text_pool_mask.to(device=language_embeddings.device, dtype=torch.bool) | |
| if language_attention_mask is not None: | |
| text_pool_mask = text_pool_mask & language_attention_mask.to(device=language_embeddings.device).bool() | |
| has_pool_token = text_pool_mask.any(dim=1, keepdim=True) | |
| if not has_pool_token.all(): | |
| fallback_mask = ( | |
| language_attention_mask.to(device=language_embeddings.device).bool() | |
| if language_attention_mask is not None | |
| else torch.ones_like(text_pool_mask, dtype=torch.bool) | |
| ) | |
| logger.warning( | |
| "text pool mask has no valid tokens for at least one sample; falling back to attention mask." | |
| ) | |
| text_pool_mask = torch.where(has_pool_token, text_pool_mask, fallback_mask) | |
| else: | |
| text_pool_mask = language_attention_mask | |
| if text_pool_mask is not None: | |
| text_pool_mask = text_pool_mask.to(device=language_embeddings.device, dtype=torch.bool) | |
| if text_pool_mask is not None: | |
| self.last_text_pool_mask = text_pool_mask.detach().bool() | |
| text_pool_weights = None | |
| if not self.use_text_summary: | |
| gate_input = self.gate_visual_ln(visual_tokens) | |
| elif self.text_pooling_mode == "mean": | |
| # Pool only valid language tokens so padding does not leak into the text summary. | |
| pooled_text = self._pool_text_embeddings(language_embeddings, text_pool_mask) # [B, D] | |
| pooled_text = pooled_text.unsqueeze(1).expand(-1, visual_tokens.shape[1], -1) # [B, N, D] | |
| elif self.text_pooling_mode == "mlp": | |
| text_pool_scores = self.text_pool(language_embeddings).squeeze(-1) # [B, T] | |
| if text_pool_mask is not None: | |
| text_pool_scores = text_pool_scores.masked_fill(~text_pool_mask.bool(), -1e9) | |
| text_pool_weights = torch.softmax(text_pool_scores, dim=1) # [B, T] | |
| if text_pool_mask is not None: | |
| text_pool_weights = text_pool_weights * text_pool_mask.to(dtype=text_pool_weights.dtype) | |
| text_pool_weights = text_pool_weights / text_pool_weights.sum(dim=1, keepdim=True).clamp(min=1e-6) | |
| pooled_text = torch.sum(text_pool_weights.unsqueeze(-1) * language_embeddings, dim=1) # [B, D] | |
| pooled_text = pooled_text.unsqueeze(1).expand(-1, visual_tokens.shape[1], -1) # [B, N, D] | |
| elif self.text_pooling_mode == "cross_attention": | |
| attn_visual_tokens = self.visual_attn_ln(visual_tokens) | |
| attn_language_embeddings = self.text_attn_ln(language_embeddings) | |
| text_queries = self.text_query(attn_visual_tokens) # [B, N, A] | |
| text_keys = self.text_key(attn_language_embeddings) # [B, T, A] | |
| text_values = self.text_value(attn_language_embeddings) # [B, T, A] | |
| batch_size, num_visual_tokens, _ = text_queries.shape | |
| num_text_tokens = text_keys.shape[1] | |
| num_heads, head_dim = self.cross_attention_heads, self.cross_attention_head_dim | |
| text_queries = text_queries.view(batch_size, num_visual_tokens, num_heads, head_dim).transpose(1, 2) | |
| text_keys = text_keys.view(batch_size, num_text_tokens, num_heads, head_dim).transpose(1, 2) | |
| text_values = text_values.view(batch_size, num_text_tokens, num_heads, head_dim).transpose(1, 2) | |
| text_pool_scores = torch.matmul(text_queries, text_keys.transpose(-1, -2)) / math.sqrt( | |
| head_dim | |
| ) # [B, H, N, T] | |
| if text_pool_mask is not None: | |
| expanded_text_pool_mask = text_pool_mask.unsqueeze(1).unsqueeze(1).bool() | |
| text_pool_scores = text_pool_scores.masked_fill(~expanded_text_pool_mask, -1e9) | |
| text_pool_weights = torch.softmax(text_pool_scores, dim=-1) # [B, H, N, T] | |
| if text_pool_mask is not None: | |
| text_pool_weights = text_pool_weights * text_pool_mask.unsqueeze(1).unsqueeze(1).to( | |
| dtype=text_pool_weights.dtype | |
| ) | |
| text_pool_weights = text_pool_weights / text_pool_weights.sum(dim=-1, keepdim=True).clamp(min=1e-6) | |
| pooled_text = torch.matmul(text_pool_weights, text_values) # [B, H, N, Dh] | |
| pooled_text = pooled_text.transpose(1, 2).contiguous().view( | |
| batch_size, num_visual_tokens, self.cross_attention_dim | |
| ) | |
| pooled_text = self.text_output(pooled_text) # [B, N, D] | |
| text_pool_weights = text_pool_weights.mean(dim=1) # [B, N, T] | |
| else: | |
| visual_norm = torch.nn.functional.normalize(visual_tokens.float(), dim=-1) | |
| text_norm = torch.nn.functional.normalize(language_embeddings.float(), dim=-1) | |
| similarity = torch.matmul(text_norm, visual_norm.transpose(1, 2)) # [B, T, N] | |
| visual_weights = torch.softmax(similarity / self.contrastive_visual_tau, dim=-1) | |
| local_evidence = torch.sum(visual_weights * similarity, dim=-1) # [B, T] | |
| global_baseline = similarity.mean(dim=-1) # [B, T] | |
| contrastive_scores = local_evidence - global_baseline | |
| self.last_contrastive_scores = contrastive_scores.detach() | |
| if text_pool_mask is not None: | |
| contrastive_scores = contrastive_scores.masked_fill(~text_pool_mask.bool(), -1e9) | |
| text_pool_weights = torch.softmax(contrastive_scores / self.contrastive_text_tau, dim=1) # [B, T] | |
| if text_pool_mask is not None: | |
| text_pool_weights = text_pool_weights * text_pool_mask.to(dtype=text_pool_weights.dtype) | |
| text_pool_weights = text_pool_weights / text_pool_weights.sum(dim=1, keepdim=True).clamp(min=1e-6) | |
| pooled_text = torch.sum( | |
| text_pool_weights.to(dtype=language_embeddings.dtype).unsqueeze(-1) * language_embeddings, | |
| dim=1, | |
| ) # [B, D] | |
| pooled_text = pooled_text.unsqueeze(1).expand(-1, visual_tokens.shape[1], -1) # [B, N, D] | |
| if self.use_text_summary: | |
| gate_text_tokens = self.gate_text_ln(pooled_text) | |
| if self.use_vision_tokens: | |
| gate_visual_tokens = self.gate_visual_ln(visual_tokens) | |
| gate_input = torch.cat([gate_visual_tokens, gate_text_tokens], dim=-1) # [B, N, 2D] | |
| else: | |
| gate_input = gate_text_tokens # [B, N, D] | |
| if self.gazing_mode == "mlp": | |
| gate_logits = self.gate(gate_input) # [B, N, 1] | |
| else: | |
| gazing_tokens = self.context_proj(gate_input) | |
| gazing_tokens = self._apply_gazing_blocks(gazing_tokens, patches_per_image) | |
| gazing_tokens = self.gazing_output_ln(gazing_tokens) | |
| gate_logits = self.gate_out(gazing_tokens) # [B, N, 1] | |
| token_gate = torch.sigmoid(gate_logits) # [B, N, 1] | |
| gated_visual_tokens = visual_tokens * token_gate | |
| return gated_visual_tokens, token_gate, text_pool_weights | |
| def _apply_gazing_blocks(self, tokens: torch.Tensor, patches_per_image: Optional[int]) -> torch.Tensor: | |
| if patches_per_image is None or patches_per_image <= 0 or tokens.shape[1] % patches_per_image != 0: | |
| if patches_per_image is not None: | |
| logger.warning( | |
| "Invalid patches_per_image for self-attention gazing; falling back to full visual sequence. " | |
| f"tokens={tokens.shape[1]}, patches_per_image={patches_per_image}" | |
| ) | |
| for block in self.gazing_blocks: | |
| tokens = block(tokens) | |
| return tokens | |
| batch_size, num_visual_tokens, dim = tokens.shape | |
| num_images = num_visual_tokens // patches_per_image | |
| tokens = tokens.reshape(batch_size * num_images, patches_per_image, dim) | |
| for block in self.gazing_blocks: | |
| tokens = block(tokens) | |
| return tokens.reshape(batch_size, num_images * patches_per_image, dim) | |
| # === Main HF Class Definitions === | |
| class PrismaticCausalLMOutputWithPast(ModelOutput): | |
| """Base class for Prismatic casual (visually-conditioned) language model outputs; also exposes visual features.""" | |
| loss: Optional[torch.FloatTensor] = None | |
| logits: torch.FloatTensor = None | |
| past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None | |
| hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None | |
| attentions: Optional[Tuple[torch.FloatTensor]] = None | |
| # Additions for VLMs | |
| projector_features: Optional[torch.FloatTensor] = None | |
| class PrismaticPreTrainedModel(PreTrainedModel): | |
| config_class: PretrainedConfig = PrismaticConfig | |
| base_model_prefix: str = "model" | |
| supports_gradient_checkpointing: bool = True | |
| _no_split_modules: ClassVar[List[str]] = ["PrismaticProjector", "TextConditionedTokenGate"] | |
| _skip_keys_device_placement: str = "past_key_values" | |
| _supports_flash_attn_2: bool = True | |
| def _init_weights(self, module: nn.Module) -> None: | |
| # Important :: this HF ported version is *not* meant for training from scratch; only inference and fine-tuning! | |
| # => As such, this init_weights code is not correct; if training VLMs from scratch, use the main codebase at | |
| # https://github.com/TRI-ML/prismatic-vlms | |
| std = ( | |
| self.config.initializer_range | |
| if hasattr(self.config, "initializer_range") | |
| else self.config.text_config.initializer_range | |
| ) | |
| if hasattr(module, "class_embedding"): | |
| module.class_embedding.data.normal_(mean=0.0, std=std) | |
| if isinstance(module, (nn.Linear, nn.Conv2d)): | |
| module.weight.data.normal_(mean=0.0, std=std) | |
| if module.bias is not None: | |
| module.bias.data.zero_() | |
| elif isinstance(module, nn.Embedding): | |
| module.weight.data.normal_(mean=0.0, std=std) | |
| if module.padding_idx is not None: | |
| module.weight.data[module.padding_idx].zero_() | |
| def _supports_sdpa(self) -> bool: | |
| """Check LLM supports SDPA Attention""" | |
| return self.language_model._supports_sdpa | |
| class PrismaticForConditionalGeneration(PrismaticPreTrainedModel): | |
| def __init__(self, config: PrismaticConfig) -> None: | |
| super().__init__(config) | |
| # [Validation] Lightweight Validate on `config` Fields + Dependency Versions | |
| if config.use_fused_vision_backbone is None: | |
| raise ValueError("Missing config field `use_fused_vision_backbone`") | |
| if timm.__version__ not in {"0.9.10", "0.9.11", "0.9.12", "0.9.16"}: | |
| raise NotImplementedError( | |
| "TIMM Version must be >= 0.9.10 and < 1.0.0 (breaking); please raise a GitHub Issue " | |
| "if you urgently need support for latest TIMM versions." | |
| ) | |
| if (transformers.__version__ != "4.40.1") or (tokenizers.__version__ != "0.19.1"): | |
| logger.warning( | |
| f"Expected `transformers==4.40.1` and `tokenizers==0.19.1` but got " | |
| f"`transformers=={transformers.__version__}` and `tokenizers=={tokenizers.__version__}`; " | |
| f"there might be inference-time regressions due to dependency changes. If in doubt, please" | |
| f"use the above versions." | |
| ) | |
| # Instantiate PrismaticVisionBackbone (w/ Potential Fused Backbone) | |
| self.vision_backbone = PrismaticVisionBackbone( | |
| config.use_fused_vision_backbone, config.image_sizes, config.timm_model_ids, config.timm_override_act_layers | |
| ) | |
| # Create Multimodal Projector | |
| self.projector = PrismaticProjector( | |
| config.use_fused_vision_backbone, | |
| vision_dim=self.vision_backbone.embed_dim, | |
| llm_dim=config.text_config.hidden_size, | |
| ) | |
| # Instantiate LLM Backbone | |
| self.language_model = AutoModelForCausalLM.from_config( | |
| config.text_config, attn_implementation=config._attn_implementation | |
| ) | |
| self.vocab_size = config.text_config.vocab_size | |
| self.pad_token_id = config.pad_token_id | |
| self.llm_dim = config.text_config.hidden_size | |
| self.use_text_token_gate = config.use_text_token_gate | |
| self.text_token_gate_use_text_summary = getattr(config, "text_token_gate_use_text_summary", True) | |
| self.text_token_gate_use_vision_tokens = getattr(config, "text_token_gate_use_vision_tokens", True) | |
| self.text_token_gate_budget = config.text_token_gate_budget | |
| self.text_token_gate_budget_loss_weight = config.text_token_gate_budget_loss_weight | |
| self.text_token_gate_linear_mean_penalty_weight = config.text_token_gate_linear_mean_penalty_weight | |
| self.text_token_gate_pool_instruction_only = getattr(config, "text_token_gate_pool_instruction_only", False) | |
| self.text_token_gate_hidden_dim = config.text_token_gate_hidden_dim | |
| self.text_token_gate_text_pooling_mode = config.text_token_gate_text_pooling_mode | |
| self.text_token_gate_text_pool_hidden_dim = config.text_token_gate_text_pool_hidden_dim | |
| self.text_token_gate_cross_attention_dim = config.text_token_gate_cross_attention_dim | |
| self.text_token_gate_cross_attention_heads = getattr(config, "text_token_gate_cross_attention_heads", 1) | |
| self.text_token_gate_mlp_depth = getattr(config, "text_token_gate_mlp_depth", 1) | |
| self.contrastive_visual_tau = getattr(config, "contrastive_visual_tau", 0.1) | |
| self.contrastive_text_tau = getattr(config, "contrastive_text_tau", 1.0) | |
| self.gazing_mode = getattr(config, "gazing_mode", "mlp") | |
| self.self_attn_dim = getattr(config, "self_attn_dim", 512) | |
| self.self_attn_heads = getattr(config, "self_attn_heads", 8) | |
| self.self_attn_layers = getattr(config, "self_attn_layers", 1) | |
| self.layer_gate_mode = getattr(config, "layer_gate_mode", "none") | |
| self.layer_gate_threshold = getattr(config, "layer_gate_threshold", 0.15) | |
| self.layer_gate_strength = getattr(config, "layer_gate_strength", 0.5) | |
| self.text_token_gate = None | |
| if self.use_text_token_gate: | |
| self.text_token_gate = TextConditionedTokenGate( | |
| llm_dim=config.text_config.hidden_size, | |
| use_text_summary=self.text_token_gate_use_text_summary, | |
| use_vision_tokens=self.text_token_gate_use_vision_tokens, | |
| hidden_dim=self.text_token_gate_hidden_dim, | |
| text_pooling_mode=self.text_token_gate_text_pooling_mode, | |
| text_pool_hidden_dim=self.text_token_gate_text_pool_hidden_dim, | |
| cross_attention_dim=self.text_token_gate_cross_attention_dim, | |
| cross_attention_heads=self.text_token_gate_cross_attention_heads, | |
| gate_mlp_depth=self.text_token_gate_mlp_depth, | |
| contrastive_visual_tau=self.contrastive_visual_tau, | |
| contrastive_text_tau=self.contrastive_text_tau, | |
| gazing_mode=self.gazing_mode, | |
| self_attn_dim=self.self_attn_dim, | |
| self_attn_heads=self.self_attn_heads, | |
| self_attn_layers=self.self_attn_layers, | |
| ) | |
| self._last_token_gate_budget_loss = None | |
| self._last_token_gate_linear_mean_penalty = None | |
| self._last_token_gate_regularization_loss = None | |
| self._debug_last_token_gate = None | |
| self._debug_last_token_gate_mean = None | |
| self._debug_last_token_gate_budget_loss = None | |
| self._debug_last_token_gate_linear_mean_penalty = None | |
| self._debug_last_text_pool_weights = None | |
| self._debug_last_text_pool_mask = None | |
| self._debug_last_contrastive_scores = None | |
| # HF Boilerplate =>> initializes weights via `_init_weights()` and sets gradient checkpointing | |
| self.post_init() | |
| # === `PreTrainedModel` Boilerplate === | |
| def get_input_embeddings(self) -> nn.Module: | |
| return self.language_model.get_input_embeddings() | |
| def set_input_embeddings(self, value: nn.Module) -> None: | |
| self.language_model.set_input_embeddings(value) | |
| def get_output_embeddings(self) -> nn.Module: | |
| return self.language_model.get_output_embeddings() | |
| def set_output_embeddings(self, new_embeddings: nn.Module) -> None: | |
| self.language_model.set_output_embeddings(new_embeddings) | |
| def get_decoder(self) -> nn.Module: | |
| return self.language_model.get_decoder() | |
| def set_decoder(self, decoder: nn.Module) -> None: | |
| self.language_model.set_decoder(decoder) | |
| def tie_weights(self) -> None: | |
| self.language_model.tie_weights() # Note: `Llama-2` and `Mistral` don't tie weights (no-op) | |
| def resize_token_embeddings( | |
| self, new_num_tokens: Optional[int] = None, pad_to_multiple_of: Optional[int] = None | |
| ) -> nn.Embedding: | |
| updated_embeddings = self.language_model.resize_token_embeddings(new_num_tokens, pad_to_multiple_of) | |
| # Update config/instance variables | |
| self.config.text_config.vocab_size = updated_embeddings.num_embeddings | |
| self.vocab_size = updated_embeddings.num_embeddings | |
| return updated_embeddings | |
| def _configure_text_token_gate_from_config(self) -> None: | |
| self.use_text_token_gate = self.config.use_text_token_gate | |
| self.text_token_gate_use_text_summary = getattr(self.config, "text_token_gate_use_text_summary", True) | |
| self.text_token_gate_use_vision_tokens = getattr(self.config, "text_token_gate_use_vision_tokens", True) | |
| self.text_token_gate_budget = self.config.text_token_gate_budget | |
| self.text_token_gate_budget_loss_weight = self.config.text_token_gate_budget_loss_weight | |
| self.text_token_gate_linear_mean_penalty_weight = self.config.text_token_gate_linear_mean_penalty_weight | |
| self.text_token_gate_pool_instruction_only = getattr(self.config, "text_token_gate_pool_instruction_only", False) | |
| self.text_token_gate_hidden_dim = self.config.text_token_gate_hidden_dim | |
| self.text_token_gate_text_pooling_mode = self.config.text_token_gate_text_pooling_mode | |
| self.text_token_gate_text_pool_hidden_dim = self.config.text_token_gate_text_pool_hidden_dim | |
| self.text_token_gate_cross_attention_dim = self.config.text_token_gate_cross_attention_dim | |
| self.text_token_gate_cross_attention_heads = getattr(self.config, "text_token_gate_cross_attention_heads", 1) | |
| self.text_token_gate_mlp_depth = getattr(self.config, "text_token_gate_mlp_depth", 1) | |
| self.contrastive_visual_tau = getattr(self.config, "contrastive_visual_tau", 0.1) | |
| self.contrastive_text_tau = getattr(self.config, "contrastive_text_tau", 1.0) | |
| self.gazing_mode = getattr(self.config, "gazing_mode", "mlp") | |
| self.self_attn_dim = getattr(self.config, "self_attn_dim", 512) | |
| self.self_attn_heads = getattr(self.config, "self_attn_heads", 8) | |
| self.self_attn_layers = getattr(self.config, "self_attn_layers", 1) | |
| self.layer_gate_mode = getattr(self.config, "layer_gate_mode", "none") | |
| self.layer_gate_threshold = getattr(self.config, "layer_gate_threshold", 0.15) | |
| self.layer_gate_strength = getattr(self.config, "layer_gate_strength", 0.5) | |
| if not self.use_text_token_gate: | |
| self.text_token_gate = None | |
| self._last_token_gate_budget_loss = None | |
| self._last_token_gate_linear_mean_penalty = None | |
| self._last_token_gate_regularization_loss = None | |
| self._debug_last_token_gate = None | |
| self._debug_last_token_gate_mean = None | |
| self._debug_last_token_gate_budget_loss = None | |
| self._debug_last_token_gate_linear_mean_penalty = None | |
| self._debug_last_text_pool_weights = None | |
| self._debug_last_text_pool_mask = None | |
| self._debug_last_contrastive_scores = None | |
| return | |
| needs_new_gate = self.text_token_gate is None | |
| if not needs_new_gate: | |
| gate_module = self.text_token_gate | |
| if hasattr(gate_module, "original_module"): | |
| gate_module = gate_module.original_module | |
| elif hasattr(gate_module, "modules_to_save") and len(gate_module.modules_to_save) > 0: | |
| gate_module = next(iter(gate_module.modules_to_save.values())) | |
| gate_hidden_dim = getattr(gate_module, "hidden_dim", None) | |
| if gate_hidden_dim is None and getattr(gate_module, "gate", None) is not None: | |
| gate_hidden_dim = gate_module.gate[0].out_features | |
| gate_text_pooling_mode = getattr(gate_module, "text_pooling_mode", None) | |
| gate_text_pool_hidden_dim = getattr(gate_module, "text_pool_hidden_dim", None) | |
| gate_cross_attention_dim = getattr(gate_module, "cross_attention_dim", None) | |
| gate_cross_attention_heads = getattr(gate_module, "cross_attention_heads", 1) | |
| gate_mlp_depth = getattr(gate_module, "gate_mlp_depth", 1) | |
| gate_contrastive_visual_tau = getattr(gate_module, "contrastive_visual_tau", 0.1) | |
| gate_contrastive_text_tau = getattr(gate_module, "contrastive_text_tau", 1.0) | |
| gate_gazing_mode = getattr(gate_module, "gazing_mode", "mlp") | |
| gate_use_text_summary = getattr(gate_module, "use_text_summary", True) | |
| gate_use_vision_tokens = getattr(gate_module, "use_vision_tokens", True) | |
| gate_self_attn_dim = getattr(gate_module, "self_attn_dim", 512) | |
| gate_self_attn_heads = getattr(gate_module, "self_attn_heads", 8) | |
| gate_self_attn_layers = getattr(gate_module, "self_attn_layers", 1) | |
| needs_new_gate = ( | |
| gate_use_text_summary != self.text_token_gate_use_text_summary | |
| or gate_use_vision_tokens != self.text_token_gate_use_vision_tokens | |
| or | |
| gate_hidden_dim != self.text_token_gate_hidden_dim | |
| or gate_text_pooling_mode != self.text_token_gate_text_pooling_mode | |
| or gate_mlp_depth != self.text_token_gate_mlp_depth | |
| or gate_gazing_mode != self.gazing_mode | |
| or ( | |
| self.text_token_gate_text_pooling_mode == "mlp" | |
| and gate_text_pool_hidden_dim != self.text_token_gate_text_pool_hidden_dim | |
| ) | |
| or ( | |
| self.text_token_gate_text_pooling_mode == "cross_attention" | |
| and gate_cross_attention_dim != self.text_token_gate_cross_attention_dim | |
| ) | |
| or ( | |
| self.text_token_gate_text_pooling_mode == "cross_attention" | |
| and gate_cross_attention_heads != self.text_token_gate_cross_attention_heads | |
| ) | |
| or ( | |
| self.text_token_gate_text_pooling_mode == "contrastive_alignment_score" | |
| and gate_contrastive_visual_tau != self.contrastive_visual_tau | |
| ) | |
| or ( | |
| self.text_token_gate_text_pooling_mode == "contrastive_alignment_score" | |
| and gate_contrastive_text_tau != self.contrastive_text_tau | |
| ) | |
| or (self.gazing_mode == "self_attention" and gate_self_attn_dim != self.self_attn_dim) | |
| or (self.gazing_mode == "self_attention" and gate_self_attn_heads != self.self_attn_heads) | |
| or (self.gazing_mode == "self_attention" and gate_self_attn_layers != self.self_attn_layers) | |
| ) | |
| if needs_new_gate: | |
| self.text_token_gate = TextConditionedTokenGate( | |
| llm_dim=self.llm_dim, | |
| use_text_summary=self.text_token_gate_use_text_summary, | |
| use_vision_tokens=self.text_token_gate_use_vision_tokens, | |
| hidden_dim=self.text_token_gate_hidden_dim, | |
| text_pooling_mode=self.text_token_gate_text_pooling_mode, | |
| text_pool_hidden_dim=self.text_token_gate_text_pool_hidden_dim, | |
| cross_attention_dim=self.text_token_gate_cross_attention_dim, | |
| cross_attention_heads=self.text_token_gate_cross_attention_heads, | |
| gate_mlp_depth=self.text_token_gate_mlp_depth, | |
| contrastive_visual_tau=self.contrastive_visual_tau, | |
| contrastive_text_tau=self.contrastive_text_tau, | |
| gazing_mode=self.gazing_mode, | |
| self_attn_dim=self.self_attn_dim, | |
| self_attn_heads=self.self_attn_heads, | |
| self_attn_layers=self.self_attn_layers, | |
| ) | |
| self.text_token_gate.apply(self._init_weights) | |
| reference_param = next(self.projector.parameters()) | |
| self.text_token_gate = self.text_token_gate.to( | |
| device=reference_param.device, | |
| dtype=reference_param.dtype, | |
| ) | |
| def _replace_input_embeddings(self, input_embeddings, all_actions_mask, noisy_action_features): | |
| """ | |
| Replace embeddings in input_embeddings at positions where all_actions_mask is True | |
| with embeddings from noisy_action_features, using vectorized operations. | |
| Args: | |
| input_embeddings: Tensor of shape (B, S, D) | |
| all_actions_mask: Boolean tensor of shape (B, S) | |
| noisy_action_features: Tensor of shape (B, K, D) where K is the number of True values in mask per sample | |
| Returns: | |
| Modified input_embeddings tensor | |
| """ | |
| # Clone input to avoid modifying the original tensor | |
| new_input_embeddings = input_embeddings.clone() | |
| # Create a tensor with the same shape of input_embeddings to hold the noisy action features | |
| repositioned_noisy_action_features = torch.zeros_like(input_embeddings) | |
| # Create batch indices for splicing | |
| batch_indices = torch.arange(input_embeddings.shape[0], device=input_embeddings.device) | |
| batch_indices = batch_indices.unsqueeze(1).expand(-1, noisy_action_features.shape[1]) | |
| # Get indices where mask is True for each sample | |
| masked_indices = torch.stack([torch.where(mask)[0] for mask in all_actions_mask]) | |
| # Move the noisy action features into their correct positions | |
| repositioned_noisy_action_features[batch_indices, masked_indices] = noisy_action_features | |
| # Combine original input embeddings and noisy action embeddings using the mask | |
| new_input_embeddings = torch.where( | |
| all_actions_mask.unsqueeze(-1), repositioned_noisy_action_features, new_input_embeddings | |
| ) | |
| return new_input_embeddings | |
| def _process_action_masks(self, labels): | |
| """Helper to get action masks from labels""" | |
| current_action_mask = get_current_action_mask(labels) | |
| next_actions_mask = get_next_actions_mask(labels) | |
| all_actions_mask = current_action_mask | next_actions_mask # (B, seq_len) | |
| return all_actions_mask | |
| def _process_vision_features(self, pixel_values, language_embeddings=None, use_film=False): | |
| """Process vision features with optional FiLM conditioning""" | |
| if use_film: | |
| # FiLM: Infuse language inputs into visual features | |
| patch_features = self.vision_backbone(pixel_values, language_embeddings) # (bsz, 256 * num_images, D) | |
| else: | |
| patch_features = self.vision_backbone(pixel_values) # (bsz, 256 * num_images, D) | |
| # Project patch embeddings into language embedding space | |
| return self.projector(patch_features) | |
| def _apply_text_token_gate( | |
| self, | |
| projected_patch_embeddings, | |
| language_embeddings, | |
| language_attention_mask=None, | |
| language_text_pool_mask=None, | |
| ): | |
| patches_per_image = self.vision_backbone.get_num_patches() | |
| gated_tokens, token_gate, text_pool_weights = self.text_token_gate( | |
| projected_patch_embeddings, | |
| language_embeddings, | |
| language_attention_mask, | |
| language_text_pool_mask, | |
| patches_per_image=patches_per_image, | |
| ) | |
| gate_mean_per_sample = token_gate.mean(dim=1).squeeze(-1) # [B], global mean for logging | |
| if patches_per_image > 0 and token_gate.shape[1] % patches_per_image == 0: | |
| num_images = token_gate.shape[1] // patches_per_image | |
| gate_mean_per_image = token_gate.reshape( | |
| token_gate.shape[0], num_images, patches_per_image, token_gate.shape[-1] | |
| ).mean(dim=2).squeeze(-1) # [B, num_images] | |
| budget_loss = ((gate_mean_per_image - self.text_token_gate_budget) ** 2).mean() | |
| linear_mean_penalty = gate_mean_per_image.mean() | |
| else: | |
| logger.warning( | |
| "Could not split token gate by image for per-camera budget loss; " | |
| f"tokens={token_gate.shape[1]}, patches_per_image={patches_per_image}. " | |
| "Falling back to global token-gate mean." | |
| ) | |
| budget_loss = ((gate_mean_per_sample - self.text_token_gate_budget) ** 2).mean() | |
| linear_mean_penalty = gate_mean_per_sample.mean() | |
| regularization_terms = [] | |
| if self.text_token_gate_budget_loss_weight != 0: | |
| regularization_terms.append(self.text_token_gate_budget_loss_weight * budget_loss) | |
| if self.text_token_gate_linear_mean_penalty_weight != 0: | |
| regularization_terms.append(self.text_token_gate_linear_mean_penalty_weight * linear_mean_penalty) | |
| regularization_loss = None | |
| if regularization_terms: | |
| regularization_loss = sum(regularization_terms) | |
| self._last_token_gate_budget_loss = budget_loss | |
| self._last_token_gate_linear_mean_penalty = linear_mean_penalty | |
| self._last_token_gate_regularization_loss = regularization_loss | |
| self._debug_last_token_gate = token_gate.detach() | |
| self._debug_last_token_gate_mean = token_gate.mean().detach() | |
| self._debug_last_token_gate_budget_loss = budget_loss.detach() | |
| self._debug_last_token_gate_linear_mean_penalty = linear_mean_penalty.detach() | |
| self._debug_last_text_pool_weights = text_pool_weights.detach() if text_pool_weights is not None else None | |
| self._debug_last_text_pool_mask = getattr(self.text_token_gate, "last_text_pool_mask", None) | |
| self._debug_last_contrastive_scores = getattr(self.text_token_gate, "last_contrastive_scores", None) | |
| return gated_tokens, token_gate, regularization_loss | |
| def _get_language_model_layers(self) -> Optional[nn.ModuleList]: | |
| model = getattr(self.language_model, "model", self.language_model) | |
| if hasattr(model, "layers"): | |
| return model.layers | |
| decoder = getattr(model, "decoder", None) | |
| if decoder is not None and hasattr(decoder, "layers"): | |
| return decoder.layers | |
| return None | |
| def _build_layer_gate(self, token_gate: Optional[torch.Tensor]) -> Optional[torch.Tensor]: | |
| if ( | |
| token_gate is None | |
| or self.layer_gate_mode == "none" | |
| or self.layer_gate_strength <= 0 | |
| or self.layer_gate_mode != "threshold" | |
| ): | |
| return None | |
| gate = token_gate.clamp(min=0.0, max=1.0) | |
| low_gate_mask = (gate < self.layer_gate_threshold).to(dtype=gate.dtype) | |
| layer_gate = 1.0 - low_gate_mask * self.layer_gate_strength * (1.0 - gate) | |
| return layer_gate | |
| def _run_language_model_with_layer_gate( | |
| self, | |
| token_gate: Optional[torch.Tensor] = None, | |
| visual_start_idx: int = 1, | |
| **language_model_kwargs, | |
| ): | |
| layer_gate = self._build_layer_gate(token_gate) | |
| if layer_gate is None: | |
| return self.language_model(**language_model_kwargs) | |
| layers = self._get_language_model_layers() | |
| if layers is None: | |
| logger.warning("Could not find language model decoder layers; skipping layer gate.") | |
| return self.language_model(**language_model_kwargs) | |
| visual_token_count = layer_gate.shape[1] | |
| def apply_to_hidden_states(hidden_states: torch.Tensor) -> torch.Tensor: | |
| visual_end_idx = visual_start_idx + visual_token_count | |
| if hidden_states.ndim != 3 or hidden_states.shape[1] < visual_end_idx: | |
| return hidden_states | |
| gate = layer_gate.to(device=hidden_states.device, dtype=hidden_states.dtype) | |
| return torch.cat( | |
| [ | |
| hidden_states[:, :visual_start_idx, :], | |
| hidden_states[:, visual_start_idx:visual_end_idx, :] * gate, | |
| hidden_states[:, visual_end_idx:, :], | |
| ], | |
| dim=1, | |
| ) | |
| def pre_hook_with_kwargs(module, args, kwargs): | |
| if args and torch.is_tensor(args[0]): | |
| args = (apply_to_hidden_states(args[0]),) + args[1:] | |
| elif "hidden_states" in kwargs and torch.is_tensor(kwargs["hidden_states"]): | |
| kwargs = dict(kwargs) | |
| kwargs["hidden_states"] = apply_to_hidden_states(kwargs["hidden_states"]) | |
| return args, kwargs | |
| def pre_hook(module, args): | |
| if args and torch.is_tensor(args[0]): | |
| return (apply_to_hidden_states(args[0]),) + args[1:] | |
| return args | |
| handles = [] | |
| try: | |
| for layer in layers: | |
| try: | |
| handles.append(layer.register_forward_pre_hook(pre_hook_with_kwargs, with_kwargs=True)) | |
| except TypeError: | |
| handles.append(layer.register_forward_pre_hook(pre_hook)) | |
| return self.language_model(**language_model_kwargs) | |
| finally: | |
| for handle in handles: | |
| handle.remove() | |
| def _process_proprio_features(self, projected_patch_embeddings, proprio, proprio_projector): | |
| """Process proprioceptive features and append to vision features""" | |
| if proprio_projector is not None and proprio is not None: | |
| # projected_patch_embeddings: (bsz, num_patches * num_images, llm_dim) | |
| # proprio: (bsz, proprio_dim) or (propro_dim,) | |
| proprio = proprio.reshape(projected_patch_embeddings.shape[0], -1) # (bsz, proprio_dim) | |
| proprio_features = proprio_projector(proprio) # (bsz, llm_dim) | |
| proprio_features = proprio_features.unsqueeze(dim=1) # (bsz, 1, llm_dim) | |
| # For simplicity, just append proprio token to the end of projected vision patch tokens | |
| return torch.cat((projected_patch_embeddings, proprio_features), dim=1) | |
| return projected_patch_embeddings | |
| def _build_multimodal_attention(self, input_embeddings, projected_patch_embeddings, attention_mask): | |
| """Build multimodal embeddings and attention mask""" | |
| # Update attention mask | |
| projected_patch_attention_mask = None | |
| if attention_mask is not None: | |
| projected_patch_attention_mask = torch.full( | |
| (projected_patch_embeddings.shape[0], projected_patch_embeddings.shape[1]), | |
| fill_value=True, | |
| dtype=attention_mask.dtype, | |
| device=attention_mask.device, | |
| ) | |
| # Build multimodal embeddings & attention mask; insert embeddings after <BOS> token (1:) | |
| multimodal_embeddings = torch.cat( | |
| [input_embeddings[:, :1, :], projected_patch_embeddings, input_embeddings[:, 1:, :]], dim=1 | |
| ) | |
| multimodal_attention_mask = None | |
| if attention_mask is not None: | |
| multimodal_attention_mask = torch.cat( | |
| [attention_mask[:, :1], projected_patch_attention_mask, attention_mask[:, 1:]], dim=1 | |
| ) | |
| return multimodal_embeddings, multimodal_attention_mask | |
| def _build_multimodal_labels(self, labels, projected_patch_embeddings): | |
| """Build multimodal labels with IGNORE_INDEX for patch embeddings""" | |
| if labels is not None: | |
| projected_patch_labels = torch.full( | |
| (projected_patch_embeddings.shape[0], projected_patch_embeddings.shape[1]), | |
| fill_value=IGNORE_INDEX, | |
| dtype=labels.dtype, | |
| device=labels.device, | |
| ) | |
| return torch.cat([labels[:, :1], projected_patch_labels, labels[:, 1:]], dim=1) | |
| return None | |
| # === Core Prismatic VLM `forward()` Logic === | |
| def forward( | |
| self, | |
| input_ids: Optional[torch.LongTensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| pixel_values: Optional[torch.FloatTensor] = None, | |
| labels: Optional[torch.LongTensor] = None, | |
| text_token_pool_mask: Optional[torch.Tensor] = None, | |
| inputs_embeds: Optional[torch.FloatTensor] = None, | |
| past_key_values: Optional[List[torch.FloatTensor]] = None, | |
| use_cache: Optional[bool] = None, | |
| output_attentions: Optional[bool] = None, | |
| output_hidden_states: Optional[bool] = None, | |
| output_projector_features: Optional[bool] = None, | |
| return_dict: Optional[bool] = None, | |
| proprio=None, | |
| proprio_projector=None, | |
| noisy_actions=None, | |
| noisy_action_projector=None, | |
| diffusion_timestep_embeddings=None, | |
| use_film: bool = False, | |
| ) -> Union[Tuple, PrismaticCausalLMOutputWithPast]: | |
| """Run a forward pass through the VLM, returning a PrismaticCausalLMOutputWithPast instance.""" | |
| output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions | |
| output_hidden_states = ( | |
| output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states | |
| ) | |
| output_projector_features = output_projector_features if output_projector_features is not None else False | |
| return_dict = return_dict if return_dict is not None else self.config.use_return_dict | |
| # Respect `use_cache` only if not training (even if `gradient_checkpointing` is off) | |
| use_cache = use_cache and not self.training | |
| # Clear token-gate stats so they never leak across forwards. | |
| self._last_token_gate_budget_loss = None | |
| self._last_token_gate_linear_mean_penalty = None | |
| self._last_token_gate_regularization_loss = None | |
| self._debug_last_token_gate = None | |
| self._debug_last_token_gate_mean = None | |
| self._debug_last_token_gate_budget_loss = None | |
| self._debug_last_token_gate_linear_mean_penalty = None | |
| self._debug_last_text_pool_weights = None | |
| self._debug_last_text_pool_mask = None | |
| self._debug_last_contrastive_scores = None | |
| # Instantiate Placeholder for Projector Features | |
| projected_patch_embeddings = None | |
| token_gate = None | |
| token_gate_regularization_loss = None | |
| # === Handle Generation with Cache (`input_ids.shape[1] == 1`) =>> requires `past_keys_values` === | |
| if input_ids.shape[1] == 1: | |
| assert input_ids.shape[0] == 1, "Generation is only currently supported for batch size of 1!" | |
| assert past_key_values is not None, "You must provide `past_key_values` during cached generation!" | |
| assert labels is None, "Unexpected key `labels` provided during cached generation!" | |
| language_model_output = self.language_model( | |
| input_ids=input_ids, | |
| attention_mask=None, | |
| position_ids=None, | |
| past_key_values=past_key_values, | |
| inputs_embeds=None, | |
| labels=None, | |
| use_cache=use_cache, | |
| output_attentions=output_attentions, | |
| output_hidden_states=output_hidden_states, | |
| return_dict=return_dict, | |
| ) | |
| # === Handle Unimodal Forward === | |
| elif pixel_values is None: | |
| assert (input_ids is not None) and (inputs_embeds is None), "Missing `input_ids` in language-only forward!" | |
| assert past_key_values is None, "Unexpected key `past_key_values` provided during language-only forward!" | |
| language_model_output = self.language_model( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| position_ids=None, | |
| past_key_values=None, | |
| inputs_embeds=None, | |
| labels=labels, | |
| use_cache=use_cache, | |
| output_attentions=output_attentions, | |
| output_hidden_states=output_hidden_states, | |
| return_dict=return_dict, | |
| ) | |
| # === Handle Multimodal Forward === | |
| elif (input_ids.shape[0] == pixel_values.shape[0]) or (inputs_embeds.shape[0] == pixel_values.shape[0]): | |
| assert past_key_values is None, "Unexpected key `past_key_values` provided during multimodal forward!" | |
| # Get input embeddings (from language model embeddings) | |
| input_embeddings = self.get_input_embeddings()(input_ids) # (B, seq_len, D) | |
| # Extract action masks | |
| all_actions_mask = self._process_action_masks(labels) | |
| # Extract the language portion of the input embeddings (i.e. remove the action tokens portion) | |
| language_embeddings = input_embeddings[~all_actions_mask].reshape( | |
| input_embeddings.shape[0], -1, input_embeddings.shape[2] | |
| ) # (B, lang_seq_len, llm_dim) | |
| language_attention_mask = None | |
| if attention_mask is not None: | |
| language_attention_mask = attention_mask[~all_actions_mask].reshape(input_embeddings.shape[0], -1) | |
| language_text_pool_mask = None | |
| if text_token_pool_mask is not None: | |
| text_token_pool_mask = text_token_pool_mask.to(device=input_embeddings.device, dtype=torch.bool) | |
| if text_token_pool_mask.shape != all_actions_mask.shape: | |
| logger.warning( | |
| "text_token_pool_mask shape does not match input/action mask shape; " | |
| f"got {text_token_pool_mask.shape} vs {all_actions_mask.shape}. Ignoring pool mask." | |
| ) | |
| else: | |
| language_text_pool_mask = text_token_pool_mask[~all_actions_mask].reshape( | |
| input_embeddings.shape[0], -1 | |
| ) | |
| # Get visual features | |
| projected_patch_embeddings = self._process_vision_features(pixel_values, language_embeddings, use_film) | |
| # Apply the gate only to projected visual tokens before proprio is appended. | |
| if self.use_text_token_gate: | |
| projected_patch_embeddings, token_gate, token_gate_regularization_loss = self._apply_text_token_gate( | |
| projected_patch_embeddings, | |
| language_embeddings, | |
| language_attention_mask, | |
| language_text_pool_mask, | |
| ) | |
| # Add proprioceptive state if provided | |
| projected_patch_embeddings = self._process_proprio_features( | |
| projected_patch_embeddings, proprio, proprio_projector | |
| ) | |
| # [Diffusion] Add diffusion timestep embedding if provided | |
| if diffusion_timestep_embeddings is not None: | |
| # For simplicity, just append diffusion timestep embedding to the end of projected vision patch tokens | |
| projected_patch_embeddings = torch.cat( | |
| (projected_patch_embeddings, diffusion_timestep_embeddings), dim=1 | |
| ) | |
| # Process action embeddings | |
| if noisy_actions is not None: | |
| # Get mask corresponding to all action tokens | |
| all_actions_mask = self._process_action_masks(labels) | |
| # Reshape noisy actions into individual action tokens | |
| # noisy_actions: (B, chunk_len, action_dim) -> (B, chunk_len * action_dim, 1) | |
| B = noisy_actions.shape[0] | |
| noisy_actions = noisy_actions.reshape(B, -1).unsqueeze(-1) | |
| # Project noisy action tokens into language model embedding space | |
| noisy_action_features = noisy_action_projector(noisy_actions) # (B, chunk_len * action_dim, llm_dim) | |
| # Replace embeddings of the action tokens with noisy action embeddings | |
| input_embeddings = self._replace_input_embeddings( | |
| input_embeddings, all_actions_mask, noisy_action_features | |
| ) | |
| else: | |
| # Replace the embeddings of the action tokens with zeros | |
| # (Later on, the positional embeddings will be added to them) | |
| all_actions_mask = all_actions_mask.unsqueeze(-1) # (B, seq_len, 1) | |
| input_embeddings = input_embeddings * ~all_actions_mask | |
| # Build multimodal embeddings & attention mask | |
| multimodal_embeddings, multimodal_attention_mask = self._build_multimodal_attention( | |
| input_embeddings, projected_patch_embeddings, attention_mask | |
| ) | |
| # Build labels for multimodal sequence if needed | |
| multimodal_labels = self._build_multimodal_labels(labels, projected_patch_embeddings) | |
| # Dispatch to language model | |
| language_model_output = self._run_language_model_with_layer_gate( | |
| token_gate=token_gate, | |
| input_ids=None, | |
| attention_mask=multimodal_attention_mask, | |
| position_ids=None, | |
| past_key_values=None, | |
| inputs_embeds=multimodal_embeddings, | |
| labels=multimodal_labels, | |
| use_cache=use_cache, | |
| output_attentions=output_attentions, | |
| output_hidden_states=output_hidden_states, | |
| return_dict=return_dict, | |
| ) | |
| # === Otherwise =>> Assume Invalid! === | |
| elif (input_ids.shape[0] != pixel_values.shape[0]) or (inputs_embeds.shape[0] != pixel_values.shape[0]): | |
| raise ValueError("Non-homogenous batch of (text, image) input -- forward() does not support mixed batches!") | |
| else: | |
| raise ValueError( | |
| "Invalid PrismaticForConditionalGeneration `forward()` call with provided arguments:\n" | |
| f"=> `input_ids` = {input_ids is not None}\n" | |
| f"=> `attention_mask` = {attention_mask is not None}\n" | |
| f"=> `pixel_values` = {pixel_values is not None}\n" | |
| f"=> `labels` = {labels is not None}\n" | |
| f"=> `input_embeds` = {inputs_embeds is not None}\n" | |
| f"=> `past_key_values` = {past_key_values is not None}\n" | |
| f"=> `use_cache` = {use_cache}" | |
| ) | |
| total_loss = language_model_output.loss if return_dict else (language_model_output[0] if labels is not None else None) | |
| if token_gate_regularization_loss is not None and total_loss is not None: | |
| total_loss = total_loss + token_gate_regularization_loss | |
| # Unpack `language_model_output` and return PrismaticCausalLMOutputWithPast (or tuple if not `return_dict`) | |
| if not return_dict: | |
| if total_loss is not None: | |
| language_model_output = (total_loss,) + tuple(language_model_output[1:]) | |
| if output_projector_features and (projected_patch_embeddings is not None): | |
| return *language_model_output, projected_patch_embeddings | |
| return language_model_output | |
| return PrismaticCausalLMOutputWithPast( | |
| loss=total_loss, | |
| logits=language_model_output.logits, | |
| past_key_values=language_model_output.past_key_values, | |
| hidden_states=language_model_output.hidden_states, | |
| attentions=language_model_output.attentions, | |
| projector_features=projected_patch_embeddings, | |
| ) | |
| # === GenerationMixin Methods === | |
| def prepare_inputs_for_generation( | |
| self, | |
| input_ids: Optional[torch.Tensor] = None, | |
| past_key_values: Optional[List[torch.FloatTensor]] = None, | |
| inputs_embeds: Optional[torch.FloatTensor] = None, | |
| pixel_values: Optional[torch.FloatTensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| **kwargs: str, | |
| ) -> Dict[str, torch.Tensor]: | |
| """Borrowed from `LlamaForCausalLM` and simplified for batch size = 1; mirrors original PrismaticVLM logic.""" | |
| if ((input_ids is not None) and (input_ids.shape[0] > 1)) or ( | |
| (inputs_embeds is not None) and (inputs_embeds.shape[0] > 1) | |
| ): | |
| raise ValueError("Generation with batch size > 1 is not currently supported!") | |
| # Handle `past_key_values` (cache) =>> assume `input_ids` just has unprocessed tokens | |
| if past_key_values is not None: | |
| input_ids = input_ids[:, -1:] | |
| # If `input_embeds` are passed, we only want to use them in the 1st generation step | |
| if inputs_embeds is not None and past_key_values is None: | |
| model_inputs = {"input_embeds": inputs_embeds} | |
| else: | |
| model_inputs = {"input_ids": input_ids} | |
| # Make sure `pixel_values` are preserved in `model_inputs` | |
| model_inputs.update( | |
| { | |
| "attention_mask": attention_mask, | |
| "pixel_values": pixel_values, | |
| "past_key_values": past_key_values, | |
| "use_cache": kwargs.get("use_cache"), | |
| } | |
| ) | |
| return model_inputs | |
| # Defer to Language Model (all handle this differently, with different return types) | |
| def _reorder_cache(self, *args, **kwargs) -> Any: | |
| return self.language_model._reorder_cache(*args, **kwargs) | |
| class OpenVLAForActionPrediction(PrismaticForConditionalGeneration): | |
| config_class: PretrainedConfig = OpenVLAConfig | |
| def __init__(self, config: OpenVLAConfig) -> None: | |
| super().__init__(config) | |
| self.norm_stats = config.norm_stats | |
| # Compute action bins | |
| self.bins = np.linspace(-1, 1, config.n_action_bins) | |
| self.bin_centers = (self.bins[:-1] + self.bins[1:]) / 2.0 | |
| # Compute vocab size for de-tokenization -- revert added "multiple of" | |
| self.vocab_size = self.config.text_config.vocab_size - self.config.pad_to_multiple_of | |
| def _prepare_input_for_action_prediction(self, input_ids, attention_mask): | |
| """Prepares input for action prediction by adding necessary tokens""" | |
| # Add (ACTION_DIM * NUM_ACTIONS_CHUNK) placeholder tokens to input_ids to simulate action tokens | |
| placeholder_action_token_ids = ( | |
| torch.ones((input_ids.shape[0], ACTION_DIM * NUM_ACTIONS_CHUNK)).to(input_ids.device).to(input_ids.dtype) | |
| ) | |
| input_ids = torch.cat([input_ids, placeholder_action_token_ids], dim=-1) | |
| # Add stop token to sequence (needed in non-causal bi-directional self-attention, as it appears at train time) | |
| stop_token_id = torch.ones((input_ids.shape[0], 1)).to(input_ids.device).to(input_ids.dtype) * STOP_INDEX | |
| input_ids = torch.cat([input_ids, stop_token_id], dim=-1) | |
| # Extend the attention mask to fit the new shape of input | |
| # Note: Only batch size == 1 supported right now | |
| mask_extension = ( | |
| torch.ones((attention_mask.shape[0], input_ids.shape[-1] - attention_mask.shape[-1])) | |
| .to(attention_mask.device) | |
| .to(attention_mask.dtype) | |
| ) | |
| attention_mask = torch.cat([attention_mask, mask_extension], dim=-1) | |
| return input_ids, attention_mask | |
| def _prepare_labels_for_action_prediction(self, labels, input_ids): | |
| """Creates labels tensor for action prediction if not provided""" | |
| # Extend labels tensor with fake action labels | |
| ARBITRARY_ACTION_TOKEN_IDX = ACTION_TOKEN_BEGIN_IDX + 1 | |
| labels_extension = ( | |
| torch.ones((labels.shape[0], input_ids.shape[-1] - labels.shape[-1])).to(labels.device).to(labels.dtype) | |
| * ARBITRARY_ACTION_TOKEN_IDX | |
| ) | |
| labels = torch.cat([labels, labels_extension], dim=-1) | |
| # Replace last label token with stop token | |
| labels[:, -1] = STOP_INDEX | |
| return labels | |
| def _unnormalize_actions(self, normalized_actions, unnorm_key=None): | |
| """Unnormalize actions using dataset statistics""" | |
| action_norm_stats = self.get_action_stats(unnorm_key) | |
| if ACTION_PROPRIO_NORMALIZATION_TYPE == NormalizationType.BOUNDS: | |
| mask = action_norm_stats.get("mask", np.ones_like(action_norm_stats["min"], dtype=bool)) | |
| action_high, action_low = np.array(action_norm_stats["max"]), np.array(action_norm_stats["min"]) | |
| elif ACTION_PROPRIO_NORMALIZATION_TYPE == NormalizationType.BOUNDS_Q99: | |
| mask = action_norm_stats.get("mask", np.ones_like(action_norm_stats["q01"], dtype=bool)) | |
| action_high, action_low = np.array(action_norm_stats["q99"]), np.array(action_norm_stats["q01"]) | |
| else: | |
| raise ValueError("Unsupported action/proprio normalization type detected!") | |
| actions = np.where( | |
| mask, | |
| 0.5 * (normalized_actions + 1) * (action_high - action_low + 1e-8) + action_low, | |
| normalized_actions, | |
| ) | |
| return actions | |
| def _run_diffusion_prediction( | |
| self, | |
| input_embeddings, | |
| all_actions_mask, | |
| noise, | |
| action_head, | |
| projected_patch_embeddings, | |
| labels, | |
| attention_mask, | |
| NUM_PATCHES, | |
| NUM_PROMPT_TOKENS, | |
| noisy_action_projector, | |
| token_gate=None, | |
| ): | |
| """Run diffusion-based action prediction""" | |
| # Clone embedding for reuse in each timestep | |
| orig_projected_patch_embeddings = projected_patch_embeddings.clone() | |
| curr_noisy_actions = noise | |
| # Reverse diffusion: Iteratively denoise to generate action prediction | |
| for t in action_head.noise_scheduler.timesteps: | |
| # Get diffusion model's noise prediction (conditioned on VLA latent embedding, current noisy action | |
| # embedding, and diffusion timestep embedding) | |
| timesteps = torch.Tensor([t]).to(labels.device) | |
| diffusion_timestep_embeddings = ( | |
| action_head.time_encoder(timesteps).to(curr_noisy_actions.dtype).to(curr_noisy_actions.device) | |
| ) # (B, llm_dim) | |
| diffusion_timestep_embeddings = diffusion_timestep_embeddings.unsqueeze(1) # (B, 1, llm_dim) | |
| # [Diffusion] Replace the embeddings of the action tokens with noisy actions | |
| # (Later on, the positional embeddings will be added to them) | |
| # For simplicity, append diffusion timestep embedding to the end of projected vision tokens | |
| projected_patch_embeddings = torch.cat( | |
| (orig_projected_patch_embeddings, diffusion_timestep_embeddings), dim=1 | |
| ) | |
| # Reshape and project noisy actions into language embedding space | |
| B = curr_noisy_actions.shape[0] | |
| orig_curr_noisy_actions_shape = curr_noisy_actions.shape | |
| curr_noisy_actions = curr_noisy_actions.reshape(B, -1).unsqueeze(-1) | |
| noisy_action_features = noisy_action_projector(curr_noisy_actions) | |
| curr_noisy_actions = curr_noisy_actions.reshape(orig_curr_noisy_actions_shape) | |
| # Replace action token embeddings with noisy action embeddings | |
| input_embeddings = self._replace_input_embeddings( | |
| input_embeddings.clone(), all_actions_mask, noisy_action_features | |
| ) | |
| # Build multimodal embeddings and attention mask | |
| multimodal_embeddings, multimodal_attention_mask = self._build_multimodal_attention( | |
| input_embeddings, projected_patch_embeddings, attention_mask | |
| ) | |
| # Forward pass through language model | |
| language_model_output = self._run_language_model_with_layer_gate( | |
| token_gate=token_gate, | |
| input_ids=None, | |
| attention_mask=multimodal_attention_mask, | |
| position_ids=None, | |
| past_key_values=None, | |
| inputs_embeds=multimodal_embeddings, | |
| labels=None, | |
| use_cache=None, | |
| output_attentions=False, | |
| output_hidden_states=True, | |
| return_dict=True, | |
| ) | |
| # Extract hidden states for action portion of response | |
| last_hidden_states = language_model_output.hidden_states[-1] # (B, seq_len, D) | |
| actions_hidden_states = last_hidden_states[ | |
| :, | |
| NUM_PATCHES + NUM_PROMPT_TOKENS : NUM_PATCHES + NUM_PROMPT_TOKENS + ACTION_DIM * NUM_ACTIONS_CHUNK, | |
| :, | |
| ] # (B, act_chunk_len, D) | |
| # Predict noise and update noisy actions: x_t -> x_{t-1} | |
| noise_pred = action_head.predict_noise(actions_hidden_states) | |
| curr_noisy_actions = action_head.noise_scheduler.step(noise_pred, t, curr_noisy_actions).prev_sample | |
| curr_noisy_actions = curr_noisy_actions.reshape(NUM_ACTIONS_CHUNK, ACTION_DIM) | |
| # Return final actions | |
| return curr_noisy_actions.float().cpu().detach().numpy(), actions_hidden_states | |
| def _regression_or_discrete_prediction( | |
| self, | |
| input_embeddings, | |
| all_actions_mask, | |
| projected_patch_embeddings, | |
| attention_mask, | |
| labels, | |
| NUM_PATCHES, | |
| NUM_PROMPT_TOKENS, | |
| action_head=None, | |
| token_gate=None, | |
| ): | |
| """Run L1 regression-based continuous action prediction or discrete action tokens prediction.""" | |
| # Zero out action token embeddings | |
| all_actions_mask = all_actions_mask.unsqueeze(-1) # (B, seq_len, 1) | |
| input_embeddings = input_embeddings * ~all_actions_mask | |
| # Build multimodal embeddings and attention mask | |
| multimodal_embeddings, multimodal_attention_mask = self._build_multimodal_attention( | |
| input_embeddings, projected_patch_embeddings, attention_mask | |
| ) | |
| # Forward pass through language model | |
| language_model_output = self._run_language_model_with_layer_gate( | |
| token_gate=token_gate, | |
| input_ids=None, | |
| attention_mask=multimodal_attention_mask, | |
| position_ids=None, | |
| past_key_values=None, | |
| inputs_embeds=multimodal_embeddings, | |
| labels=None, | |
| use_cache=None, | |
| output_attentions=False, | |
| output_hidden_states=True, | |
| return_dict=True, | |
| ) | |
| # Extract hidden states for action tokens | |
| last_hidden_states = language_model_output.hidden_states[-1] # (B, seq_len, D) | |
| actions_hidden_states = last_hidden_states[ | |
| :, | |
| NUM_PATCHES + NUM_PROMPT_TOKENS : NUM_PATCHES + NUM_PROMPT_TOKENS + ACTION_DIM * NUM_ACTIONS_CHUNK, | |
| :, | |
| ] # (B, act_chunk_len, D) | |
| # Handle different prediction methods | |
| if action_head is not None: | |
| # L1 regression prediction | |
| normalized_actions = action_head.predict_action(actions_hidden_states) | |
| normalized_actions = normalized_actions.reshape(NUM_ACTIONS_CHUNK, ACTION_DIM) | |
| normalized_actions = normalized_actions.float().cpu().detach().numpy() | |
| else: | |
| # Discrete token-based prediction | |
| predicted_action_token_ids = ( | |
| language_model_output.logits[ | |
| :, | |
| NUM_PATCHES + NUM_PROMPT_TOKENS : NUM_PATCHES + NUM_PROMPT_TOKENS + ACTION_DIM * NUM_ACTIONS_CHUNK, | |
| ] | |
| .argmax(dim=2) | |
| .cpu() | |
| .numpy() | |
| ) | |
| discretized_actions = self.vocab_size - predicted_action_token_ids | |
| discretized_actions = np.clip(discretized_actions - 1, a_min=0, a_max=self.bin_centers.shape[0] - 1) | |
| normalized_actions = self.bin_centers[discretized_actions] | |
| normalized_actions = normalized_actions.reshape(NUM_ACTIONS_CHUNK, ACTION_DIM) | |
| return normalized_actions, actions_hidden_states | |
| def predict_action( | |
| self, | |
| input_ids: Optional[torch.LongTensor] = None, | |
| unnorm_key: Optional[str] = None, | |
| proprio=None, | |
| proprio_projector=None, | |
| action_head=None, | |
| noisy_action_projector=None, | |
| use_film: bool = False, | |
| text_token_pool_mask: Optional[torch.Tensor] = None, | |
| **kwargs: str, | |
| ) -> np.ndarray: | |
| """Predict actions from input sequence, with options for different prediction methods. | |
| Args: | |
| input_ids: Input token ids | |
| unnorm_key: Key for unnormalization statistics | |
| proprio: Proprioceptive features | |
| proprio_projector: Projector for proprioceptive features | |
| action_head: Optional head for L1 regression or diffusion-based prediction | |
| noisy_action_projector: Projector for noisy actions in diffusion-based prediction | |
| use_film: Whether to use FiLM conditioning | |
| **kwargs: Additional arguments including pixel_values and attention_mask | |
| Returns: | |
| Tuple of (unnormalized_actions, action_hidden_states) | |
| """ | |
| self._last_token_gate_budget_loss = None | |
| self._last_token_gate_linear_mean_penalty = None | |
| self._last_token_gate_regularization_loss = None | |
| self._debug_last_token_gate = None | |
| self._debug_last_token_gate_mean = None | |
| self._debug_last_token_gate_budget_loss = None | |
| self._debug_last_token_gate_linear_mean_penalty = None | |
| self._debug_last_text_pool_weights = None | |
| self._debug_last_text_pool_mask = None | |
| self._debug_last_contrastive_scores = None | |
| # If the special empty token ('') does not already appear after the colon (':') token in the prompt | |
| # (after "OUT:" or "ASSISTANT:"), insert it to match the inputs seen at training time | |
| if text_token_pool_mask is not None: | |
| text_token_pool_mask = text_token_pool_mask.to(device=input_ids.device, dtype=torch.bool) | |
| if not torch.all(input_ids[:, -1] == 29871): | |
| input_ids = torch.cat( | |
| (input_ids, torch.unsqueeze(torch.Tensor([29871]).long(), dim=0).to(input_ids.device)), dim=1 | |
| ) | |
| if text_token_pool_mask is not None: | |
| text_token_pool_mask = torch.cat( | |
| [ | |
| text_token_pool_mask, | |
| torch.zeros((text_token_pool_mask.shape[0], 1), device=input_ids.device, dtype=torch.bool), | |
| ], | |
| dim=1, | |
| ) | |
| pixel_values = kwargs["pixel_values"] | |
| attention_mask = kwargs["attention_mask"] | |
| # Create fake labels tensor (needed for action mask) | |
| labels = input_ids.clone() | |
| labels[:] = IGNORE_INDEX | |
| # Get number of tokens in prompt (excluding the start token) | |
| NUM_PROMPT_TOKENS = input_ids.shape[-1] - 1 # Subtract action tokens and stop token | |
| # Prepare inputs by adding necessary tokens | |
| input_ids, attention_mask = self._prepare_input_for_action_prediction(input_ids, attention_mask) | |
| if text_token_pool_mask is not None: | |
| mask_extension_len = input_ids.shape[-1] - text_token_pool_mask.shape[-1] | |
| if mask_extension_len > 0: | |
| text_token_pool_mask = torch.cat( | |
| [ | |
| text_token_pool_mask, | |
| torch.zeros( | |
| (text_token_pool_mask.shape[0], mask_extension_len), | |
| device=input_ids.device, | |
| dtype=torch.bool, | |
| ), | |
| ], | |
| dim=1, | |
| ) | |
| # Update labels tensor for action mask computation later | |
| labels = self._prepare_labels_for_action_prediction(labels, input_ids) | |
| # Get input embeddings and action masks | |
| input_embeddings = self.get_input_embeddings()(input_ids) | |
| all_actions_mask = self._process_action_masks(labels) | |
| # Extract language embeddings | |
| language_embeddings = input_embeddings[~all_actions_mask].reshape( | |
| input_embeddings.shape[0], -1, input_embeddings.shape[2] | |
| ) | |
| language_attention_mask = None | |
| if attention_mask is not None: | |
| language_attention_mask = attention_mask[~all_actions_mask].reshape(input_embeddings.shape[0], -1) | |
| language_text_pool_mask = None | |
| if text_token_pool_mask is not None: | |
| if text_token_pool_mask.shape != all_actions_mask.shape: | |
| logger.warning( | |
| "text_token_pool_mask shape does not match eval input/action mask shape; " | |
| f"got {text_token_pool_mask.shape} vs {all_actions_mask.shape}. Ignoring pool mask." | |
| ) | |
| else: | |
| language_text_pool_mask = text_token_pool_mask[~all_actions_mask].reshape(input_embeddings.shape[0], -1) | |
| # Process vision features | |
| projected_patch_embeddings = self._process_vision_features(pixel_values, language_embeddings, use_film) | |
| # Apply the gate only to projected visual tokens before proprio is appended. | |
| token_gate = None | |
| if self.use_text_token_gate: | |
| projected_patch_embeddings, token_gate, _ = self._apply_text_token_gate( | |
| projected_patch_embeddings, language_embeddings, language_attention_mask, language_text_pool_mask | |
| ) | |
| # Add proprioceptive features if provided | |
| use_proprio = proprio_projector is not None and proprio is not None | |
| if use_proprio: | |
| proprio = torch.Tensor(proprio).to(projected_patch_embeddings.device, dtype=projected_patch_embeddings.dtype) | |
| projected_patch_embeddings = self._process_proprio_features( | |
| projected_patch_embeddings, proprio, proprio_projector | |
| ) | |
| # Use diffusion if provided, otherwise use regression or discrete prediction | |
| use_diffusion = noisy_action_projector is not None and hasattr(action_head, "noise_scheduler") | |
| # Calculate number of patches (including proprio token and/or diffusion timestep embedding if present) | |
| NUM_PATCHES = self.vision_backbone.get_num_patches() * self.vision_backbone.get_num_images_in_input() | |
| if use_proprio: | |
| NUM_PATCHES += 1 | |
| if use_diffusion: | |
| NUM_PATCHES += 1 | |
| if use_diffusion: | |
| # Sample random noise with shape equal to output action, used as the starting state for reverse diffusion | |
| noise = torch.randn( | |
| size=(1, NUM_ACTIONS_CHUNK, ACTION_DIM), device=input_embeddings.device, dtype=input_embeddings.dtype | |
| ) | |
| # Run diffusion-based prediction | |
| normalized_actions, actions_hidden_states = self._run_diffusion_prediction( | |
| input_embeddings, | |
| all_actions_mask, | |
| noise, | |
| action_head, | |
| projected_patch_embeddings, | |
| labels, | |
| attention_mask, | |
| NUM_PATCHES, | |
| NUM_PROMPT_TOKENS, | |
| noisy_action_projector, | |
| token_gate=token_gate, | |
| ) | |
| else: | |
| # Run regression or discrete token-based prediction | |
| normalized_actions, actions_hidden_states = self._regression_or_discrete_prediction( | |
| input_embeddings, | |
| all_actions_mask, | |
| projected_patch_embeddings, | |
| attention_mask, | |
| labels, | |
| NUM_PATCHES, | |
| NUM_PROMPT_TOKENS, | |
| action_head, | |
| token_gate=token_gate, | |
| ) | |
| # Unnormalize predicted actions | |
| actions = self._unnormalize_actions(normalized_actions, unnorm_key) | |
| return actions, actions_hidden_states | |
| def _check_unnorm_key(norm_stats: Dict[str, Dict[str, Any]], unnorm_key: Optional[str]) -> str: | |
| """Validate and resolve the unnormalization key for action statistics""" | |
| if unnorm_key is None: | |
| assert len(norm_stats) == 1, ( | |
| f"Your model was trained on more than one dataset, " | |
| f"please pass a `unnorm_key` from the following options to choose the statistics " | |
| f"used for un-normalizing actions: {norm_stats.keys()}" | |
| ) | |
| unnorm_key = next(iter(norm_stats.keys())) | |
| assert unnorm_key in norm_stats, ( | |
| f"The `unnorm_key` you chose is not in the set of available dataset statistics, " | |
| f"please choose from: {norm_stats.keys()}" | |
| ) | |
| return unnorm_key | |
| def get_action_dim(self, unnorm_key: Optional[str] = None) -> int: | |
| """Get the dimensionality of the policy's action space.""" | |
| unnorm_key = self._check_unnorm_key(self.norm_stats, unnorm_key) | |
| return len(self.norm_stats[unnorm_key]["action"]["min"]) | |
| def get_action_stats(self, unnorm_key: Optional[str] = None) -> Dict[str, Any]: | |
| """Get all the logged statistics for the given dataset.""" | |
| unnorm_key = self._check_unnorm_key(self.norm_stats, unnorm_key) | |
| return self.norm_stats[unnorm_key]["action"] | |