from __future__ import annotations from typing import Any import torch import torch.nn as nn from transformers import Dinov2Config, Dinov2Model from competitive_vla_model import ( DEFAULT_PHASE_DIM, CompetitivePolicyHead, normalize_images, pool_dinov2_tokens, rgb_grid_tokens, ) def _make_head(config: dict[str, Any]) -> CompetitivePolicyHead: vision_config = config["vision_config"] return CompetitivePolicyHead( vision_dim=int(vision_config["hidden_size"]), proprio_dim=int(config["proprio_dim"]), text_dim=int(config["text_dim"]), phase_dim=int(config.get("phase_dim", DEFAULT_PHASE_DIM)), num_tasks=len(config["task_to_id"]), num_difficulties=len(config["difficulty_to_id"]), hidden_dim=int(config["hidden_dim"]), history=int(config["history"]), action_chunk=int(config["action_chunk"]), ensemble_heads=int(config["ensemble_heads"]), dropout=float(config.get("dropout", 0.10)), ) class HybridCompetitiveVLAModel(nn.Module): """One frozen DINOv2 encoder shared by a stable and specialist policy head.""" def __init__(self, config: dict[str, Any]) -> None: super().__init__() self.policy_config = config self.route_configs = { str(name): dict(value) for name, value in config.get("route_configs", {}).items() } if self.route_configs: first_config = next(iter(self.route_configs.values())) self.base_config = None self.specialist_config = first_config else: self.base_config = dict(config["base_config"]) self.specialist_config = dict(config["specialist_config"]) first_config = self.specialist_config self.spatial_grid = int(first_config["spatial_grid"]) self.vision = Dinov2Model( Dinov2Config(**first_config["vision_config"]) ) if self.route_configs: self.heads = nn.ModuleDict( {name: _make_head(value) for name, value in self.route_configs.items()} ) else: assert self.base_config is not None self.base_head = _make_head(self.base_config) self.specialist_head = _make_head(self.specialist_config) def encode_images(self, images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: config = self.specialist_config image_size = int(config.get("image_size", 224)) rgb_tokens = rgb_grid_tokens( images, spatial_grid=self.spatial_grid, image_size=image_size ) pixels = normalize_images(images, image_size=image_size) target_dtype = next(self.vision.parameters()).dtype pixels = pixels.to(dtype=target_dtype) rgb_tokens = rgb_tokens.to(dtype=target_dtype) layer_average = max(1, int(config.get("vision_layer_average", 1))) output = self.vision( pixel_values=pixels, output_hidden_states=layer_average > 1, ) if layer_average > 1: if output.hidden_states is None: raise RuntimeError("DINOv2 did not return requested hidden states") hidden = torch.stack(output.hidden_states[-layer_average:], dim=0).mean(dim=0) else: hidden = output.last_hidden_state return pool_dinov2_tokens(hidden, self.spatial_grid), rgb_tokens def forward_cached( self, route: str, visual_tokens: torch.Tensor, rgb_tokens: torch.Tensor, proprio: torch.Tensor, phase: torch.Tensor, text_features: torch.Tensor, task_ids: torch.Tensor, difficulty_ids: torch.Tensor, ) -> torch.Tensor: if self.route_configs: if route not in self.heads: raise ValueError(f"unknown policy route: {route}") head = self.heads[route] elif route == "base": head = self.base_head elif route == "specialist": head = self.specialist_head else: raise ValueError(f"unknown policy route: {route}") return head.forward_cached( visual_tokens, rgb_tokens, proprio, phase, text_features, task_ids, difficulty_ids, )