File size: 10,069 Bytes
34f3bc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations
import logging
from dataclasses import dataclass, field, replace
from typing import Any

from transformers.models.qwen3_vl import Qwen3VLProcessor

import torch

from src.utils.constants import (
    ACTION,
    OBS_IMAGE, OBS_IMAGES, OBS_STATE, OBS_STR, NUM_IMAGE_SLOTS,
)
from src.transforms.core import DataTransformFn, DataDict


@DataTransformFn.register_subclass("labvla_processor")
@dataclass
class Qwen3_VLProcessorTransformFn(DataTransformFn):
    pretrained_model_name_or_path: str = 'Qwen/Qwen3-VL-4B-Instruct'
    max_length: int = 48
    task_key: str = "task"
    padding_side: str = "right"
    padding: str = "max_length"
    truncation: bool = True

    spatial_merge_size: int = 2

    vision_start_token_id: int = 151652
    vision_end_token_id: int = 151653
    image_token_id: int = 151655

    process: Any = field(default=None, init=False, repr=False)

    def __post_init__(self):
        self.processor = Qwen3VLProcessor.from_pretrained(self.pretrained_model_name_or_path)
        self.vision_start_token_id = self.processor.vision_start_token_id
        self.vision_end_token_id = self.processor.vision_end_token_id
        self.image_token_id = self.processor.image_token_id

        # Validate spatial_merge_size against the actual processor: a silent
        # mismatch between the dataclass default and the HF checkpoint scales
        # image-token counts by (actual/expected)^2. `merge_size` is the standard
        # name on Qwen3-VL's image_processor; warn_once on older HF that lacks it.
        image_proc = getattr(self.processor, "image_processor", None)
        actual_merge_size = getattr(image_proc, "merge_size", None)
        if actual_merge_size is None:
            from src.utils.logging_utils import warn_once
            import logging as _logging
            warn_once(
                _logging.getLogger(__name__),
                ("qwen3_vl_merge_size_unavailable",
                 self.pretrained_model_name_or_path),
                "[Qwen3_VLProcessorTransformFn] processor.image_processor has "
                "no `merge_size` attribute (old HF transformers?); skipping "
                "spatial_merge_size validation. Configured value=%d.",
                self.spatial_merge_size,
            )
        elif int(actual_merge_size) != int(self.spatial_merge_size):
            raise ValueError(
                f"[Qwen3_VLProcessorTransformFn] spatial_merge_size mismatch: "
                f"configured={self.spatial_merge_size} but processor reports "
                f"merge_size={int(actual_merge_size)}. Image-token counts would "
                f"be wrong by a factor of "
                f"({int(actual_merge_size)}/{self.spatial_merge_size})^2. "
                f"Align the config or use a different checkpoint."
            )

    def _get_num_img_tokens(self, grid_thw: torch.Tensor) -> int:
        """Compute the number of image tokens per camera = prod(grid_thw) / spatial_merge_size^2."""
        return int(torch.prod(grid_thw) // self.spatial_merge_size ** 2)

    def __call__(self, data: DataDict) -> DataDict:
        input_ids = []
        attention_mask = []
        pixel_values = []
        image_grid_thw = []
        first_valid_img_inputs = None
        first_valid_num_img_token = None
        first_valid_idx = None  # slot index processed by the probe loop
        for i in range(NUM_IMAGE_SLOTS):
            k = f"{OBS_IMAGES}.image{i}"
            if data[f"{k}_mask"]:
                first_valid_img_inputs = self.processor.image_processor(
                    data[k],
                    do_rescale=False,
                )
                first_valid_num_img_token = self._get_num_img_tokens(
                    first_valid_img_inputs.image_grid_thw[-1]
                )
                first_valid_idx = i  # reuse this slot's result in the main loop
                break

        if first_valid_img_inputs is None:
            # Every schema-mapped camera fell back to an invalid zero-frame
            # (missing/corrupt video or transient read), so this sample has no
            # visual evidence. Fail loud rather than feed a zero-frame with all
            # vision attention zeroed, which would silently train on a vision-less
            # sample. SkipBadSamplesDataset (the default wrapper) skips it;
            # without that wrapper the error surfaces the data problem at source.
            raise ValueError(
                "[Qwen3_VLProcessorTransformFn] all schema-mapped cameras are "
                f"invalid for this sample (checked {NUM_IMAGE_SLOTS} image "
                "slots; none had a True *_mask). This sample carries no visual "
                "evidence — treating it as a bad sample. Fix the source data "
                "(missing/corrupt video or read failure) or rely on "
                "SkipBadSamplesDataset to drop it."
            )

        for i in range(NUM_IMAGE_SLOTS):
            k = f"{OBS_IMAGES}.image{i}"
            if data[f"{k}_mask"]:
                if i == first_valid_idx:
                    # The probe loop above already ran image_processor on this
                    # slot; reuse its result instead of decoding/resizing twice.
                    img_inputs = first_valid_img_inputs
                    num_img_token = first_valid_num_img_token
                else:
                    img_inputs = self.processor.image_processor(
                        data[k],
                        do_rescale=False,
                    )
                    num_img_token = self._get_num_img_tokens(img_inputs.image_grid_thw[-1])
                pixel_values.append(img_inputs.pixel_values)
                image_grid_thw.append(img_inputs.image_grid_thw)
                input_ids += [self.vision_start_token_id] + [self.image_token_id] * num_img_token + [self.vision_end_token_id]
                attention_mask += [1] * (num_img_token + 2)
            else:
                pixel_values.append(first_valid_img_inputs.pixel_values)
                image_grid_thw.append(first_valid_img_inputs.image_grid_thw)
                input_ids += [self.vision_start_token_id] + [self.image_token_id] * first_valid_num_img_token + [self.vision_end_token_id]
                attention_mask += [0] * (first_valid_num_img_token + 2)

        data[f"{OBS_STR}.pixel_values"] = torch.cat(pixel_values)
        data[f"{OBS_STR}.image_grid_thw"] = torch.cat(image_grid_thw)

        lang_inputs = self.processor.tokenizer(
            data[self.task_key],
            max_length=self.max_length,
            padding_side=self.padding_side,
            padding=self.padding,
            truncation=self.truncation,
        )

        input_ids += lang_inputs.input_ids
        attention_mask += lang_inputs.attention_mask
        data[f"{OBS_STR}.input_ids"] = torch.tensor(input_ids)
        data[f"{OBS_STR}.attention_mask"] = torch.tensor(attention_mask)

        return data


@DataTransformFn.register_subclass("unify_labvla_inputs")
@dataclass
class UnifyLabVLAInputsTransformFn(DataTransformFn):
    # Populated by hydrate_all() from schema.action_keys. Default ("action",)
    # matches the canonical single-action-key path (legacy single-key datasets);
    # multi-key schemas like robointer_droid need this overridden so per-sub-key
    # `_is_pad` tensors can be OR-aggregated into the unified `action_is_pad`.
    action_keys: tuple[str, ...] = ("action",)

    def __call__(self, data: DataDict) -> DataDict:
        # Aggregate per-action-key `_is_pad` → unified `action_is_pad`. The
        # adapter writes `{key}_is_pad` per action sub-key; ComposeFieldsTransform
        # merges the per-key *values* into `action` but leaves the `_is_pad`
        # tensors alone (not in its src_keys). OR across sub-keys: any sub-key
        # padded → whole frame padded. Skip when canonical `action_is_pad` is
        # already present (single-key path where the adapter wrote it directly).
        if "action_is_pad" not in data:
            pad_keys = [f"{k}_is_pad" for k in self.action_keys if f"{k}_is_pad" in data]
            if pad_keys:
                merged = data[pad_keys[0]].clone()
                for pk in pad_keys[1:]:
                    merged = merged | data[pk]
                data["action_is_pad"] = merged

        out = {
            OBS_STATE: data[OBS_STATE],
            ACTION: data[ACTION],
            f"{OBS_STR}.pixel_values": data[f"{OBS_STR}.pixel_values"],
            f"{OBS_STR}.image_grid_thw": data[f"{OBS_STR}.image_grid_thw"],
            f"{OBS_STR}.input_ids": data[f"{OBS_STR}.input_ids"],
            f"{OBS_STR}.attention_mask": data[f"{OBS_STR}.attention_mask"],
        }
        for i in range(NUM_IMAGE_SLOTS):
            out[f"{OBS_IMAGES}.image{i}"] = data[f"{OBS_IMAGES}.image{i}"]
            out[f"{OBS_IMAGES}.image{i}_mask"] = data[f"{OBS_IMAGES}.image{i}_mask"]
        # Preserve optional keys when present (FAST tokens produced pre-Pad; pad mask; task).
        for k in ("fast_action_tokens", "fast_action_mask", "action_is_pad", "task"):
            if k in data:
                out[k] = data[k]
        # Preserve per-annotation token/mask/weight tensors (prefix-matched).
        # Empty in OXE-style datasets → no keys forwarded → model sees no
        # annotations and takes the pure-MSE fast path.
        for k, v in data.items():
            if (k.startswith("annotation_tokens__")
                    or k.startswith("annotation_mask__")
                    or k.startswith("annotation_weight__")):
                out[k] = v
        return out

    def hydrate(self, ctx) -> "UnifyLabVLAInputsTransformFn":
        # Supply action_keys so UnifyLabVLAInputsTransformFn can locate
        # per-sub-key `_is_pad` tensors (OR-aggregation → action_is_pad).
        t = replace(self, action_keys=tuple(ctx.schema.action_keys))
        logging.info(
            f"Hydrated {t.__class__.__name__} with {len(t.action_keys)} "
            f"action_keys ({ctx.schema.schema_id})"
        )
        return t