Add Wall-X serving and Turtle2 TCP WebSocket bridge
Pre-commit / pre-commit (push) Canceled after 0s
Pre-commit / pre-commit (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,47 @@
|
||||
"""Checkpoint-specific multimodal prompt contract, without loading model weights."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from wall_x._vendor.harrix.serving._wallx_infer.infer_config import InferConfig
|
||||
from wall_x._vendor.harrix.serving._wallx_infer.model_wrapper import WallxModelWrapper
|
||||
from wall_x.trainer.trainer_utils import load_wallx_processors
|
||||
|
||||
|
||||
def test_three_square_cameras_preserve_all_image_tokens():
|
||||
"""Truncating an image placeholder must never reach the VLA as valid input."""
|
||||
if "WALLX_TEST_CHECKPOINT" not in os.environ:
|
||||
pytest.skip("set WALLX_TEST_CHECKPOINT to run checkpoint-specific input tests")
|
||||
checkpoint = Path(os.environ["WALLX_TEST_CHECKPOINT"])
|
||||
config = InferConfig(
|
||||
checkpoint_path=str(checkpoint),
|
||||
train_config_path=str(checkpoint / "config.yml"),
|
||||
model_device="cpu", action_horizon=32,
|
||||
)
|
||||
processor = load_wallx_processors(config.train_config, device="cpu")["processor"]
|
||||
wrapper = WallxModelWrapper.__new__(WallxModelWrapper)
|
||||
wrapper.config = config
|
||||
wrapper.model = SimpleNamespace(processor=processor)
|
||||
wrapper.cam_names = config.cam_names
|
||||
wrapper.norm_key = "pick_paper_lerobot_v21_3cam"
|
||||
wrapper.logger = logging.getLogger(__name__)
|
||||
wrapper.role_start_symbol = "<|im_start|>"
|
||||
wrapper.role_end_symbol = "<|im_end|>"
|
||||
wrapper.vision_start_symbol = "<|vision_start|>"
|
||||
wrapper.vision_end_symbol = "<|vision_end|>"
|
||||
wrapper.image_pad_symbol = "<|image_pad|>"
|
||||
wrapper.propri_symbol = "<|propri|>"
|
||||
wrapper.action_symbol = "<|action|>"
|
||||
prefix, postfix = wrapper.get_text_for_action("pick up the paper towel")
|
||||
image = np.full((448, 448, 3), 127, dtype=np.uint8)
|
||||
observation = [{name: image for name in wrapper.cam_names}]
|
||||
inputs = wrapper.construct_model_input(observation, prefix, postfix)
|
||||
merge = processor.image_processor.merge_size ** 2
|
||||
expected = int((inputs["image_grid_thw"].prod(dim=1) // merge).sum())
|
||||
actual = int((inputs["input_ids"] == processor.tokenizer.convert_tokens_to_ids("<|image_pad|>")).sum())
|
||||
assert actual == expected == 768
|
||||
Reference in New Issue
Block a user