| 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, |
| ) |
|
|