File size: 12,803 Bytes
372c993
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
"""
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:
        # fallback: 旧版 transformers 可能没有 AutoModelForImageTextToText
        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",  # Windows/4060: sdpa; Linux/A10+: flash_attention_2
    ):
        super().__init__()
        
        self.action_dim = action_dim
        self.action_mean = action_mean
        self.action_std = action_std
        
        # ============================================================
        # 1. Qwen3-VL (冻结,自带 Vision Encoder + LLM)
        # ============================================================
        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()
        
        # 获取 hidden_size(Qwen3-VL 的 config 结构不同于 Qwen2-VL)
        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:
            # fallback: 通过 model structure
            hidden_size = self.qwen3vl.model.language_model.config.hidden_size
        print(f"[Model] Qwen3-VL hidden size: {hidden_size}")
        
        # ============================================================
        # 2. MLP Projector (可训练)
        # ============================================================
        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)
        
        # ============================================================
        # 3. Action Head (可训练)
        # ============================================================
        self.action_head = nn.Linear(mlp_hidden_dim, action_dim)
        
        # 可学习的 action 缩放因子
        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

        # 兼容旧路径:如果传入 tensor,先转回 PIL(但建议 dataset 直接返回 PIL)
        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)

        # 确保 images 是 PIL Image list
        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)}"

        # ============================================================
        # 1. 构造 prompts(含 <image> 占位符,Qwen3-VL 标准用法)
        # ============================================================
        prompts = [
            f"<|vision_start|><|image_pad|><|vision_end|>\nWhat action should the robot take to: {instr}?"
            for instr in instructions
        ]

        # ============================================================
        # 2. 使用 AutoProcessor 处理图像+文本(PAI 上验证通过的方式)
        # ============================================================
        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()}

        # ============================================================
        # 3. Qwen3-VL 前向 (冻结)
        #    用 self.qwen3vl.model (裸 Transformer) 只取最后一层 hidden state
        #    避免 output_hidden_states=True 存全部 28 层导致 OOM(省 ~7GB)
        # ============================================================
        with torch.no_grad():
            outputs = self.qwen3vl.model(
                **inputs,
                return_dict=True,
            )
            # last_hidden_state: (B, seq_len, hidden_size)
            # 取的是 final layer norm 之后的值,和 logits 输入一致
            hidden_states = outputs.last_hidden_state  # (B, seq_len, hidden_size)

        # ============================================================
        # 4. 取最后一个 token 的 hidden state 作为 action feature
        # ============================================================
        action_features = hidden_states[:, -1, :]  # (B, hidden_size)

        # ============================================================
        # 5. MLP Projector + Action Head (可训练)
        # ============================================================
        projected = self.mlp_projector(action_features)  # (B, mlp_hidden_dim)
        action = self.action_head(projected)              # (B, action_dim)

        # 应用可学习缩放
        action[:, :3] *= self.translation_scale  # 位移缩放
        action[:, 3:6] *= self.rotation_scale      # 旋转缩放
        # gripper 不缩放

        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()  # (B, 3, H, W) uint8
        
        pil_images = []
        for img in images:
            img = img.transpose(1, 2, 0)  # (H, W, 3)
            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) ===")
    
    # 使用 Qwen2-VL 2B 测试(避免下载大模型)
    model = FrozenVLA(
        llm_name="Qwen/Qwen2-VL-2B-Instruct",
        mlp_hidden_dim=256,
        mlp_depth=2,
        action_dim=7,
    )
    
    # 模拟 PIL Image 输入(官方推荐方式)
    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}")  # (2, 7)
    print(f"Output: {output}")
    print(f"Trainable params: {model.count_trainable_params():,}")
    print("=== Test OK ===")