Files
VLA/wall_x/model/qact/tokenizer_mixin.py

963 lines
32 KiB
Python

"""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]()