task23-dino / metaclip_action_chunk_model.py
jerseyjerry's picture
Uploading model files
7d66dc6 verified
Raw
History Blame Contribute Delete
28.2 kB
"""Self-contained MetaCLIP temporal action-chunk policy definition.
This module deliberately contains no model-hub access. A complete
``clip_config`` dictionary is embedded in the policy configuration, so
constructing :class:`MetaCLIPActionChunkModel` only creates modules. Callers are
responsible for loading a local state dict afterwards.
The image, phase, and token helpers are shared by feature-cache creation,
training, and the submission adapter. Keeping those operations here prevents
subtle train/deployment preprocessing drift.
"""
from __future__ import annotations
import copy
import math
from collections.abc import Mapping
from typing import Any
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import CLIPConfig, CLIPModel
ACTION_DIM = 7
DEFAULT_TEXT_DIM = 512
DEFAULT_PHASE_DIM = 4
DEFAULT_SPATIAL_HEADS = 8
DEFAULT_DIFFICULTIES = ("low", "medium", "hard", "very_high")
TEXT_FEATURE_VERSION = "metaclip_clip_bpe_projected_l2_text_v2"
METACLIP_IMAGE_MEAN = (0.48145466, 0.4578275, 0.40821073)
METACLIP_IMAGE_STD = (0.26862954, 0.26130258, 0.27577711)
_REQUIRED_CONFIG_KEYS = (
"clip_config",
"image_size",
"spatial_grid",
"proprio_dim",
"text_dim",
"phase_dim",
"hidden_dim",
"history",
"action_chunk",
"ensemble_heads",
"dropout",
"task_to_id",
"difficulty_to_id",
)
def _require_plain_int(value: Any, name: str, *, minimum: int = 1) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < minimum:
raise ValueError(f"{name} must be an integer >= {minimum}, got {value!r}")
return int(value)
def _require_probability(value: Any, name: str) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{name} must be a number in [0, 1), got {value!r}")
value = float(value)
if not math.isfinite(value) or not 0.0 <= value < 1.0:
raise ValueError(f"{name} must be a finite number in [0, 1), got {value!r}")
return value
def _validate_id_map(value: Any, name: str) -> dict[str, int]:
if not isinstance(value, Mapping) or not value:
raise ValueError(f"{name} must be a non-empty mapping of strings to IDs")
result: dict[str, int] = {}
for key, item in value.items():
if not isinstance(key, str) or not key.strip():
raise ValueError(f"{name} contains an invalid key: {key!r}")
if isinstance(item, bool) or not isinstance(item, int) or item < 0:
raise ValueError(f"{name}[{key!r}] must be a non-negative integer")
if key in result:
raise ValueError(f"{name} contains duplicate key {key!r}")
result[key] = int(item)
expected = list(range(len(result)))
actual = sorted(result.values())
if actual != expected:
raise ValueError(
f"{name} IDs must be unique and contiguous 0..{len(result) - 1}, got {actual}"
)
return result
def validate_policy_config(config: Mapping[str, Any]) -> dict[str, Any]:
"""Validate and return an isolated copy of a policy configuration.
Deployment/training metadata outside the architectural keys is retained,
but all fields that affect tensor shapes or preprocessing are checked. The
returned object can therefore be safely stored on a model without later
mutations by the caller changing its behavior.
"""
if not isinstance(config, Mapping):
raise TypeError(f"config must be a mapping, got {type(config).__name__}")
missing = [key for key in _REQUIRED_CONFIG_KEYS if key not in config]
if missing:
raise ValueError(f"policy config is missing required keys: {missing}")
validated = copy.deepcopy(dict(config))
clip_config = validated["clip_config"]
if not isinstance(clip_config, Mapping) or not clip_config:
raise ValueError("clip_config must be a non-empty mapping")
clip_config = copy.deepcopy(dict(clip_config))
if clip_config.get("model_type") != "clip":
raise ValueError(
"clip_config.model_type must be 'clip', got "
f"{clip_config.get('model_type')!r}"
)
if (
_require_plain_int(
clip_config.get("projection_dim"), "clip_config.projection_dim"
)
!= DEFAULT_TEXT_DIM
):
raise ValueError(
f"MetaCLIP projection_dim must be {DEFAULT_TEXT_DIM}, "
f"got {clip_config.get('projection_dim')!r}"
)
vision_config = clip_config.get("vision_config")
if not isinstance(vision_config, Mapping) or not vision_config:
raise ValueError("clip_config.vision_config must be a non-empty mapping")
vision_config = copy.deepcopy(dict(vision_config))
for key in (
"hidden_size",
"intermediate_size",
"image_size",
"patch_size",
"num_hidden_layers",
"num_attention_heads",
):
if key not in vision_config:
raise ValueError(
f"clip_config.vision_config is missing required key {key!r}"
)
_require_plain_int(vision_config[key], f"clip_config.vision_config.{key}")
if vision_config.get("model_type") != "clip_vision_model":
raise ValueError(
"clip_config.vision_config.model_type must be 'clip_vision_model', got "
f"{vision_config.get('model_type')!r}"
)
if (
int(vision_config["image_size"]) != 224
or int(vision_config["patch_size"]) != 16
):
raise ValueError(
"The audited MetaCLIP B/16 contract requires image_size=224 and "
f"patch_size=16, got {vision_config['image_size']!r} and "
f"{vision_config['patch_size']!r}"
)
if int(vision_config["hidden_size"]) % int(vision_config["num_attention_heads"]):
raise ValueError(
"vision_config.hidden_size must be divisible by "
"vision_config.num_attention_heads"
)
if "num_channels" in vision_config and int(vision_config["num_channels"]) != 3:
raise ValueError(
"only three-channel MetaCLIP vision configurations are supported"
)
text_config = clip_config.get("text_config")
if not isinstance(text_config, Mapping) or not text_config:
raise ValueError("clip_config.text_config must be a non-empty mapping")
text_config = copy.deepcopy(dict(text_config))
for key in (
"hidden_size",
"intermediate_size",
"max_position_embeddings",
"num_hidden_layers",
"num_attention_heads",
"vocab_size",
):
if key not in text_config:
raise ValueError(f"clip_config.text_config is missing required key {key!r}")
_require_plain_int(text_config[key], f"clip_config.text_config.{key}")
if text_config.get("model_type") != "clip_text_model":
raise ValueError(
"clip_config.text_config.model_type must be 'clip_text_model', got "
f"{text_config.get('model_type')!r}"
)
if int(text_config["max_position_embeddings"]) != 77:
raise ValueError(
"MetaCLIP text max_position_embeddings must be 77, got "
f"{text_config['max_position_embeddings']!r}"
)
clip_config["vision_config"] = vision_config
clip_config["text_config"] = text_config
validated["clip_config"] = clip_config
image_size = _require_plain_int(validated["image_size"], "image_size")
if image_size != int(vision_config["image_size"]):
raise ValueError(
"policy image_size must match vision_config.image_size, got "
f"{image_size} and {vision_config['image_size']!r}"
)
patch_size = int(vision_config["patch_size"])
if image_size % patch_size:
raise ValueError(
f"image_size={image_size} must be divisible by MetaCLIP patch_size={patch_size}"
)
patch_side = image_size // patch_size
spatial_grid = _require_plain_int(validated["spatial_grid"], "spatial_grid")
if spatial_grid > patch_side:
raise ValueError(
f"spatial_grid={spatial_grid} cannot exceed the {patch_side}x{patch_side} "
"MetaCLIP input patch grid"
)
for key in (
"proprio_dim",
"text_dim",
"phase_dim",
"hidden_dim",
"history",
"action_chunk",
"ensemble_heads",
):
validated[key] = _require_plain_int(validated[key], key)
if validated["text_dim"] != int(clip_config["projection_dim"]):
raise ValueError(
"text_dim must equal clip_config.projection_dim, got "
f"{validated['text_dim']} and {clip_config['projection_dim']!r}"
)
if validated["phase_dim"] != DEFAULT_PHASE_DIM:
raise ValueError(
f"phase_dim must be {DEFAULT_PHASE_DIM} for phase_vector(), "
f"got {validated['phase_dim']}"
)
if validated["hidden_dim"] % DEFAULT_SPATIAL_HEADS:
raise ValueError(
f"hidden_dim must be divisible by {DEFAULT_SPATIAL_HEADS} spatial heads"
)
validated["dropout"] = _require_probability(validated["dropout"], "dropout")
validated["task_to_id"] = _validate_id_map(validated["task_to_id"], "task_to_id")
validated["difficulty_to_id"] = _validate_id_map(
validated["difficulty_to_id"], "difficulty_to_id"
)
if "action_dim" in validated and validated["action_dim"] != ACTION_DIM:
raise ValueError(
f"action_dim must be {ACTION_DIM}, got {validated['action_dim']!r}"
)
if "rgb_dim" in validated and validated["rgb_dim"] != 3:
raise ValueError(f"rgb_dim must be 3, got {validated['rgb_dim']!r}")
if (
"text_feature_version" in validated
and validated["text_feature_version"] != TEXT_FEATURE_VERSION
):
raise ValueError(
f"text_feature_version must be {TEXT_FEATURE_VERSION!r}, got "
f"{validated['text_feature_version']!r}"
)
return validated
def phase_vector(step: int, horizon: int) -> "np.ndarray":
"""Encode episode progress using the exact train/runtime four-vector."""
if isinstance(step, bool) or not isinstance(step, (int, np.integer)):
raise ValueError(f"step must be an integer, got {step!r}")
if isinstance(horizon, bool) or not isinstance(horizon, (int, np.integer)):
raise ValueError(f"horizon must be an integer, got {horizon!r}")
horizon_f = max(float(horizon), 1.0)
progress = float(np.clip(float(step) / horizon_f, 0.0, 1.0))
return np.asarray(
[
progress,
1.0 - progress,
math.sin(math.pi * progress),
math.cos(math.pi * progress),
],
dtype=np.float32,
)
def _images_to_nchw_rgb(images: torch.Tensor) -> torch.Tensor:
if not isinstance(images, torch.Tensor):
raise TypeError(f"images must be a torch.Tensor, got {type(images).__name__}")
if images.ndim != 4:
raise ValueError(
f"expected a four-dimensional image tensor, got {tuple(images.shape)}"
)
# Prefer an unambiguous channel-first interpretation, then NHWC. Normal
# robotics images are 224x224, so both layouts are unambiguous in practice.
if images.shape[1] in (1, 3, 4) and images.shape[-1] not in (1, 3, 4):
nchw = images
elif images.shape[-1] in (1, 3, 4):
nchw = images.permute(0, 3, 1, 2)
elif images.shape[1] in (1, 3, 4):
nchw = images
else:
raise ValueError(
f"cannot determine image channels for shape {tuple(images.shape)}"
)
if nchw.shape[1] == 1:
nchw = nchw.repeat(1, 3, 1, 1)
elif nchw.shape[1] == 4:
nchw = nchw[:, :3]
if nchw.shape[1] != 3:
raise ValueError(
f"expected one, three, or four image channels, got {nchw.shape[1]}"
)
return nchw
def images_to_unit_rgb(images: torch.Tensor, image_size: int = 224) -> torch.Tensor:
"""Convert NHWC/NCHW uint8-like images to resized NCHW RGB in ``[0, 1]``."""
image_size = _require_plain_int(image_size, "image_size")
rgb = _images_to_nchw_rgb(images).float()
if rgb.numel() and float(rgb.detach().amax().cpu()) > 2.0:
rgb = rgb / 255.0
if not bool(torch.isfinite(rgb).all().detach().cpu()):
raise ValueError("images contain NaN or Inf")
if tuple(rgb.shape[-2:]) != (image_size, image_size):
rgb = F.interpolate(
rgb,
size=(image_size, image_size),
mode="bicubic",
align_corners=False,
antialias=True,
)
return rgb
def normalize_images(images: torch.Tensor, image_size: int = 224) -> torch.Tensor:
"""Prepare image pixels for the MetaCLIP vision encoder."""
rgb = images_to_unit_rgb(images, image_size=image_size)
mean = rgb.new_tensor(METACLIP_IMAGE_MEAN).view(1, 3, 1, 1)
std = rgb.new_tensor(METACLIP_IMAGE_STD).view(1, 3, 1, 1)
return (rgb - mean) / std
def rgb_grid_tokens(
images: torch.Tensor, spatial_grid: int, image_size: int = 224
) -> torch.Tensor:
"""Return global RGB plus a row-major spatial grid, shaped ``[B,1+G²,3]``."""
spatial_grid = _require_plain_int(spatial_grid, "spatial_grid")
rgb = images_to_unit_rgb(images, image_size=image_size)
global_rgb = rgb.mean(dim=(-2, -1)).unsqueeze(1)
grid_rgb = F.adaptive_avg_pool2d(rgb, (spatial_grid, spatial_grid))
grid_rgb = grid_rgb.flatten(2).transpose(1, 2)
return torch.cat([global_rgb, grid_rgb], dim=1)
def pool_metaclip_tokens(
hidden_states: torch.Tensor,
spatial_grid: int,
) -> torch.Tensor:
"""Pool MetaCLIP patch tokens to ``G x G`` and retain the CLS token."""
spatial_grid = _require_plain_int(spatial_grid, "spatial_grid")
if not isinstance(hidden_states, torch.Tensor):
raise TypeError("hidden_states must be a torch.Tensor")
if hidden_states.ndim != 3 or hidden_states.shape[1] <= 1:
raise ValueError(
f"unexpected MetaCLIP output shape: {tuple(hidden_states.shape)}"
)
cls_token = hidden_states[:, :1]
patches = hidden_states[:, 1:]
side = math.isqrt(int(patches.shape[1]))
if side * side != int(patches.shape[1]):
raise ValueError(f"MetaCLIP patch count {patches.shape[1]} is not a square")
patches = patches.transpose(1, 2).reshape(
patches.shape[0], patches.shape[2], side, side
)
patches = F.adaptive_avg_pool2d(patches, (spatial_grid, spatial_grid))
patches = patches.flatten(2).transpose(1, 2)
return torch.cat([cls_token, patches], dim=1)
class MetaCLIPActionChunkHead(nn.Module):
"""Task-conditioned spatial pooling followed by a short temporal policy."""
def __init__(
self,
*,
vision_dim: int,
proprio_dim: int,
text_dim: int,
phase_dim: int,
num_tasks: int,
num_difficulties: int,
hidden_dim: int = 256,
history: int = 4,
action_chunk: int = 8,
ensemble_heads: int = 3,
dropout: float = 0.10,
) -> None:
super().__init__()
self.vision_dim = _require_plain_int(vision_dim, "vision_dim")
self.proprio_dim = _require_plain_int(proprio_dim, "proprio_dim")
self.text_dim = _require_plain_int(text_dim, "text_dim")
self.phase_dim = _require_plain_int(phase_dim, "phase_dim")
self.hidden_dim = _require_plain_int(hidden_dim, "hidden_dim")
self.history = _require_plain_int(history, "history")
self.action_chunk = _require_plain_int(action_chunk, "action_chunk")
self.ensemble_heads = _require_plain_int(ensemble_heads, "ensemble_heads")
num_tasks = _require_plain_int(num_tasks, "num_tasks")
num_difficulties = _require_plain_int(num_difficulties, "num_difficulties")
dropout = _require_probability(dropout, "dropout")
if self.hidden_dim % DEFAULT_SPATIAL_HEADS:
raise ValueError(
f"hidden_dim must be divisible by {DEFAULT_SPATIAL_HEADS} spatial heads"
)
self.vision_proj = nn.Sequential(
nn.LayerNorm(self.vision_dim), nn.Linear(self.vision_dim, self.hidden_dim)
)
self.rgb_proj = nn.Sequential(nn.Linear(3, self.hidden_dim), nn.SiLU())
self.proprio_proj = nn.Sequential(
nn.LayerNorm(self.proprio_dim + self.phase_dim),
nn.Linear(self.proprio_dim + self.phase_dim, self.hidden_dim),
nn.SiLU(),
nn.Dropout(dropout),
)
self.text_proj = nn.Sequential(
nn.LayerNorm(self.text_dim),
nn.Linear(self.text_dim, self.hidden_dim),
nn.SiLU(),
)
# The last row of each embedding is the trained unknown/fallback ID.
self.task_embedding = nn.Embedding(num_tasks + 1, self.hidden_dim)
self.difficulty_embedding = nn.Embedding(num_difficulties + 1, self.hidden_dim)
self.task_scale = nn.Parameter(torch.tensor(0.5))
self.difficulty_scale = nn.Parameter(torch.tensor(0.25))
self.condition_norm = nn.LayerNorm(self.hidden_dim)
self.spatial_attention = nn.MultiheadAttention(
embed_dim=self.hidden_dim,
num_heads=DEFAULT_SPATIAL_HEADS,
dropout=dropout,
batch_first=True,
)
self.frame_fusion = nn.Sequential(
nn.Linear(self.hidden_dim * 2, self.hidden_dim),
nn.SiLU(),
nn.LayerNorm(self.hidden_dim),
nn.Dropout(dropout),
)
self.temporal_gru = nn.GRU(
input_size=self.hidden_dim,
hidden_size=self.hidden_dim,
num_layers=2,
dropout=dropout,
batch_first=True,
)
self.output_heads = nn.ModuleList(
[
nn.Sequential(
nn.LayerNorm(self.hidden_dim),
nn.Linear(self.hidden_dim, self.hidden_dim),
nn.SiLU(),
nn.Dropout(dropout),
nn.Linear(self.hidden_dim, self.action_chunk * ACTION_DIM),
)
for _ in range(self.ensemble_heads)
]
)
def _validate_inputs(
self,
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,
) -> tuple[int, int, int]:
if visual_tokens.ndim != 4:
raise ValueError(
f"visual_tokens must have shape [B,T,V,D], got {tuple(visual_tokens.shape)}"
)
batch, timesteps, token_count, vision_dim = visual_tokens.shape
if timesteps != self.history:
raise ValueError(f"expected history={self.history}, got {timesteps}")
if vision_dim != self.vision_dim:
raise ValueError(f"expected vision_dim={self.vision_dim}, got {vision_dim}")
if rgb_tokens.shape != (batch, timesteps, token_count, 3):
raise ValueError(
"rgb_tokens must align with visual_tokens and end in RGB, got "
f"{tuple(rgb_tokens.shape)}"
)
if proprio.shape != (batch, timesteps, self.proprio_dim):
raise ValueError(
f"proprio must have shape {(batch, timesteps, self.proprio_dim)}, "
f"got {tuple(proprio.shape)}"
)
if phase.shape != (batch, timesteps, self.phase_dim):
raise ValueError(
f"phase must have shape {(batch, timesteps, self.phase_dim)}, "
f"got {tuple(phase.shape)}"
)
if text_features.shape != (batch, self.text_dim):
raise ValueError(
f"text_features must have shape {(batch, self.text_dim)}, "
f"got {tuple(text_features.shape)}"
)
for name, ids in (("task_ids", task_ids), ("difficulty_ids", difficulty_ids)):
if ids.shape != (batch,):
raise ValueError(
f"{name} must have shape {(batch,)}, got {tuple(ids.shape)}"
)
if ids.dtype not in (torch.int32, torch.int64):
raise ValueError(
f"{name} must contain integer IDs, got dtype={ids.dtype}"
)
return batch, timesteps, token_count
def forward_cached(
self,
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:
"""Return raw action logits shaped ``[ensemble, batch, chunk, 7]``."""
batch, timesteps, token_count = self._validate_inputs(
visual_tokens,
rgb_tokens,
proprio,
phase,
text_features,
task_ids,
difficulty_ids,
)
visual = self.vision_proj(visual_tokens) + self.rgb_proj(rgb_tokens)
state = self.proprio_proj(torch.cat([proprio, phase], dim=-1))
text = self.text_proj(text_features)
task = self.task_embedding(task_ids)
difficulty = self.difficulty_embedding(difficulty_ids)
condition = self.condition_norm(
state
+ text[:, None]
+ torch.tanh(self.task_scale) * task[:, None]
+ torch.tanh(self.difficulty_scale) * difficulty[:, None]
)
flat_visual = visual.reshape(batch * timesteps, token_count, self.hidden_dim)
flat_query = condition.reshape(batch * timesteps, 1, self.hidden_dim)
attended, _ = self.spatial_attention(
flat_query, flat_visual, flat_visual, need_weights=False
)
attended = attended.reshape(batch, timesteps, self.hidden_dim)
frames = self.frame_fusion(torch.cat([attended, condition], dim=-1))
temporal, _ = self.temporal_gru(frames)
final = temporal[:, -1]
outputs = [
head(final).reshape(batch, self.action_chunk, ACTION_DIM)
for head in self.output_heads
]
return torch.stack(outputs, dim=0)
def forward(
self,
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:
return self.forward_cached(
visual_tokens,
rgb_tokens,
proprio,
phase,
text_features,
task_ids,
difficulty_ids,
)
class MetaCLIPActionChunkModel(nn.Module):
"""Complete submission model containing frozen MetaCLIP and the policy head."""
def __init__(self, config: Mapping[str, Any]) -> None:
super().__init__()
self.policy_config = validate_policy_config(config)
self.spatial_grid = int(self.policy_config["spatial_grid"])
# Offline construction only: this creates a model from the embedded
# architecture. It never resolves a repository or downloads weights.
clip_config = CLIPConfig.from_dict(self.policy_config["clip_config"])
self.clip = CLIPModel(clip_config)
self.head = MetaCLIPActionChunkHead(
vision_dim=int(clip_config.vision_config.hidden_size),
proprio_dim=int(self.policy_config["proprio_dim"]),
text_dim=int(self.policy_config["text_dim"]),
phase_dim=int(self.policy_config["phase_dim"]),
num_tasks=len(self.policy_config["task_to_id"]),
num_difficulties=len(self.policy_config["difficulty_to_id"]),
hidden_dim=int(self.policy_config["hidden_dim"]),
history=int(self.policy_config["history"]),
action_chunk=int(self.policy_config["action_chunk"]),
ensemble_heads=int(self.policy_config["ensemble_heads"]),
dropout=float(self.policy_config["dropout"]),
)
def freeze_backbone(self) -> None:
"""Freeze both MetaCLIP towers and keep them in inference mode."""
self.clip.requires_grad_(False)
self.clip.eval()
def train(self, mode: bool = True) -> "MetaCLIPActionChunkModel":
# A caller may train the complete wrapper for convenience. If the
# backbone has been frozen, do not accidentally switch it back to train
# mode through nn.Module.train() recursion.
super().train(mode)
if not any(parameter.requires_grad for parameter in self.clip.parameters()):
self.clip.eval()
return self
def encode_images(self, images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Encode an image batch into aligned MetaCLIP and raw-RGB tokens."""
image_size = int(self.policy_config["image_size"])
rgb_tokens = rgb_grid_tokens(
images, spatial_grid=self.spatial_grid, image_size=image_size
)
pixels = normalize_images(images, image_size=image_size)
vision_parameter = next(self.clip.vision_model.parameters())
pixels = pixels.to(device=vision_parameter.device, dtype=vision_parameter.dtype)
rgb_tokens = rgb_tokens.to(
device=vision_parameter.device, dtype=vision_parameter.dtype
)
hidden = self.clip.vision_model(pixel_values=pixels).last_hidden_state
hidden = self.clip.vision_model.post_layernorm(hidden)
return (
pool_metaclip_tokens(hidden, self.spatial_grid),
rgb_tokens,
)
def encode_text(
self, input_ids: torch.Tensor, attention_mask: torch.Tensor
) -> torch.Tensor:
"""Return normalized projected MetaCLIP instruction embeddings."""
parameter = next(self.clip.text_model.parameters())
input_ids = input_ids.to(device=parameter.device)
attention_mask = attention_mask.to(device=parameter.device)
outputs = self.clip.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
return_dict=True,
)
features = self.clip.text_projection(outputs.pooler_output)
return F.normalize(features.float(), dim=-1).to(dtype=parameter.dtype)
def forward_cached(
self,
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:
return self.head.forward_cached(
visual_tokens,
rgb_tokens,
proprio,
phase,
text_features,
task_ids,
difficulty_ids,
)
def forward(
self,
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:
return self.forward_cached(
visual_tokens,
rgb_tokens,
proprio,
phase,
text_features,
task_ids,
difficulty_ids,
)
# Compatibility aliases make reference checkpoints/code easy to compare while
# retaining descriptive names in the new trainer.
CompetitivePolicyHead = MetaCLIPActionChunkHead
CompetitiveVLAModel = MetaCLIPActionChunkModel
__all__ = [
"ACTION_DIM",
"DEFAULT_DIFFICULTIES",
"DEFAULT_PHASE_DIM",
"DEFAULT_SPATIAL_HEADS",
"DEFAULT_TEXT_DIM",
"METACLIP_IMAGE_MEAN",
"METACLIP_IMAGE_STD",
"TEXT_FEATURE_VERSION",
"CompetitivePolicyHead",
"CompetitiveVLAModel",
"MetaCLIPActionChunkHead",
"MetaCLIPActionChunkModel",
"images_to_unit_rgb",
"normalize_images",
"phase_vector",
"pool_metaclip_tokens",
"rgb_grid_tokens",
"validate_policy_config",
]