Update Wall-X to 1.1.0 (#104)
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""QAct (Qwen-VLA) model family."""
|
||||
@@ -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
|
||||
@@ -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
@@ -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]()
|
||||
Reference in New Issue
Block a user