Update Wall-X to 1.1.0 (#104)

This commit is contained in:
Starrick Liu
2026-06-15 11:40:00 +08:00
committed by GitHub
parent e23a586846
commit 72834e7de5
200 changed files with 33916 additions and 16771 deletions
+1
View File
@@ -0,0 +1 @@
"""QAct (Qwen-VLA) model family."""
+2
View File
@@ -0,0 +1,2 @@
from .configuration_qwen2_5_vl import Qwen2_5_VLConfig
from .modeling_qwen2_5_vl_act import Qwen2_5_VLMoEForAction, Qwen2_5_VLMoEModel
+56
View File
@@ -0,0 +1,56 @@
"""Qwen2.5 VLA adapter - variant-specific overrides on top of VLAdapter."""
from wall_x.model.registry import register_model
from wall_x.trainer.adapters.vla_model_adapter import VLAdapter
@register_model("qwen2_5")
class Qwen2_5Adapter(VLAdapter):
MODEL_TYPE = "qwen2_5"
@classmethod
def model_class(cls):
from wall_x.model.qact.qwen2_5 import Qwen2_5_VLMoEForAction
return Qwen2_5_VLMoEForAction
@classmethod
def config_class(cls):
from wall_x.model.qact.qwen2_5 import Qwen2_5_VLConfig
return Qwen2_5_VLConfig
@classmethod
def inference_model_class(cls):
return cls.model_class()
def get_transformer_layer_cls(self):
layer_classes = set()
try:
from transformers.models.qwen2_vl.modeling_qwen2_vl import (
Qwen2VLDecoderLayer,
)
layer_classes.add(Qwen2VLDecoderLayer)
except ImportError:
pass
try:
from wall_x.model.qact.qwen2_5.modeling_qwen2_5_vl import (
Qwen2_5_VLDecoderLayer,
)
layer_classes.add(Qwen2_5_VLDecoderLayer)
except ImportError:
pass
return layer_classes if layer_classes else None
@staticmethod
def log_attention_implementation(logger, model):
logger.info(
f"*** model attention implementation: "
f"{model.model._attn_implementation} ***"
)
logger.info(
f"*** model.visual attention implementation: "
f"{model.visual.config._attn_implementation} ***"
)
@@ -0,0 +1,357 @@
# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!
# This file was automatically generated from src/transformers/models/qwen2_5_vl/modular_qwen2_5_vl.py.
# Do NOT edit this file manually as any edits will be overwritten by the generation of
# the file from the modular. If any change should be done, please apply the change to the
# modular_qwen2_5_vl.py file directly. One of our CI enforces this.
# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!
# coding=utf-8
# Copyright 2025 The Qwen Team and The HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import logging
from transformers.configuration_utils import PretrainedConfig
from transformers.modeling_rope_utils import rope_config_validation
logger = logging.getLogger(__name__)
class Qwen2_5_VLVisionConfig(PretrainedConfig):
model_type = "qwen2_5_vl"
base_config_key = "vision_config"
def __init__(
self,
depth=32,
hidden_size=3584,
hidden_act="silu",
intermediate_size=3420,
num_heads=16,
in_channels=3,
patch_size=14,
spatial_merge_size=2,
temporal_patch_size=2,
tokens_per_second=4,
window_size=112,
out_hidden_size=3584,
fullatt_block_indexes=[7, 15, 23, 31],
initializer_range=0.02,
_attn_implementation="flash_attention_2",
attn_deterministic=False,
**kwargs,
):
super().__init__(**kwargs)
self.depth = depth
self.hidden_size = hidden_size
self.hidden_act = hidden_act
self.intermediate_size = intermediate_size
self.num_heads = num_heads
self.in_channels = in_channels
self.patch_size = patch_size
self.spatial_merge_size = spatial_merge_size
self.temporal_patch_size = temporal_patch_size
self.tokens_per_second = tokens_per_second
self.window_size = window_size
self.fullatt_block_indexes = fullatt_block_indexes
self.out_hidden_size = out_hidden_size
self.initializer_range = initializer_range
self._attn_implementation = _attn_implementation
self.attn_deterministic = attn_deterministic
class Qwen2_5_VLConfig(PretrainedConfig):
r"""
This is the configuration class to store the configuration of a [`Qwen2_5_VLModel`]. It is used to instantiate a
Qwen2-VL model according to the specified arguments, defining the model architecture. Instantiating a configuration
with the defaults will yield a similar configuration to that of
Qwen2-VL-7B-Instruct [Qwen/Qwen2-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2-VL-7B-Instruct).
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
documentation from [`PretrainedConfig`] for more information.
Args:
vocab_size (`int`, *optional*, defaults to 152064):
Vocabulary size of the Qwen2_5_VL model. Defines the number of different tokens that can be represented by the
`inputs_ids` passed when calling [`Qwen2_5_VLModel`]
hidden_size (`int`, *optional*, defaults to 8192):
Dimension of the hidden representations.
intermediate_size (`int`, *optional*, defaults to 29568):
Dimension of the MLP representations.
num_hidden_layers (`int`, *optional*, defaults to 80):
Number of hidden layers in the Transformer encoder.
num_attention_heads (`int`, *optional*, defaults to 64):
Number of attention heads for each attention layer in the Transformer encoder.
num_key_value_heads (`int`, *optional*, defaults to 8):
This is the number of key_value heads that should be used to implement Grouped Query Attention. If
`num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if
`num_key_value_heads=1` the model will use Multi Query Attention (MQA) otherwise GQA is used. When
converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed
by meanpooling all the original heads within that group. For more details checkout [this
paper](https://arxiv.org/pdf/2305.13245.pdf). If it is not specified, will default to `32`.
hidden_act (`str` or `function`, *optional*, defaults to `"silu"`):
The non-linear activation function (function or string) in the decoder.
max_position_embeddings (`int`, *optional*, defaults to 32768):
The maximum sequence length that this model might ever be used with.
initializer_range (`float`, *optional*, defaults to 0.02):
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
rms_norm_eps (`float`, *optional*, defaults to 1e-05):
The epsilon used by the rms normalization layers.
use_cache (`bool`, *optional*, defaults to `True`):
Whether or not the model should return the last key/values attentions (not used by all models). Only
relevant if `config.is_decoder=True`.
tie_word_embeddings (`bool`, *optional*, defaults to `False`):
Whether the model's input and output word embeddings should be tied.
rope_theta (`float`, *optional*, defaults to 1000000.0):
The base period of the RoPE embeddings.
use_sliding_window (`bool`, *optional*, defaults to `False`):
Whether to use sliding window attention.
sliding_window (`int`, *optional*, defaults to 4096):
Sliding window attention (SWA) window size. If not specified, will default to `4096`.
max_window_layers (`int`, *optional*, defaults to 80):
The number of layers that use SWA (Sliding Window Attention). The bottom layers use SWA while the top use full attention.
attention_dropout (`float`, *optional*, defaults to 0.0):
The dropout ratio for the attention probabilities.
vision_config (`Dict`, *optional*):
The config for the visual encoder initialization.
rope_scaling (`Dict`, *optional*):
Dictionary containing the scaling configuration for the RoPE embeddings. NOTE: if you apply new rope type
and you expect the model to work on longer `max_position_embeddings`, we recommend you to update this value
accordingly.
Expected contents:
`rope_type` (`str`):
The sub-variant of RoPE to use. Can be one of ['default', 'linear', 'dynamic', 'yarn', 'longrope',
'llama3'], with 'default' being the original RoPE implementation.
`factor` (`float`, *optional*):
Used with all rope types except 'default'. The scaling factor to apply to the RoPE embeddings. In
most scaling types, a `factor` of x will enable the model to handle sequences of length x *
original maximum pre-trained length.
`original_max_position_embeddings` (`int`, *optional*):
Used with 'dynamic', 'longrope' and 'llama3'. The original max position embeddings used during
pretraining.
`attention_factor` (`float`, *optional*):
Used with 'yarn' and 'longrope'. The scaling factor to be applied on the attention
computation. If unspecified, it defaults to value recommended by the implementation, using the
`factor` field to infer the suggested value.
`beta_fast` (`float`, *optional*):
Only used with 'yarn'. Parameter to set the boundary for extrapolation (only) in the linear
ramp function. If unspecified, it defaults to 32.
`beta_slow` (`float`, *optional*):
Only used with 'yarn'. Parameter to set the boundary for interpolation (only) in the linear
ramp function. If unspecified, it defaults to 1.
`short_factor` (`List[float]`, *optional*):
Only used with 'longrope'. The scaling factor to be applied to short contexts (<
`original_max_position_embeddings`). Must be a list of numbers with the same length as the hidden
size divided by the number of attention heads divided by 2
`long_factor` (`List[float]`, *optional*):
Only used with 'longrope'. The scaling factor to be applied to long contexts (<
`original_max_position_embeddings`). Must be a list of numbers with the same length as the hidden
size divided by the number of attention heads divided by 2
`low_freq_factor` (`float`, *optional*):
Only used with 'llama3'. Scaling factor applied to low frequency components of the RoPE
`high_freq_factor` (`float`, *optional*):
Only used with 'llama3'. Scaling factor applied to high frequency components of the RoPE
```python
>>> from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2_5_VLConfig
>>> # Initializing a Qwen2_5_VL style configuration
>>> configuration = Qwen2_5_VLConfig()
>>> # Initializing a model from the Qwen2-VL-7B style configuration
>>> model = Qwen2_5_VLForConditionalGeneration(configuration)
>>> # Accessing the model configuration
>>> configuration = model.config
```"""
model_type = "qwen2_5_vl"
sub_configs = {"vision_config": Qwen2_5_VLVisionConfig}
keys_to_ignore_at_inference = ["past_key_values"]
# Default tensor parallel plan for base model `Qwen2_5_VL`
base_model_tp_plan = {
"layers.*.self_attn.q_proj": "colwise",
"layers.*.self_attn.k_proj": "colwise",
"layers.*.self_attn.v_proj": "colwise",
"layers.*.self_attn.o_proj": "rowwise",
"layers.*.mlp.gate_proj": "colwise",
"layers.*.mlp.up_proj": "colwise",
"layers.*.mlp.down_proj": "rowwise",
}
base_model_pp_plan = {
"embed_tokens": (["input_ids"], ["inputs_embeds"]),
"layers": (["hidden_states", "attention_mask"], ["hidden_states"]),
"norm": (["hidden_states"], ["hidden_states"]),
}
def __init__(
self,
vocab_size=152064,
hidden_size=8192,
action_hidden_size=2048,
state_hidden_size=2048,
intermediate_size=29568,
num_hidden_layers=80,
num_attention_heads=64,
num_key_value_heads=8,
hidden_act="silu",
max_position_embeddings=32768,
initializer_range=0.02,
rms_norm_eps=1e-05,
use_cache=True,
tie_word_embeddings=False,
rope_theta=1000000.0,
use_sliding_window=False,
sliding_window=4096,
max_window_layers=80,
attention_dropout=0.0,
vision_config=None,
rope_scaling=None,
num_experts=4,
experts=None,
dof_config=None,
noise_scheduler=None,
dim_inputs=(1536, 1536),
attention_moe=False,
mlp_moe=False,
norm_moe=False,
mot_opt=False,
ar_loss_weight=1.0,
use_state_string_representation=False,
use_adarms=False,
proj_with_mask=True,
adarms_cond_dim=None,
action_horizon_flow=32,
causal_action_attention_mask=False,
use_flow_action_expert=True,
use_x_pred=False,
attn_deterministic=False,
use_x_loss=False,
**kwargs,
):
# Compatibility with newer transformers versions (5.x):
# - Older versions: super() sets self.pad_token_id; override it with the saved value so kwargs are preserved
# - Newer versions: super() does not set self.pad_token_id; assign it afterward
_pad_token_id = kwargs.pop("pad_token_id", None)
self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings
self.hidden_size = hidden_size
self.action_hidden_size = action_hidden_size
self.state_hidden_size = state_hidden_size
self.intermediate_size = intermediate_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.use_sliding_window = use_sliding_window
self.sliding_window = sliding_window
self.max_window_layers = max_window_layers
# for backward compatibility
if num_key_value_heads is None:
num_key_value_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.hidden_act = hidden_act
self.initializer_range = initializer_range
self.rms_norm_eps = rms_norm_eps
self.use_cache = use_cache
self.rope_theta = rope_theta
self.attention_dropout = attention_dropout
self.rope_scaling = rope_scaling
self.num_experts = num_experts
self.experts = experts
self.dof_config = dof_config
self.noise_scheduler = noise_scheduler
self.dim_inputs = tuple(dim_inputs)
self.attention_moe = attention_moe
self.mlp_moe = mlp_moe
self.norm_moe = norm_moe
self.mot_opt = mot_opt
self.ar_loss_weight = ar_loss_weight
self.use_state_string_representation = use_state_string_representation
self.use_adarms = use_adarms
self.adarms_cond_dim = adarms_cond_dim
self.proj_with_mask = proj_with_mask
self.use_flow_action_expert = use_flow_action_expert
self.action_horizon_flow = action_horizon_flow
self.causal_action_attention_mask = causal_action_attention_mask
self.use_flow_action_expert = use_flow_action_expert
self.use_x_pred = use_x_pred
self.attn_deterministic = attn_deterministic
self.use_x_loss = use_x_loss
# Validate the correctness of rotary position embeddings parameters
# BC: if there is a 'type' field, move it to 'rope_type'.
# and change type from 'mrope' to 'default' because `mrope` does defeault RoPE calculations
# one can set it to "linear"/"dynamic" etc. to have scaled RoPE
# TODO: @raushan update config in the hub
if self.rope_scaling is not None and "type" in self.rope_scaling:
if self.rope_scaling["type"] == "mrope":
self.rope_scaling["type"] = "default"
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
rope_config_validation(self, ignore_keys={"mrope_section"})
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
# Assign pad_token_id as described above, overriding older super() values or filling newer missing attributes
self.pad_token_id = _pad_token_id
# move vision config initialization after super init to avoid recursively set in latest transformers version
# TODO: make it better
if isinstance(vision_config, dict):
self.vision_config = self.sub_configs["vision_config"](**vision_config)
elif vision_config is None:
self.vision_config = self.sub_configs["vision_config"]()
def update_model_config(self, train_config):
"""Update model configuration from training config.
This method updates the model configuration with training-specific
settings such as action horizon, DOF config, attention implementation, etc.
Args:
train_config: dict containing training configuration parameters.
"""
self.use_state_string_representation = train_config["data"].get(
"use_state_string_representation", False
)
self.ar_loss_weight = train_config.get("ar_loss_weight", 1.0)
self.dof_config = train_config["dof_config"]
self.agent_pos_config = train_config["agent_pos_config"]
self.action_horizon_flow = train_config["data"].get("action_horizon_flow", 32)
if train_config.get("_attn_implementation", None) is not None:
self._attn_implementation = train_config["_attn_implementation"]
if train_config.get("attn_deterministic", None) is not None:
self.attn_deterministic = train_config["attn_deterministic"]
self.vision_config.attn_deterministic = train_config["attn_deterministic"]
logger.debug("Attention is using deterministic kernel for this run")
else:
self.attn_deterministic = True
self.vision_config.attn_deterministic = True
if train_config.get("noise_scheduler", None) is not None:
self.noise_scheduler = train_config["noise_scheduler"]
__all__ = ["Qwen2_5_VLConfig"]
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+962
View File
@@ -0,0 +1,962 @@
"""Action tokenizer mixin for loading tokenizers and action mappings."""
from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
from transformers import AutoProcessor
from wall_x.utils.constant import is_action_dataset_name
# Delay imports so missing optional packages do not fail at import time
try:
from spatial_tokenizer.spatial_tokenizer import SpatialActionTokenizer
except ImportError:
SpatialActionTokenizer = None
class ActionTokenizerMixin(ABC):
"""Base class for action tokenizers"""
def __init__(self):
self._tokenizer = None
self.action_normalizer = None
self._tokenizer_type: str = ""
self._action_mapper_cache: Optional[Dict] = None # Cached action_mapper
self.dllm = False
self.input_placeholder_flag = False
@property
def tokenizer_type(self) -> str:
"""Return the tokenizer type identifier"""
return self._tokenizer_type
@property
def tokenizer(self):
"""Return the underlying tokenizer instance"""
return self._tokenizer
@property
def action_mapper(self) -> Optional[Dict]:
"""Return the cached action_mapper"""
return self._action_mapper_cache
@abstractmethod
def load_tokenizer(self, config: dict, normalizer, device: str = "cpu") -> Any:
"""
Load the tokenizer instance
Args:
config: Configuration dictionary
normalizer: Normalizer
device: Device, usually "cpu" for training and "cuda" for inference
Returns:
tokenizer instance
"""
pass
@abstractmethod
def get_val_tokenizer(self, config: dict) -> Any:
"""
Get the tokenizer used for validation/inference
Args:
config: Configuration dictionary
Returns:
validation tokenizer instance
"""
pass
@abstractmethod
def get_special_tokens(self) -> List[str]:
"""
Return special tokens to add to the vocabulary
Returns:
list of token strings
"""
pass
def get_all_special_tokens(self) -> Tuple[List[str], Optional[List[str]]]:
"""
Return all special tokens and keep <|action_token_0|> before AR action tokens
Returns:
(new_tokens, special_tokens) tuple
"""
tokens, special_tokens = self.get_special_tokens()
if "<|action_token_0|>" not in tokens:
tokens.insert(0, "<|action_token_0|>")
return tokens, special_tokens
@abstractmethod
def build_action_mapper(self, processor) -> Optional[Dict]:
"""
Build action_mapper
Args:
processor: HuggingFace processor,used to convert tokens to IDs
Returns:
action_mapper dictionary; format depends on tokenizer type
"""
pass
@abstractmethod
def get_action_token_list(self, processor) -> List[int]:
"""
Get the action token ID list
Args:
processor: HuggingFace processor
Returns:
action token ID list
"""
pass
@abstractmethod
def decode_action(
self,
output_ids: torch.Tensor,
action_mapper: Dict,
action_horizon: int,
action_dim: int,
device: torch.device,
proprioception: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_id: Optional[int] = None,
state: Optional[torch.Tensor] = None,
) -> Tuple[Optional[Union[np.ndarray, torch.Tensor]], bool]:
"""
Unified decoding interface
Args:
output_ids: model output token IDs [1, seq_len]
action_mapper: action_mapper dictionary
action_horizon: action horizon
action_dim: action dimension
device: Device
proprioception: proprioception, normalized when required by a tokenizer
dof_mask: DOF mask when required by a tokenizer
robot_type_id: robot type ID when required by a tokenizer
state: state when required by fast/spatial tokenizers
Returns:
(predict_action, decode_success)
- predict_action: decoded action [T, action_dim] or None
- decode_success: whether decoding succeeded
"""
pass
@abstractmethod
def compute_accuracy(
self,
logits: torch.Tensor,
labels: torch.Tensor,
action_mapper: Dict,
action_token_id_set: Dict,
) -> Dict[str, torch.Tensor]:
"""
Compute accuracy metrics
Args:
logits: model output logits
labels: Labels
action_mapper: action_mapper dictionary
action_token_id_set: set of action token IDs
Returns:
accuracy metric dictionary, such as {"action_accuracy": tensor, ...}
"""
pass
@abstractmethod
def get_accuracy_keys(self) -> List[str]:
"""
Return accuracy metric keys for logging
Returns:
list of key strings
"""
pass
@property
@abstractmethod
def vocab_size(self) -> int:
"""Return the vocabulary size"""
pass
@property
@abstractmethod
def uses_dof_mask_for_unnorm(self) -> bool:
"""
Whether unnormalization requires dof_mask
Returns:
True: dof_mask is required by fast/spatial tokenizers
False: dof_mask is not required and full dimensions are returned
"""
pass
@property
@abstractmethod
def needs_action_crop(self) -> bool:
"""
Whether actions must be clipped before encoding
Returns:
True: clip by chunk_size and dof_mask for fast/spatial tokenizers
False: no clipping; the encoder handles the full sequence internally
"""
pass
@abstractmethod
def encode_to_tokens(
self,
actions: torch.Tensor,
obs_state: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_ids: Optional[List[int]] = None,
is_train: bool = True,
) -> List[List[str]]:
"""
Encode actions into token strings for training data processing
Args:
actions: Normalized actions
- fast/spatial: List[Tensor], each [T, D] (clipped)
- v3.1 delta: Tensor [B, T, D] (full sequence)
obs_state: observation state when required by a tokenizer[B, obs_horizon, D]
dof_mask: DOF mask[B, T, D]
robot_type_ids: robot type ID list when required by a tokenizer
Returns:
List[List[str]]: token string list for each sample
"""
pass
def init_inference(self, robot_type: Optional[str] = None) -> None:
"""
Inference initialization hook; subclasses may override
Args:
robot_type: robot type name
"""
pass
def prepare_action_for_ar_encoding(
self,
ar_actionchunk: torch.Tensor,
dataset_names: List[str],
agent_pos: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Prepare action data before AR encoding; subclasses may override
SpatialVLA needs waypoint selection; other tokenizers return normalized_action directly
Args:
ar_actionchunk: [B, T, D] actions before normalization
dataset_names: [str] dataset name list
agent_pos: [B, T, D] agent positions required by SpatialVLA
Returns:
prepared action data
"""
if self.action_normalizer is not None:
ar_actionchunk = self.action_normalizer.normalize_data(
ar_actionchunk, dataset_names
)
return ar_actionchunk
def get_robot_type_ids(
self,
uids: List[Optional[str]],
dataset_names: List[str],
) -> Optional[List[int]]:
"""
Get robot_type_ids; subclasses may override
Only the v3.1 delta tokenizer needs this; other tokenizers return None
Args:
uids: UID list
dataset_names: dataset name list
Returns:
robot_type_id list or None
"""
return None
def modify_inputs_for_dllm(
self,
inputs: Dict[str, torch.Tensor],
processor,
sample_time: torch.Tensor,
dataset_names: List[str],
):
return NotImplementedError
class FastTokenizerMixin(ActionTokenizerMixin):
"""Fast tokenizer implementation"""
def __init__(self):
super().__init__()
self._tokenizer_type = "fast"
def load_tokenizer(self, config: dict, normalizer, device: str = "cpu") -> Any:
"""Load the fast tokenizer"""
self._tokenizer = AutoProcessor.from_pretrained(
config["action_tokenizer_path"], trust_remote_code=True
)
self.action_normalizer = normalizer
return self._tokenizer
def get_val_tokenizer(self, config: dict) -> Any:
return self._tokenizer
def get_special_tokens(self) -> List[str]:
"""Return special tokens for the fast tokenizer"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
# tokens = ["<|ar_action|>", "<|ar_pad|>"] #temporary compatibility with existing checkpoints
tokens = []
for i in range(self._tokenizer.vocab_size):
tokens.append(f"<|action_token_{i}|>")
return tokens, None
def build_action_mapper(self, processor) -> Dict[int, int]:
"""
Build the cached action_mapper for the fast tokenizer
Returns:
Dict[token_id, action_idx]
"""
# Return cached value if available
if self._action_mapper_cache is not None:
return self._action_mapper_cache
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
action_mapper = {}
for i in range(self._tokenizer.vocab_size):
token = f"<|action_token_{i}|>"
token_id = processor.tokenizer.convert_tokens_to_ids(token)
action_mapper[token_id] = i
self._action_mapper_cache = action_mapper
return action_mapper
def get_action_token_list(self, processor) -> List[int]:
"""Get the action token ID list"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
action_token_list = []
for i in range(self._tokenizer.vocab_size):
token_id = processor.tokenizer.convert_tokens_to_ids(
f"<|action_token_{i}|>"
)
action_token_list.append(token_id)
return action_token_list
def decode_action(
self,
output_ids: torch.Tensor,
action_mapper: Dict,
action_horizon: int,
action_dim: int,
device: torch.device,
proprioception: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_id: Optional[int] = None,
state: Optional[torch.Tensor] = None,
) -> Tuple[Optional[Union[np.ndarray, torch.Tensor]], bool]:
"""Fast tokenizer decoding"""
action_id = []
for token_id_i in output_ids[0]:
if token_id_i.item() in action_mapper:
action_id.append(action_mapper[token_id_i.item()])
if len(action_id) == 0:
return np.zeros((action_horizon, action_dim)), False
predict_action = self._tokenizer.decode(
[action_id], time_horizon=action_horizon, action_dim=action_dim
)
# Check whether decoding succeeded
decode_success = False
if isinstance(predict_action, np.ndarray):
decode_success = np.sum(predict_action) != 0
elif isinstance(predict_action, torch.Tensor):
decode_success = predict_action.sum().item() != 0
return predict_action, decode_success
def compute_accuracy(
self,
logits: torch.Tensor,
labels: torch.Tensor,
action_mapper: Dict,
action_token_id_set: Dict,
) -> Dict[str, torch.Tensor]:
"""Compute fast tokenizer accuracy"""
result = {}
if len(action_token_id_set.get("action_token_list", [])) > 0:
shift_logits = logits[..., :-1, :].contiguous()
action_preds = shift_logits.argmax(dim=-1)
shift_labels = labels[..., 1:].contiguous()
action_mask = shift_labels > action_token_id_set["action_token_list"][0]
correct_preds = (action_preds == shift_labels) & action_mask
action_accuracy = correct_preds.sum().float() / action_mask.sum().float()
result["action_accuracy"] = action_accuracy
return result
def get_accuracy_keys(self) -> List[str]:
"""Return fast tokenizer accuracy keys"""
return ["action_accuracy"]
@property
def vocab_size(self) -> int:
if self._tokenizer is None:
return 0
return self._tokenizer.vocab_size
@property
def uses_dof_mask_for_unnorm(self) -> bool:
"""fast tokenizer requires dof_mask"""
return True
@property
def needs_action_crop(self) -> bool:
return True
@property
def inference_ar_steps_for_dllm(self) -> int:
return self.max_length
def encode_to_tokens(
self,
actions: List,
obs_state: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_ids: Optional[List[int]] = None,
is_train: bool = True,
) -> List[List[str]]:
"""
Fast tokenizer encoding
Args:
actions: List[Tensor/ndarray], each [T, D] (clipped)
"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
all_action_tokens = []
for i in range(len(actions)):
action = actions[i]
if isinstance(action, torch.Tensor):
action = action.cpu().numpy()
token_id = self._tokenizer(action)
action_tokens = [f"<|action_token_{idx}|>" for idx in token_id[0]]
all_action_tokens.append(action_tokens)
return all_action_tokens
def modify_inputs_for_dllm(
self,
inputs: Dict[str, torch.Tensor],
processor,
sample_time: torch.Tensor,
dataset_names: List[str],
):
# Untested
input_ids = inputs["input_ids"]
labels = inputs["labels"]
prefix_length = inputs["prefix_length"]
bs, seqlen = input_ids.shape
ar_token_length = self.max_length
ar_step_num = ar_token_length
device = input_ids.device
dtype = input_ids.dtype
placeholder_ids = torch.tensor(
processor.placeholder_seq, device=device, dtype=dtype
)
if not torch.is_tensor(sample_time):
sample_time = torch.tensor(sample_time, device=device, dtype=torch.float32)
else:
sample_time = sample_time.to(device=device, dtype=torch.float32)
sample_time = sample_time.clamp(0.0, 1.0)
noisy_steps_per_sample = (
torch.ceil((1.0 - sample_time) * (ar_step_num + 1)).long() - 1
)
noisy_steps_per_sample = noisy_steps_per_sample.clamp(min=0, max=ar_step_num)
start = prefix_length - ar_token_length - 2
end = prefix_length - 2
# Prepare the noise sequence for each sample
noise_seqs = placeholder_ids.repeat(bs, 1)
ar_len = end - start
if noise_seqs.size(1) != ar_len:
# These should usually match; defensively truncate to ar_len
noise_seqs = noise_seqs[:, :ar_len]
rand = torch.rand(bs, ar_len, device=device)
perm = rand.argsort(dim=-1)
ranks = perm.argsort(dim=-1)
noisy_mask = ranks < noisy_steps_per_sample.view(-1, 1)
ar_input = input_ids[:, start:end]
ar_labels = labels[:, start + 1 : end + 1]
ar_input[noisy_mask] = noise_seqs[noisy_mask]
ar_labels[~noisy_mask] = -100
inputs["input_ids"] = input_ids
inputs["labels"] = labels
return inputs
def update_placeholder_mask(self, processor, prefix_length, input_ids):
# Untested
ar_action_mask = torch.zeros_like(input_ids)
inc = torch.arange(
1,
self.max_length + 1,
)
inc = inc.unsqueeze(0).expand(ar_action_mask.size(0), -1) # [bs, ar_len]
ar_action_mask[
:,
prefix_length - self.max_length - 2 : prefix_length - 2,
] = inc
ar_action_mask[:, prefix_length - 2 : prefix_length] = -1 # eos
return {"ar_action_mask": ar_action_mask}
def get_placeholder_for_dllm(self):
placeholder_seq = ["<|ar_action|>"] * self.max_length
return placeholder_seq
class SpatialVLATokenizerMixin(ActionTokenizerMixin):
"""SpatialVLA tokenizer implementation"""
def __init__(self):
super().__init__()
self._tokenizer_type = "spatialvla"
def load_tokenizer(self, config: dict, normalizer, device: str = "cpu") -> Any:
"""Load the SpatialVLA tokenizer"""
if SpatialActionTokenizer is None:
raise ImportError(
"SpatialActionTokenizer is not installed. "
"Please install spatial_tokenizer package."
)
self._tokenizer = SpatialActionTokenizer(
normalizer=normalizer,
augment_ratio=config.get("augment_ratio", 0.0),
max_waypoints=config.get("max_waypoints", 5),
with_gripper=config.get("with_gripper", True),
single_arm=config.get("single_arm", False),
)
self._val_tokenizer = None
self.config = config
self.dllm = config.get("dllm", False)
self.input_placeholder_flag = config.get("input_placeholder_flag", False)
self.action_normalizer = normalizer
self.with_gripper = config.get("with_gripper", True)
return self._tokenizer
def get_placeholder_for_dllm(self):
if self.with_gripper:
placeholder_seq = [
"<|left_xyz|>",
"<|left_rpy|>",
"<|left_gripper|>",
"<|right_xyz|>",
"<|right_rpy|>",
"<|right_gripper|>",
]
else:
placeholder_seq = [
"<|left_xyz|>",
"<|left_rpy|>",
"<|right_xyz|>",
"<|right_rpy|>",
]
if self._tokenizer.single_arm:
placeholder_seq = placeholder_seq[len(placeholder_seq) // 2 :]
placeholder_seq = placeholder_seq * self._tokenizer.max_waypoints
return placeholder_seq
def get_val_tokenizer(self, config: dict) -> Any:
if self._val_tokenizer:
return self._val_tokenizer
self._val_tokenizer = SpatialActionTokenizer(
normalizer=self.action_normalizer,
augment_ratio=0,
max_waypoints=self.config.get("max_waypoints", 5),
with_gripper=self.config.get("with_gripper", True),
single_arm=self.config.get("single_arm", False),
)
return self._val_tokenizer
def get_special_tokens(self) -> List[str]:
"""Return special tokens for the SpatialVLA tokenizer"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
tokens = [
"<|step|>",
"<|left|>",
"<|right|>",
"<|move|>",
] # only for compatibility with existing checkpoints
if self.input_placeholder_flag:
special_tokens = [
"<|left_xyz|>",
"<|left_rpy|>",
"<|left_gripper|>",
"<|right_xyz|>",
"<|right_rpy|>",
"<|right_gripper|>",
]
tokens += special_tokens
if not self._tokenizer.with_gripper:
indices = [0, 1, 3, 4]
special_tokens = [special_tokens[i] for i in indices]
if self._tokenizer.single_arm:
special_tokens = special_tokens[len(special_tokens) // 2 :]
for i in range(self._tokenizer.vocab_size):
tokens.append(f"<|action_token_{i}|>")
return tokens, special_tokens
def build_action_mapper(self, processor) -> Dict[int, int]:
"""
Build the cached action_mapper for the SpatialVLA tokenizer
Returns:
Dict[token_id, action_idx]
"""
# Return cached value if available
if self._action_mapper_cache is not None:
return self._action_mapper_cache
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
action_mapper = {}
for i in range(self._tokenizer.vocab_size):
token = f"<|action_token_{i}|>"
token_id = processor.tokenizer.convert_tokens_to_ids(token)
action_mapper[token_id] = i
self._action_mapper_cache = action_mapper
return action_mapper
def get_action_token_list(self, processor) -> List[int]:
"""Get the action token ID list"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
action_token_list = []
for i in range(self._tokenizer.vocab_size):
token_id = processor.tokenizer.convert_tokens_to_ids(
f"<|action_token_{i}|>"
)
action_token_list.append(token_id)
return action_token_list
def decode_action(
self,
output_ids: torch.Tensor,
action_mapper: Dict,
action_horizon: int,
action_dim: int,
device: torch.device,
proprioception: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_id: Optional[int] = None,
state: Optional[torch.Tensor] = None,
) -> Tuple[Optional[Union[np.ndarray, torch.Tensor]], bool]:
"""SpatialVLA tokenizer decoding"""
action_id = []
for token_id_i in output_ids[0]:
if token_id_i.item() in action_mapper:
action_id.append(action_mapper[token_id_i.item()])
if len(action_id) == 0:
return np.zeros((action_horizon, action_dim)), False
if state is not None:
predict_action = self._tokenizer.decode(
[action_id],
state=state[0, 0, :action_dim],
time_horizon=action_horizon,
action_dim=action_dim,
)
else:
predict_action = self._tokenizer.decode(
[action_id], time_horizon=action_horizon, action_dim=action_dim
)
# Check whether decoding succeeded
decode_success = False
if isinstance(predict_action, np.ndarray):
decode_success = np.sum(predict_action) != 0
elif isinstance(predict_action, torch.Tensor):
decode_success = predict_action.sum().item() != 0
return predict_action, decode_success
def compute_accuracy(
self,
logits: torch.Tensor,
labels: torch.Tensor,
action_mapper: Dict,
action_token_id_set: Dict,
) -> Dict[str, torch.Tensor]:
"""Compute SpatialVLA tokenizer accuracy, same as fast"""
result = {}
if len(action_token_id_set.get("action_token_list", [])) > 0:
shift_logits = logits[..., :-1, :].contiguous()
action_preds = shift_logits.argmax(dim=-1)
shift_labels = labels[..., 1:].contiguous()
action_mask = shift_labels > action_token_id_set["action_token_list"][0]
correct_preds = (action_preds == shift_labels) & action_mask
action_accuracy = correct_preds.sum().float() / action_mask.sum().float()
result["action_accuracy"] = action_accuracy
return result
def get_accuracy_keys(self) -> List[str]:
"""Return SpatialVLA tokenizer accuracy keys"""
return ["action_accuracy"]
@property
def vocab_size(self) -> int:
if self._tokenizer is None:
return 0
return self._tokenizer.vocab_size
@property
def uses_dof_mask_for_unnorm(self) -> bool:
"""SpatialVLA tokenizer requires dof_mask"""
return True
@property
def needs_action_crop(self) -> bool:
return False
@property
def inference_ar_steps_for_dllm(self) -> int:
return self._tokenizer.max_waypoints
def prepare_action_for_ar_encoding(
self,
ar_actionchunk: torch.Tensor,
dataset_names: List[str],
agent_pos: Optional[torch.Tensor] = None,
) -> List[torch.Tensor]:
step = ar_actionchunk.shape[1] // self._tokenizer.max_waypoints
indices = np.arange(0, ar_actionchunk.shape[1], step)[
: self._tokenizer.max_waypoints
]
ar_action = ar_actionchunk[:, indices, :]
ar_action = self.action_normalizer.normalize_data(ar_action, dataset_names)
return ar_action
def encode_to_tokens(
self,
actions: List,
obs_state: Optional[torch.Tensor] = None,
dof_mask: Optional[torch.Tensor] = None,
robot_type_ids: Optional[List[int]] = None,
is_train: bool = True,
) -> List[List[str]]:
"""
SpatialVLA tokenizer encoding
Args:
actions: List[Tensor/ndarray], each [T, D] (clipped)
"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer not loaded. Call load_tokenizer first.")
if isinstance(actions, torch.Tensor):
# Convert torch tensors to numpy arrays
actions = actions.cpu().numpy()
tokenizer = self._tokenizer if is_train else self.get_val_tokenizer({})
token_ids_group = tokenizer.batch_encode(actions)
all_action_tokens = []
for i in range(len(token_ids_group)):
token_ids = np.array(token_ids_group[i]).reshape(-1)
action_token = [f"<|action_token_{i}|>" for i in token_ids]
all_action_tokens.append(action_token)
return all_action_tokens
def modify_inputs_for_dllm(
self,
inputs: Dict[str, torch.Tensor],
processor,
sample_time: torch.Tensor,
dataset_names: List[str],
):
"""
Prepare AR DLLM inputs by encoding sample_time into input_ids and labels
Noise injection strategy:
- Split [0, 1] into N equal parts (N = ar_step_num)
- For each sample, compute the number of noisy steps k in [1, N]
k = clamp(ceil((1 - t) * N), 1, N)
-> smaller t means more noise; values closer to 1 mean less noise, with at least one noised step
- seq is the noise sequence formed by concatenating ar_step_num placeholder_seq blocks
Each step maps to len(placeholder_seq) consecutive tokens
- Noise injection overwrites matching tokens in input_ids / labels step by step from seq
"""
input_ids = inputs["input_ids"]
labels = inputs["labels"]
prefix_length = inputs["prefix_length"] # int or tensor(1,)
bs, seqlen = input_ids.shape
# Total number of AR steps
ar_step_num = self._tokenizer.max_waypoints # N
step_token_len = len(processor.placeholder_seq) # Number of tokens per step
ar_token_length = ar_step_num * step_token_len # Total AR token length
device = input_ids.device
dtype = input_ids.dtype
placeholder_ids = torch.tensor(
processor.placeholder_seq, device=device, dtype=dtype
) # (step_token_len,)
noise_seq = placeholder_ids.repeat(ar_step_num) # (ar_token_length,)
if not torch.is_tensor(sample_time):
sample_time = torch.tensor(sample_time, device=device, dtype=torch.float32)
else:
sample_time = sample_time.to(device=device, dtype=torch.float32)
sample_time = sample_time.clamp(0.0, 1.0)
noisy_steps_per_sample = (
torch.ceil((1.0 - sample_time) * (ar_step_num + 1)).long() - 1
)
noisy_steps_per_sample = noisy_steps_per_sample.clamp(min=0, max=ar_step_num)
start = prefix_length - ar_token_length - 2
end = prefix_length - 2
action_idx = 0 # action sample counter
for i in range(bs):
if not is_action_dataset_name(dataset_names[i]):
continue
k = noisy_steps_per_sample[action_idx].item()
action_idx += 1
perm = torch.randperm(ar_step_num, device=device)
chosen_steps = perm[:k] # (k,)
step_mask = torch.zeros(ar_step_num, dtype=torch.bool, device=device)
step_mask[chosen_steps] = True
token_mask = step_mask.repeat_interleave(
step_token_len
) # (ar_token_length,)
ignore_mask = ~token_mask
# Take a view of the current sample's AR span for mask-based replacement
cur_input_view = input_ids[i, start:end]
cur_label_view = labels[
i, start + 1 : end + 1
] # shift placeholders one position to the right
# Overwrite selected tokens with noise
cur_input_view[token_mask] = noise_seq[token_mask]
cur_label_view[ignore_mask] = -100
inputs["input_ids"] = input_ids
inputs["labels"] = labels
return inputs
def update_positional_masks_for_dllm(
self, positional_masks, inputs, processor, visible_predict_ar_ratio=1
):
if self.dllm and self.input_placeholder_flag:
mask = self.update_placeholder_mask(
processor,
inputs["prefix_length"],
inputs["input_ids"],
)
positional_masks.update(mask)
if (
np.random.rand() < visible_predict_ar_ratio
): # FIXME ar_visible is temporarily decided per batch; mixed settings are untested and need optimization.
positional_masks["ar_visible"] = True
if "ar_predict_token_positions" in positional_masks:
del positional_masks["ar_predict_token_positions"]
# positional_masks["ar_predict_token_positions"] = None
else:
positional_masks["ar_visible"] = False
return positional_masks
def update_placeholder_mask(self, processor, prefix_length, input_ids):
ar_action_mask = torch.zeros_like(input_ids)
current_index = 1
step_len = len(processor.placeholder_seq)
for b, seq in enumerate(input_ids):
for i in range(self._tokenizer.max_waypoints):
start_idx = (
prefix_length - (self._tokenizer.max_waypoints - i) * step_len - 2
)
end_idx = start_idx + step_len
ar_action_mask[b, start_idx:end_idx] = current_index
current_index += 1
ar_action_mask[b, end_idx:prefix_length] = -1 # eos
return {"ar_action_mask": ar_action_mask}
# ============================================================
# Factory function
# ============================================================
_TOKENIZER_REGISTRY: Dict[str, type] = {
"fast": FastTokenizerMixin,
"spatialvla": SpatialVLATokenizerMixin,
}
def get_action_tokenizer_mixin(tokenizer_type: str) -> ActionTokenizerMixin:
"""
Get the mixin instance for a tokenizer type
Args:
tokenizer_type: tokenizer type, supports "fast", "spatialvla"
Returns:
ActionTokenizerMixin instance
Raises:
ValueError: Unsupported tokenizer type
"""
if tokenizer_type not in _TOKENIZER_REGISTRY:
raise ValueError(
f"Unsupported action tokenizer type: {tokenizer_type}. "
f"Supported types: {list(_TOKENIZER_REGISTRY.keys())}"
)
return _TOKENIZER_REGISTRY[tokenizer_type]()