tiao55's picture
Upload calibrated DINOv2 multi-backbone VLA policy for task 23
f353578 verified
Raw
History Blame Contribute Delete
4.35 kB
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,
)