| """ |
| model.py |
| ======== |
| Frozen VLA 模型定义(Qwen3-VL 版本) |
| |
| 架构: |
| - Qwen3-VL-2B-Instruct: 冻结,自带 Vision Encoder + LLM |
| - AutoProcessor: 处理图像+文本(由模型自带) |
| - MLP Projector: 可训练,将 LLM 输出映射到动作空间 |
| - Action Head: 可训练,输出 7D EEF delta |
| |
| 注意: |
| - 不需要额外加载 SigLIP,Qwen3-VL 自带视觉编码器 |
| - 所有 Qwen3-VL 参数的 requires_grad = False |
| - 只优化 MLP Projector + Action Head 的参数 |
| - 输出保持 BridgeData 原始格式 (EEF delta),不需要改 |
| - 严格遵循 Qwen3-VL 官方用法: apply_chat_template + AutoProcessor 完整流程 |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| from typing import Optional, Dict, Any, List, Union |
| import numpy as np |
| from PIL import Image |
|
|
| try: |
| from transformers import AutoModelForImageTextToText, AutoProcessor |
| HAS_QWEN3VL = True |
| except ImportError: |
| HAS_QWEN3VL = False |
| try: |
| |
| from transformers import Qwen3VLForConditionalGeneration, AutoProcessor |
| print("[WARN] AutoModelForImageTextToText not found, using Qwen3VLForConditionalGeneration") |
| except ImportError: |
| try: |
| from transformers import Qwen2VLForConditionalGeneration, AutoProcessor |
| print("[WARN] Qwen3VL not found, using Qwen2VL") |
| except ImportError: |
| raise ImportError( |
| "Qwen2VL or Qwen3VL not found. Please upgrade transformers:\n" |
| " pip install transformers>=4.51.0" |
| ) |
|
|
|
|
| class FrozenVLA(nn.Module): |
| """ |
| Frozen Vision-Language-Action 模型(Qwen3-VL 基座) |
| |
| Args: |
| llm_name: HuggingFace model name (Qwen3-VL or Qwen2-VL) |
| mlp_hidden_dim: MLP projector 隐藏层维度 |
| mlp_depth: MLP projector 层数 |
| action_dim: 输出动作维度 (BridgeData=7) |
| action_mean: action 标准化均值 |
| action_std: action 标准化标准差 |
| """ |
| |
| def __init__( |
| self, |
| llm_name: str = "Qwen/Qwen3-VL-2B-Instruct", |
| mlp_hidden_dim: int = 512, |
| mlp_depth: int = 2, |
| action_dim: int = 7, |
| action_mean: Optional[torch.Tensor] = None, |
| action_std: Optional[torch.Tensor] = None, |
| attn_implementation: str = "sdpa", |
| ): |
| super().__init__() |
| |
| self.action_dim = action_dim |
| self.action_mean = action_mean |
| self.action_std = action_std |
| |
| |
| |
| |
| print(f"[Model] Loading Qwen3-VL: {llm_name}") |
| |
| if HAS_QWEN3VL: |
| ModelClass = AutoModelForImageTextToText |
| else: |
| try: |
| ModelClass = Qwen3VLForConditionalGeneration |
| except NameError: |
| ModelClass = Qwen2VLForConditionalGeneration |
| |
| self.qwen3vl = ModelClass.from_pretrained( |
| llm_name, |
| torch_dtype=torch.bfloat16, |
| trust_remote_code=True, |
| attn_implementation=attn_implementation, |
| ) |
| self.processor = AutoProcessor.from_pretrained( |
| llm_name, |
| trust_remote_code=True, |
| ) |
| |
| |
| for param in self.qwen3vl.parameters(): |
| param.requires_grad = False |
| self.qwen3vl.eval() |
| |
| |
| if hasattr(self.qwen3vl.config, "hidden_size"): |
| hidden_size = self.qwen3vl.config.hidden_size |
| elif hasattr(self.qwen3vl.config, "text_config") and hasattr(self.qwen3vl.config.text_config, "hidden_size"): |
| hidden_size = self.qwen3vl.config.text_config.hidden_size |
| else: |
| |
| hidden_size = self.qwen3vl.model.language_model.config.hidden_size |
| print(f"[Model] Qwen3-VL hidden size: {hidden_size}") |
| |
| |
| |
| |
| print(f"[Model] Building MLP Projector: {hidden_size} -> {mlp_hidden_dim}x{mlp_depth} -> {action_dim}") |
| |
| mlp_layers = [] |
| in_dim = hidden_size |
| for _ in range(mlp_depth): |
| mlp_layers.extend([ |
| nn.Linear(in_dim, mlp_hidden_dim), |
| nn.GELU(), |
| nn.Dropout(0.1), |
| ]) |
| in_dim = mlp_hidden_dim |
| self.mlp_projector = nn.Sequential(*mlp_layers) |
| |
| |
| |
| |
| self.action_head = nn.Linear(mlp_hidden_dim, action_dim) |
| |
| |
| self.translation_scale = nn.Parameter(torch.ones(1)) |
| self.rotation_scale = nn.Parameter(torch.ones(1)) |
| |
| print(f"[Model] Total trainable params: {self.count_trainable_params():,}") |
| print(f"[Model] Total frozen params: {self.count_frozen_params():,}") |
| |
| def count_trainable_params(self): |
| return sum(p.numel() for p in self.parameters() if p.requires_grad) |
| |
| def count_frozen_params(self): |
| return sum(p.numel() for p in self.parameters() if not p.requires_grad) |
| |
| def forward(self, images: Union[List[Image.Image], torch.Tensor], instructions: List[str]): |
| """ |
| 前向传播 |
| |
| Args: |
| images: list of PIL Images (推荐) 或 Tensor (B, 3, H, W)(兼容旧路径,会警告) |
| instructions: list of B strings |
| |
| Returns: |
| action: (B, action_dim) EEF delta |
| """ |
| device = next(self.qwen3vl.parameters()).device |
|
|
| |
| if isinstance(images, torch.Tensor): |
| print("[WARN] Received torch.Tensor images in forward(). " |
| "Please set dataset use_processor=True to return PIL Images directly.") |
| images = self._tensor_to_pil(images) |
|
|
| |
| assert isinstance(images, list) and len(images) == len(instructions), \ |
| f"images must be a list of PIL Images with length {len(instructions)}, got {type(images)}" |
|
|
| |
| |
| |
| prompts = [ |
| f"<|vision_start|><|image_pad|><|vision_end|>\nWhat action should the robot take to: {instr}?" |
| for instr in instructions |
| ] |
|
|
| |
| |
| |
| inputs = self.processor( |
| text=prompts, |
| images=images, |
| return_tensors="pt", |
| padding=True, |
| ) |
| inputs = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in inputs.items()} |
|
|
| |
| |
| |
| |
| |
| with torch.no_grad(): |
| outputs = self.qwen3vl.model( |
| **inputs, |
| return_dict=True, |
| ) |
| |
| |
| hidden_states = outputs.last_hidden_state |
|
|
| |
| |
| |
| action_features = hidden_states[:, -1, :] |
|
|
| |
| |
| |
| projected = self.mlp_projector(action_features) |
| action = self.action_head(projected) |
|
|
| |
| action[:, :3] *= self.translation_scale |
| action[:, 3:6] *= self.rotation_scale |
| |
|
|
| return action |
| |
| def _tensor_to_pil(self, images_tensor): |
| """ |
| 【兼容旧路径】将 torchvision normalize 后的 tensor 转回 PIL Image。 |
| 仅在 dataset 返回 tensor 时应急使用。建议 dataset 直接返回 PIL。 |
| images_tensor: (B, 3, H, W) |
| """ |
| mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1).to(images_tensor.device) |
| std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1).to(images_tensor.device) |
| |
| images = images_tensor * std + mean |
| images = torch.clamp(images, 0, 1) |
| images = (images * 255).byte().cpu().numpy() |
| |
| pil_images = [] |
| for img in images: |
| img = img.transpose(1, 2, 0) |
| pil_images.append(Image.fromarray(img)) |
| return pil_images |
| |
| def get_trainable_params(self): |
| """返回所有可训练参数""" |
| return [p for p in self.parameters() if p.requires_grad] |
| |
| def save_trainable_dict(self) -> dict: |
| """返回可训练参数的字典(用于 checkpoint)""" |
| state = { |
| "epoch": 0, |
| "mlp_projector": self.mlp_projector.state_dict(), |
| "action_head": self.action_head.state_dict(), |
| "translation_scale": self.translation_scale.data, |
| "rotation_scale": self.rotation_scale.data, |
| } |
| if self.action_mean is not None: |
| state["action_mean"] = self.action_mean |
| state["action_std"] = self.action_std |
| return state |
| |
| def load_trainable_dict(self, state: dict): |
| """从字典加载可训练参数""" |
| self.mlp_projector.load_state_dict(state["mlp_projector"]) |
| self.action_head.load_state_dict(state["action_head"]) |
| self.translation_scale.data = state["translation_scale"] |
| self.rotation_scale.data = state["rotation_scale"] |
| if "action_mean" in state: |
| self.action_mean = state["action_mean"] |
| self.action_std = state["action_std"] |
| |
| def save_trainable(self, path: str): |
| """只保存可训练参数 (checkpoint 很小)""" |
| state = self.save_trainable_dict() |
| torch.save(state, path) |
| print(f"[Model] Saved trainable params to {path}") |
| |
| def load_trainable(self, path: str): |
| """加载可训练参数""" |
| state = torch.load(path, map_location="cpu") |
| self.load_trainable_dict(state) |
| print(f"[Model] Loaded trainable params from {path}") |
|
|
|
|
| def build_model(config: Dict[str, Any]) -> FrozenVLA: |
| """从配置构建模型""" |
| model = FrozenVLA( |
| llm_name=config["model"]["llm"], |
| mlp_hidden_dim=config["model"]["mlp_hidden_dim"], |
| mlp_depth=config["model"]["mlp_depth"], |
| action_dim=config["model"]["action_dim"], |
| ) |
| return model |
|
|
|
|
| if __name__ == "__main__": |
| |
| print("=== Model Test (Qwen3-VL) ===") |
| |
| |
| model = FrozenVLA( |
| llm_name="Qwen/Qwen2-VL-2B-Instruct", |
| mlp_hidden_dim=256, |
| mlp_depth=2, |
| action_dim=7, |
| ) |
| |
| |
| dummy_images = [Image.new("RGB", (224, 224), color=(128, 128, 128)) for _ in range(2)] |
| dummy_text = ["put the red block on the blue plate", "move the spoon to the bowl"] |
| |
| with torch.no_grad(): |
| output = model(dummy_images, dummy_text) |
| |
| print(f"Output shape: {output.shape}") |
| print(f"Output: {output}") |
| print(f"Trainable params: {model.count_trainable_params():,}") |
| print("=== Test OK ===") |
|
|