mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-25 10:46:01 +00:00
fix: nanoVLM processing
This commit is contained in:
@@ -52,9 +52,14 @@ lerobot-train \
|
||||
--dataset.repo_id=<dataset_repo_id> \
|
||||
--output_dir=outputs/vf_nanovlm_probe \
|
||||
--steps=5000 \
|
||||
--batch_size=16
|
||||
--batch_size=1
|
||||
```
|
||||
|
||||
The released checkpoint's native preprocessing resizes the long image side to
|
||||
2048 and creates 512px global/split views. A 480x640 camera therefore produces
|
||||
13 vision inputs and roughly 832 image placeholders; use batch size 1 initially
|
||||
for a three-camera setup.
|
||||
|
||||
Then load the probe checkpoint and selectively fine-tune the projector/decoder
|
||||
at a lower learning rate.
|
||||
|
||||
@@ -107,5 +112,8 @@ uv run python scripts/overfit_vf_variant.py \
|
||||
--steps=500
|
||||
```
|
||||
|
||||
For `nanovlm_value_function`, start with `--num_samples=2` because all overfit
|
||||
samples are held in one batch and native image tiling is memory intensive.
|
||||
|
||||
Compare runs using held-out episode NLL/MAE, per-episode return rank correlation,
|
||||
terminal success/failure separation, and the matched-versus-shuffled image loss gap.
|
||||
|
||||
@@ -20,6 +20,11 @@ from lerobot.rewards.factory import make_reward_model, make_reward_pre_post_proc
|
||||
from lerobot.rewards.nanovlm_value_function.configuration_nanovlm_value_function import (
|
||||
NanoVLMVFConfig,
|
||||
)
|
||||
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
|
||||
NANOVLM_ATTENTION_MASK,
|
||||
NANOVLM_IMAGES,
|
||||
NANOVLM_INPUT_IDS,
|
||||
)
|
||||
from lerobot.rewards.temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
|
||||
TemporalSiglipVFConfig,
|
||||
)
|
||||
@@ -132,6 +137,20 @@ def main():
|
||||
|
||||
|
||||
def _collate_processed(samples):
|
||||
if NANOVLM_IMAGES in samples[0]:
|
||||
max_length = max(sample[NANOVLM_INPUT_IDS].shape[1] for sample in samples)
|
||||
input_ids = []
|
||||
attention_masks = []
|
||||
for sample in samples:
|
||||
padding = max_length - sample[NANOVLM_INPUT_IDS].shape[1]
|
||||
input_ids.append(torch.nn.functional.pad(sample[NANOVLM_INPUT_IDS], (padding, 0)))
|
||||
attention_masks.append(torch.nn.functional.pad(sample[NANOVLM_ATTENTION_MASK], (padding, 0)))
|
||||
return {
|
||||
NANOVLM_IMAGES: [sample[NANOVLM_IMAGES][0] for sample in samples],
|
||||
NANOVLM_INPUT_IDS: torch.cat(input_ids),
|
||||
NANOVLM_ATTENTION_MASK: torch.cat(attention_masks),
|
||||
}
|
||||
|
||||
keys = {
|
||||
*[key for key in samples[0] if key.startswith("observation.images.")],
|
||||
OBS_LANGUAGE_TOKENS,
|
||||
@@ -192,8 +211,11 @@ def _image_shuffle_diagnostic(model, batch, camera_keys):
|
||||
matched_loss, _ = model(batch)
|
||||
shuffled = dict(batch)
|
||||
permutation = torch.roll(torch.arange(batch["mc_return"].shape[0], device=matched_loss.device), 1)
|
||||
if NANOVLM_IMAGES in batch:
|
||||
shuffled[NANOVLM_IMAGES] = [batch[NANOVLM_IMAGES][index] for index in permutation.cpu().tolist()]
|
||||
for key in camera_keys:
|
||||
shuffled[key] = batch[key][permutation]
|
||||
if key in batch:
|
||||
shuffled[key] = batch[key][permutation]
|
||||
shuffled_loss, _ = model(shuffled)
|
||||
print(
|
||||
f"\nvisual dependence: matched_loss={matched_loss.item():.4f} "
|
||||
|
||||
@@ -12,9 +12,9 @@ from lerobot.optim import AdamWConfig, CosineDecayWithWarmupSchedulerConfig
|
||||
class NanoVLMVFConfig(RewardModelConfig):
|
||||
nanovlm_pretrained_path: str = "lusxvr/nanoVLM-460M-8k"
|
||||
nanovlm_code_path: str = "third_party/nanoVLM"
|
||||
tokenizer_path: str = "HuggingFaceTB/SmolLM2-360M-Instruct"
|
||||
image_resolution: tuple[int, int] = (512, 512)
|
||||
tokenizer_max_length: int = 256
|
||||
# The checkpoint was aligned with an 8k context. Native image tiling can
|
||||
# require thousands of placeholder tokens for several robot cameras.
|
||||
tokenizer_max_length: int = 8192
|
||||
num_value_bins: int = 201
|
||||
value_support_min: float = -1.0
|
||||
value_support_max: float = 0.0
|
||||
|
||||
@@ -9,14 +9,14 @@ from typing import Any
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.configs.types import FeatureType
|
||||
from lerobot.rewards.distributional_value_function.common import DistributionalValueMixin
|
||||
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import ValueHead
|
||||
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
|
||||
IMAGE_MASK_SUFFIX,
|
||||
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
|
||||
NANOVLM_ATTENTION_MASK,
|
||||
NANOVLM_IMAGES,
|
||||
NANOVLM_INPUT_IDS,
|
||||
)
|
||||
from lerobot.rewards.pretrained import PreTrainedRewardModel
|
||||
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
|
||||
|
||||
from .configuration_nanovlm_value_function import NanoVLMVFConfig
|
||||
|
||||
@@ -31,9 +31,6 @@ class NanoVLMVFRewardModel(DistributionalValueMixin, PreTrainedRewardModel):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
config.validate_features()
|
||||
self.image_keys = [
|
||||
key for key, feature in config.input_features.items() if feature.type == FeatureType.VISUAL
|
||||
]
|
||||
code_path = Path(config.nanovlm_code_path)
|
||||
if not code_path.is_absolute():
|
||||
code_path = Path(__file__).resolve().parents[4] / code_path
|
||||
@@ -80,25 +77,33 @@ class NanoVLMVFRewardModel(DistributionalValueMixin, PreTrainedRewardModel):
|
||||
return self._distributional_forward(batch)
|
||||
|
||||
def _get_value_readout(self, batch: dict[str, Tensor]) -> Tensor:
|
||||
batch_size = batch[OBS_LANGUAGE_TOKENS].shape[0]
|
||||
image_tokens = []
|
||||
image_masks = []
|
||||
for key in self.image_keys:
|
||||
image = batch[key]
|
||||
mask = batch[key + IMAGE_MASK_SUFFIX].bool()
|
||||
features = self.nanovlm.MP(self.nanovlm.vision_encoder(image))
|
||||
image_tokens.append(features * mask[:, None, None].to(features.dtype))
|
||||
image_masks.append(mask[:, None].expand(batch_size, features.shape[1]))
|
||||
|
||||
text_tokens = self.nanovlm.decoder.token_embedding(batch[OBS_LANGUAGE_TOKENS])
|
||||
input_ids = batch[NANOVLM_INPUT_IDS]
|
||||
attention_mask = batch[NANOVLM_ATTENTION_MASK].bool()
|
||||
batch_size = input_ids.shape[0]
|
||||
images = self.nanovlm._process_images(batch[NANOVLM_IMAGES], input_ids.device)
|
||||
text_tokens = self.nanovlm.decoder.token_embedding(input_ids)
|
||||
if images is not None:
|
||||
image_tokens = self.nanovlm.MP(self.nanovlm.vision_encoder(images))
|
||||
placeholder_count = (input_ids == self.nanovlm.tokenizer.image_token_id).sum().item()
|
||||
image_token_count = image_tokens.shape[0] * image_tokens.shape[1]
|
||||
if placeholder_count != image_token_count:
|
||||
raise ValueError(
|
||||
"nanoVLM image placeholders do not match projected image tokens: "
|
||||
f"{placeholder_count} placeholders versus {image_token_count} tokens. "
|
||||
"The prompt may have been truncated; increase tokenizer_max_length."
|
||||
)
|
||||
text_tokens = self.nanovlm._replace_img_tokens_with_embd(
|
||||
input_ids,
|
||||
text_tokens,
|
||||
image_tokens,
|
||||
)
|
||||
query = self.value_query(torch.zeros(batch_size, 1, dtype=torch.long, device=text_tokens.device)).to(
|
||||
text_tokens.dtype
|
||||
)
|
||||
inputs = torch.cat([*image_tokens, text_tokens, query], dim=1)
|
||||
inputs = torch.cat([text_tokens, query], dim=1)
|
||||
attention_mask = torch.cat(
|
||||
[
|
||||
*image_masks,
|
||||
batch[OBS_LANGUAGE_ATTENTION_MASK].bool(),
|
||||
attention_mask,
|
||||
torch.ones(batch_size, 1, dtype=torch.bool, device=text_tokens.device),
|
||||
],
|
||||
dim=1,
|
||||
|
||||
@@ -1,30 +1,180 @@
|
||||
"""Processor for the nanoVLM value-function experiment."""
|
||||
"""Processor using nanoVLM's native image splitting and chat-token layout."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
|
||||
from lerobot.configs import FeatureType
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStepRegistry,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
batch_to_transition,
|
||||
policy_action_to_transition,
|
||||
transition_to_batch,
|
||||
)
|
||||
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
|
||||
DistributionalVFImagePreprocessorStep,
|
||||
DistributionalVFPrepareTaskPromptStep,
|
||||
)
|
||||
from lerobot.types import TransitionKey
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_nanovlm_value_function import NanoVLMVFConfig
|
||||
|
||||
NANOVLM_IMAGES = "observation.nanovlm.images"
|
||||
NANOVLM_INPUT_IDS = "observation.nanovlm.input_ids"
|
||||
NANOVLM_ATTENTION_MASK = "observation.nanovlm.attention_mask"
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="nanovlm_native_processor")
|
||||
@dataclass
|
||||
class NanoVLMNativeProcessorStep(ComplementaryDataProcessorStep):
|
||||
pretrained_path: str
|
||||
code_path: str
|
||||
image_keys: tuple[str, ...]
|
||||
max_length: int
|
||||
_tokenizer: Any = field(default=None, init=False, repr=False)
|
||||
_image_processor: Any = field(default=None, init=False, repr=False)
|
||||
_get_image_string: Any = field(default=None, init=False, repr=False)
|
||||
_mp_image_token_length: int = field(default=64, init=False, repr=False)
|
||||
|
||||
def __post_init__(self):
|
||||
code_path = Path(self.code_path)
|
||||
if not code_path.is_absolute():
|
||||
code_path = Path(__file__).resolve().parents[4] / code_path
|
||||
if str(code_path) not in sys.path:
|
||||
sys.path.insert(0, str(code_path))
|
||||
|
||||
from data.processors import get_image_processor, get_image_string, get_tokenizer
|
||||
|
||||
config_path = _resolve_checkpoint_file(self.pretrained_path, "config.json")
|
||||
config = json.loads(Path(config_path).read_text())
|
||||
if self.max_length > config["lm_max_length"]:
|
||||
raise ValueError(
|
||||
f"tokenizer_max_length={self.max_length} exceeds nanoVLM's "
|
||||
f"lm_max_length={config['lm_max_length']}"
|
||||
)
|
||||
self._tokenizer = get_tokenizer(
|
||||
config["lm_tokenizer"],
|
||||
config["vlm_extra_tokens"],
|
||||
config["lm_chat_template"],
|
||||
)
|
||||
self._image_processor = get_image_processor(
|
||||
config["max_img_size"],
|
||||
config["vit_img_size"],
|
||||
config["resize_to_max_side_len"],
|
||||
)
|
||||
self._get_image_string = get_image_string
|
||||
self._mp_image_token_length = config["mp_image_token_length"]
|
||||
|
||||
def complementary_data(self, complementary_data):
|
||||
raw_tasks = complementary_data.get("task")
|
||||
if raw_tasks is None:
|
||||
raise ValueError("Task is required for nanoVLM value processing")
|
||||
observation = self.transition[TransitionKey.OBSERVATION]
|
||||
present_image_keys = [key for key in self.image_keys if key in observation]
|
||||
if not present_image_keys:
|
||||
raise ValueError("No configured nanoVLM image key is present in the observation")
|
||||
batch_size = observation[present_image_keys[0]].shape[0]
|
||||
tasks = [raw_tasks] * batch_size if isinstance(raw_tasks, str) else list(raw_tasks)
|
||||
if len(tasks) != batch_size:
|
||||
raise ValueError(f"Received {len(tasks)} tasks for an image batch of size {batch_size}")
|
||||
|
||||
processed_batch = []
|
||||
input_rows = []
|
||||
attention_rows = []
|
||||
for batch_index in range(batch_size):
|
||||
processed_images = []
|
||||
split_counts = []
|
||||
for key in self.image_keys:
|
||||
if key not in observation:
|
||||
continue
|
||||
image = observation[key][batch_index]
|
||||
if image.ndim != 3:
|
||||
raise ValueError(f"nanoVLM expects CHW images, got {tuple(image.shape)} for {key}")
|
||||
if image.shape[0] not in (1, 3, 4) and image.shape[-1] in (1, 3, 4):
|
||||
image = image.permute(2, 0, 1)
|
||||
if image.dtype != torch.uint8:
|
||||
image = image.float()
|
||||
if image.min() < -1e-6 or image.max() > 1.0 + 1e-6:
|
||||
raise ValueError(
|
||||
f"nanoVLM expects uint8 [0,255] or float [0,1] images; "
|
||||
f"{key} has range [{image.min().item()}, {image.max().item()}]"
|
||||
)
|
||||
image = image.clamp(0, 1)
|
||||
pil_image = to_pil_image(image.cpu()).convert("RGB")
|
||||
processed, split_count = self._image_processor(pil_image)
|
||||
processed_images.append(processed)
|
||||
split_counts.append(split_count)
|
||||
|
||||
image_string = self._get_image_string(
|
||||
self._tokenizer,
|
||||
split_counts,
|
||||
self._mp_image_token_length,
|
||||
)
|
||||
prompt = self._tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": image_string + f"Task: {tasks[batch_index]}."}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
tokenized = self._tokenizer(
|
||||
prompt,
|
||||
truncation=False,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
if len(tokenized["input_ids"]) > self.max_length:
|
||||
raise ValueError(
|
||||
f"nanoVLM prompt has {len(tokenized['input_ids'])} tokens, exceeding "
|
||||
f"tokenizer_max_length={self.max_length}. The native nanoVLM collator "
|
||||
"discards over-length examples instead of truncating image placeholders."
|
||||
)
|
||||
input_rows.append(tokenized["input_ids"])
|
||||
attention_rows.append(tokenized.get("attention_mask", [1] * len(tokenized["input_ids"])))
|
||||
processed_batch.append(processed_images)
|
||||
|
||||
max_length = max(map(len, input_rows))
|
||||
for input_ids, attention_mask in zip(input_rows, attention_rows, strict=True):
|
||||
padding = max_length - len(input_ids)
|
||||
input_ids[:0] = [self._tokenizer.pad_token_id] * padding
|
||||
attention_mask[:0] = [0] * padding
|
||||
|
||||
observation = dict(observation)
|
||||
observation[NANOVLM_IMAGES] = processed_batch
|
||||
observation[NANOVLM_INPUT_IDS] = torch.tensor(input_rows, dtype=torch.long)
|
||||
observation[NANOVLM_ATTENTION_MASK] = torch.tensor(attention_rows, dtype=torch.bool)
|
||||
self.transition[TransitionKey.OBSERVATION] = observation
|
||||
return complementary_data
|
||||
|
||||
def transform_features(
|
||||
self,
|
||||
features: dict[PipelineFeatureType, dict[str, PolicyFeature]],
|
||||
):
|
||||
return features
|
||||
|
||||
def get_config(self):
|
||||
return {
|
||||
"pretrained_path": self.pretrained_path,
|
||||
"code_path": self.code_path,
|
||||
"image_keys": self.image_keys,
|
||||
"max_length": self.max_length,
|
||||
}
|
||||
|
||||
|
||||
def _resolve_checkpoint_file(repo_id_or_path: str, filename: str) -> str:
|
||||
local_path = Path(repo_id_or_path) / filename
|
||||
if local_path.exists():
|
||||
return str(local_path)
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
return hf_hub_download(repo_id=repo_id_or_path, filename=filename)
|
||||
|
||||
|
||||
def make_nanovlm_vf_pre_post_processors(
|
||||
config: NanoVLMVFConfig,
|
||||
@@ -45,16 +195,11 @@ def make_nanovlm_vf_pre_post_processors(
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DistributionalVFImagePreprocessorStep(
|
||||
image_resolution=config.image_resolution,
|
||||
NanoVLMNativeProcessorStep(
|
||||
pretrained_path=config.nanovlm_pretrained_path,
|
||||
code_path=config.nanovlm_code_path,
|
||||
image_keys=image_keys,
|
||||
),
|
||||
DistributionalVFPrepareTaskPromptStep(),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name=config.tokenizer_path,
|
||||
max_length=config.tokenizer_max_length,
|
||||
padding_side="right",
|
||||
padding="max_length",
|
||||
),
|
||||
DeviceProcessorStep(device=config.device or "cpu"),
|
||||
],
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
import json
|
||||
import sys
|
||||
from types import ModuleType, SimpleNamespace
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch import nn
|
||||
|
||||
from lerobot.configs import FeatureType, PolicyFeature
|
||||
@@ -9,7 +11,13 @@ from lerobot.rewards.factory import get_reward_model_class, make_reward_model_co
|
||||
from lerobot.rewards.nanovlm_value_function.configuration_nanovlm_value_function import (
|
||||
NanoVLMVFConfig,
|
||||
)
|
||||
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
|
||||
from lerobot.rewards.nanovlm_value_function.processor_nanovlm_value_function import (
|
||||
NANOVLM_ATTENTION_MASK,
|
||||
NANOVLM_IMAGES,
|
||||
NANOVLM_INPUT_IDS,
|
||||
NanoVLMNativeProcessorStep,
|
||||
)
|
||||
from lerobot.types import TransitionKey
|
||||
|
||||
CAMERA = "observation.images.top"
|
||||
|
||||
@@ -17,6 +25,7 @@ CAMERA = "observation.images.top"
|
||||
def test_config_and_factory_registration():
|
||||
config = make_reward_model_config("nanovlm_value_function")
|
||||
assert isinstance(config, NanoVLMVFConfig)
|
||||
assert config.tokenizer_max_length == 8192
|
||||
assert get_reward_model_class("nanovlm_value_function").__name__ == "NanoVLMVFRewardModel"
|
||||
|
||||
|
||||
@@ -56,6 +65,15 @@ def test_nanovlm_model_forward(monkeypatch):
|
||||
self.vision_encoder = FakeVision()
|
||||
self.MP = FakeProjector()
|
||||
self.decoder = FakeDecoder()
|
||||
self.tokenizer = SimpleNamespace(image_token_id=99)
|
||||
|
||||
def _process_images(self, images, device):
|
||||
return torch.cat([image for sample in images for image in sample]).to(device)
|
||||
|
||||
def _replace_img_tokens_with_embd(self, input_ids, token_embd, image_embd):
|
||||
token_embd = token_embd.clone()
|
||||
token_embd[input_ids == self.tokenizer.image_token_id] = image_embd.flatten(0, 1)
|
||||
return token_embd
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, path):
|
||||
@@ -72,13 +90,67 @@ def test_nanovlm_model_forward(monkeypatch):
|
||||
config.input_features = {CAMERA: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 16, 16))}
|
||||
model = NanoVLMVFRewardModel(config)
|
||||
batch = {
|
||||
CAMERA: torch.rand(1, 3, 16, 16),
|
||||
CAMERA + ".mask": torch.ones(1, dtype=torch.bool),
|
||||
OBS_LANGUAGE_TOKENS: torch.ones(1, 4, dtype=torch.long),
|
||||
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(1, 4, dtype=torch.bool),
|
||||
NANOVLM_IMAGES: [[torch.rand(1, 3, 16, 16)]],
|
||||
NANOVLM_INPUT_IDS: torch.tensor([[99, 99, 99, 99, 1]]),
|
||||
NANOVLM_ATTENTION_MASK: torch.ones(1, 5, dtype=torch.bool),
|
||||
"mc_return": torch.tensor([-0.5]),
|
||||
"is_terminal": torch.tensor([False]),
|
||||
}
|
||||
loss, metrics = model(batch)
|
||||
assert torch.isfinite(loss)
|
||||
assert -1.0 <= metrics["predicted_value_mean"] <= 0.0
|
||||
|
||||
|
||||
def test_native_processor_uses_checkpoint_layout_and_left_padding(monkeypatch, tmp_path):
|
||||
config = {
|
||||
"lm_tokenizer": "fake",
|
||||
"vlm_extra_tokens": {},
|
||||
"lm_chat_template": "fake",
|
||||
"lm_max_length": 8192,
|
||||
"max_img_size": 2048,
|
||||
"vit_img_size": 512,
|
||||
"resize_to_max_side_len": True,
|
||||
"mp_image_token_length": 4,
|
||||
}
|
||||
(tmp_path / "config.json").write_text(json.dumps(config))
|
||||
|
||||
class FakeTokenizer:
|
||||
pad_token_id = 0
|
||||
image_token_id = 99
|
||||
|
||||
def apply_chat_template(self, messages, tokenize, add_generation_prompt):
|
||||
assert not tokenize and add_generation_prompt
|
||||
return messages[0]["content"]
|
||||
|
||||
def __call__(self, prompt, truncation, add_special_tokens):
|
||||
assert not truncation and not add_special_tokens
|
||||
suffix = [1, 2] if "long" in prompt else [1]
|
||||
return {"input_ids": [99] * 4 + suffix, "attention_mask": [1] * (4 + len(suffix))}
|
||||
|
||||
def fake_image_processor(image):
|
||||
assert isinstance(image, Image.Image) and image.mode == "RGB"
|
||||
return torch.rand(1, 3, 512, 512), (1, 1)
|
||||
|
||||
processors = ModuleType("data.processors")
|
||||
processors.get_tokenizer = lambda *args: FakeTokenizer()
|
||||
processors.get_image_processor = lambda *args: fake_image_processor
|
||||
processors.get_image_string = lambda *args: "<image>"
|
||||
monkeypatch.setitem(sys.modules, "data.processors", processors)
|
||||
|
||||
step = NanoVLMNativeProcessorStep(
|
||||
pretrained_path=str(tmp_path),
|
||||
code_path="third_party/nanoVLM",
|
||||
image_keys=(CAMERA,),
|
||||
max_length=8192,
|
||||
)
|
||||
transition = {
|
||||
TransitionKey.OBSERVATION: {CAMERA: torch.rand(2, 3, 16, 16)},
|
||||
TransitionKey.COMPLEMENTARY_DATA: {"task": ["short", "long"]},
|
||||
}
|
||||
output = step(transition)[TransitionKey.OBSERVATION]
|
||||
|
||||
assert len(output[NANOVLM_IMAGES]) == 2
|
||||
assert output[NANOVLM_INPUT_IDS].shape == (2, 6)
|
||||
assert output[NANOVLM_INPUT_IDS][0, 0] == 0
|
||||
assert not output[NANOVLM_ATTENTION_MASK][0, 0]
|
||||
assert output[NANOVLM_ATTENTION_MASK][1].all()
|
||||
|
||||
Reference in New Issue
Block a user