# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto. Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.
import torch
import torch.nn.functional as F
import torch.distributions as dists
from typing import Dict, Optional
def get_token_ids_from_config(config) -> Dict[str, int]:
"""Extract all token IDs from the configuration object.
Args:
config: Configuration object (LocateAnythingConfig or similar)
Returns:
Dictionary containing all token IDs
"""
token_ids = {}
# Get from main config
token_ids['box_start_token_id'] = getattr(config, 'box_start_token_id', 151668)
token_ids['box_end_token_id'] = getattr(config, 'box_end_token_id', 151669)
token_ids['grasp_start_token_id'] = getattr(
config, 'grasp_start_token_id', token_ids['box_start_token_id']
)
token_ids['grasp_end_token_id'] = getattr(
config, 'grasp_end_token_id', token_ids['box_end_token_id']
)
token_ids['grasp_rect_start_token_id'] = getattr(
config, 'grasp_rect_start_token_id', token_ids['box_start_token_id']
)
token_ids['grasp_rect_end_token_id'] = getattr(
config, 'grasp_rect_end_token_id', token_ids['box_end_token_id']
)
token_ids['coord_start_token_id'] = getattr(config, 'coord_start_token_id', 151677)
token_ids['coord_end_token_id'] = getattr(config, 'coord_end_token_id', 152677)
token_ids['ref_start_token_id'] = getattr(config, 'ref_start_token_id', 151672)
token_ids['ref_end_token_id'] = getattr(config, 'ref_end_token_id', 151673)
token_ids['none_token_id'] = getattr(config, 'none_token_id', 4064)
# Get from text_config
text_config = getattr(config, 'text_config', None)
if text_config is not None:
token_ids['null_token_id'] = getattr(text_config, 'null_token_id', 152678)
token_ids['im_end_token_id'] = getattr(text_config, 'eos_token_id', 151645)
token_ids['switch_token_id'] = getattr(text_config, 'switch_token_id', 152679)
token_ids['default_mask_token_id'] = getattr(text_config, 'text_mask_token_id', 151676)
else:
token_ids['null_token_id'] = 152678
token_ids['im_end_token_id'] = 151645
token_ids['switch_token_id'] = 152679
token_ids['default_mask_token_id'] = 151676
return token_ids
def top_p_logits(
logits: torch.Tensor,
top_p: float = None
) -> torch.Tensor:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
# Shift the indices to the right to keep the first token above the threshold
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
mask = torch.zeros_like(logits, dtype=torch.bool, device=logits.device)
mask = mask.scatter_(-1, sorted_indices, sorted_indices_to_remove)
logits = logits.masked_fill(mask, torch.finfo(logits.dtype).min)
return logits
def top_k_logits(
logits: torch.Tensor,
top_k: int = None
) -> torch.Tensor:
top_k = min(top_k, logits.size(-1)) # Safety check
# Remove all tokens with a probability less than the last token of the top-k
indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
logits = logits.masked_fill(indices_to_remove, torch.finfo(logits.dtype).min)
return logits
def apply_repetition_penalty(
logits: torch.Tensor,
input_ids: torch.Tensor,
repetition_penalty: float = 1.0
) -> torch.Tensor:
"""
Apply repetition penalty to logits.
Args:
logits: Shape [batch_size, seq_len, vocab_size] or [batch_size, vocab_size]
input_ids: Previously generated token ids, shape [batch_size, seq_len]
repetition_penalty: Penalty factor. > 1.0 penalizes repetition, < 1.0 encourages it.
Returns:
Modified logits with repetition penalty applied.
"""
if repetition_penalty == 1.0:
return logits
# Convert to 3D for vectorized computation
if logits.dim() == 2:
logits = logits.unsqueeze(1) # [B, 1, V]
squeeze_back = True
else:
squeeze_back = False
batch_size, seq_len, vocab_size = logits.shape
# Construct [B, V] bool mask marking tokens that have appeared in each batch
device = logits.device
token_mask = torch.zeros(batch_size, vocab_size, dtype=torch.bool, device=device)
for b in range(batch_size):
# Apply penalty only based on tokens already generated in this batch
unique_tokens = input_ids[b].unique()
# Prevent out-of-bounds: only keep IDs within vocab range
valid_tokens = unique_tokens[(unique_tokens >= 0) & (unique_tokens < vocab_size)]
if valid_tokens.numel() > 0:
token_mask[b, valid_tokens] = True
# Expand to [B, L, V] to align with logits
token_mask = token_mask.unsqueeze(1).expand(-1, seq_len, -1)
# Divide positive values by penalty, multiply negative values by penalty
positive = logits > 0
negative = ~positive
# Apply penalty only at mask positions
logits = torch.where(token_mask & positive, logits / repetition_penalty, logits)
logits = torch.where(token_mask & negative, logits * repetition_penalty, logits)
if squeeze_back:
logits = logits.squeeze(1)
return logits
def sample_tokens(
logits: torch.Tensor,
generated: torch.Tensor,
token_ids: Dict[str, int],
**generate_kwargs,
):
batch_size, seq_len, vocab_size = logits.shape
repetition_penalty = generate_kwargs.get('repetition_penalty', 1.0)
temperature = generate_kwargs.get('temperature', 0)
top_p = generate_kwargs.get('top_p', None)
top_k = generate_kwargs.get('top_k', None)
# Apply repetition penalty based on all previously generated tokens
if repetition_penalty != 1.0:
logits = apply_repetition_penalty(logits, generated, repetition_penalty)
if temperature > 0:
logits = logits / temperature
if top_p is not None and top_p < 1:
logits = top_p_logits(logits, top_p)
if top_k is not None:
logits = top_k_logits(logits, top_k)
probs = torch.softmax(logits, dim=-1)
if temperature > 0:
try:
x0 = dists.Categorical(probs=probs).sample()
confidence = torch.gather(probs, -1, x0.unsqueeze(-1)).squeeze(-1)
except Exception:
confidence, x0 = probs.max(dim=-1)
else:
confidence, x0 = probs.max(dim=-1)
if seq_len == 1:
return probs, confidence, x0, None, None
box_avg = []
structured_decode_failed = []
fallback_box = torch.zeros(1, dtype=x0.dtype, device=x0.device)
for b in range(batch_size):
geometry_type = generate_kwargs.get('geometry_type', 'bbox')
contact_frame_expected = (
geometry_type == 'contact'
and generated[b, -1].item() == token_ids['ref_end_token_id']
)
grasp_rect_frame_expected = (
geometry_type == 'grasp_rect'
and generated[b, -1].item() == token_ids['ref_end_token_id']
)
if contact_frame_expected:
decoded_box = decode_contact_pair(
logits[b],
probs[b],
token_ids,
keep_k=generate_kwargs.get('contact_keep_k', 4),
image_size=generate_kwargs.get('image_size'),
minimum_width_diagonal=generate_kwargs.get(
'contact_minimum_width_diagonal', 1e-4
),
maximum_width_diagonal=generate_kwargs.get(
'contact_maximum_width_diagonal', 1.0
),
coord_mass_threshold=generate_kwargs.get(
'contact_coord_mass_threshold', 1e-4
),
force_frame=True,
)
elif grasp_rect_frame_expected:
decoded_box = decode_grasp_rectangle(
logits[b],
probs[b],
token_ids,
keep_k=generate_kwargs.get('grasp_rect_keep_k', 4),
image_size=generate_kwargs.get('image_size'),
minimum_width_diagonal=generate_kwargs.get(
'grasp_rect_minimum_width_diagonal', 1e-4
),
gripper_depth_pixels=generate_kwargs.get(
'grasp_rect_gripper_depth_pixels', 40.0
),
coord_mass_threshold=generate_kwargs.get(
'grasp_rect_coord_mass_threshold', 1e-4
),
coord_entropy_threshold=generate_kwargs.get(
'grasp_rect_coord_entropy_threshold', 1.0
),
force_frame=True,
)
elif geometry_type in ('contact', 'grasp_rect'):
decoded_box = None
else:
decoded_box = decode_bbox_avg(
logits[b], probs[b], token_ids,
keep_k=generate_kwargs.get('keep_k_avg', 4),
generation_mode=generate_kwargs.get('generation_mode', 'hybrid'),
)
decode_failed = bool(
(contact_frame_expected or grasp_rect_frame_expected)
and decoded_box is None
)
structured_decode_failed.append(decode_failed)
if decode_failed:
box_avg.append(fallback_box)
elif decoded_box is not None:
box_avg.append(decoded_box)
else:
out_ref = decode_ref(logits[b], probs[b], token_ids)
if out_ref is not None:
box_avg.append(torch.tensor(out_ref, dtype=x0.dtype, device=x0.device))
else:
box_avg.append(fallback_box)
box_avg = torch.stack(box_avg)
return (
probs,
confidence,
x0,
box_avg,
torch.tensor(
structured_decode_failed,
dtype=torch.bool,
device=x0.device,
),
)
def structured_decode_failure_pattern(
token_ids: Dict[str, int],
generation_mode: str,
geometry_type: str,
):
"""Return an explicit transition after structured joint decode failure."""
if geometry_type == 'contact':
start = token_ids['grasp_start_token_id']
end = token_ids['grasp_end_token_id']
error_type = 'contact_decode_error'
elif geometry_type == 'grasp_rect':
start = token_ids['grasp_rect_start_token_id']
end = token_ids['grasp_rect_end_token_id']
error_type = 'grasp_rect_decode_error'
else:
raise ValueError(
'structured decode failure is only valid for contact or grasp_rect'
)
if generation_mode == 'fast':
return {
'type': error_type,
'tokens': [start, end],
'need_switch_to_ar': False,
'is_terminal': True,
}
if generation_mode == 'hybrid':
return {
'type': 'error_box',
'tokens': [start],
'need_switch_to_ar': True,
'is_terminal': False,
}
raise ValueError(
'structured MTP decode failure is invalid in slow generation mode'
)
def decode_contact_pair(
logits,
probs,
token_ids: Dict[str, int],
keep_k=4,
start_thresh=0.7,
end_thresh=0.2,
image_size=None,
minimum_width_diagonal=1e-4,
maximum_width_diagonal=1.0,
coord_mass_threshold=1e-4,
force_frame=False,
):
"""Jointly decode four contact coordinates under image-space constraints."""
del logits
coord_start = token_ids['coord_start_token_id']
coord_end = token_ids['coord_end_token_id']
grasp_start = token_ids.get(
'grasp_start_token_id', token_ids['box_start_token_id']
)
grasp_end = token_ids.get(
'grasp_end_token_id', token_ids['box_end_token_id']
)
none_token = token_ids['none_token_id']
null_token = token_ids['null_token_id']
if coord_end - coord_start != 1000:
raise ValueError('contact decoding requires 1001 contiguous coordinate tokens')
if not force_frame and probs[0, grasp_start] < start_thresh:
return None
if force_frame and probs[1, none_token] >= probs[
1, coord_start : coord_end + 1
].max():
return torch.tensor(
[grasp_start, none_token, grasp_end, null_token, null_token, null_token],
dtype=torch.long,
device=probs.device,
)
contact_token_ids = dict(token_ids)
contact_token_ids['box_start_token_id'] = grasp_start
contact_token_ids['box_end_token_id'] = grasp_end
box_type = is_valid_box_frame(
probs,
contact_token_ids,
start_thresh=0.0 if force_frame else start_thresh,
end_thresh=0.0 if force_frame else end_thresh,
topk=keep_k,
)
if box_type == 'empty_box':
return torch.tensor(
[grasp_start, none_token, grasp_end, null_token, null_token, null_token],
dtype=torch.long,
device=probs.device,
)
if box_type == 'illegal_box':
return None
coordinate_probs = probs[1:5, coord_start : coord_end + 1]
coordinate_mass = coordinate_probs.sum(dim=-1)
if (coordinate_mass < coord_mass_threshold).any():
return None
keep_k = max(1, min(int(keep_k), coordinate_probs.shape[-1]))
top_probs, top_values = coordinate_probs.topk(keep_k, dim=-1)
choice_axis = torch.arange(keep_k, device=probs.device)
combinations = torch.cartesian_prod(
choice_axis, choice_axis, choice_axis, choice_axis
)
if combinations.ndim == 1:
combinations = combinations.unsqueeze(0)
positions = torch.arange(4, device=probs.device).unsqueeze(1)
choices = combinations.transpose(0, 1)
candidate_values = top_values[positions, choices].transpose(0, 1).float()
candidate_log_scores = (
top_probs[positions, choices].clamp_min(1e-30).log().sum(dim=0)
)
if image_size is None:
image_width = image_height = 1.0
else:
size = torch.as_tensor(image_size).flatten()
if size.numel() != 2:
raise ValueError('image_size must be (width, height) for contact decoding')
image_width = float(size[0].item())
image_height = float(size[1].item())
if image_width <= 0 or image_height <= 0:
raise ValueError('image_size values must be positive')
dx = (candidate_values[:, 2] - candidate_values[:, 0]) * image_width
dy = (candidate_values[:, 3] - candidate_values[:, 1]) * image_height
width_diagonal = torch.sqrt(dx.square() + dy.square()) / (
1000.0 * (image_width ** 2 + image_height ** 2) ** 0.5
)
valid = (
(width_diagonal >= float(minimum_width_diagonal))
& (width_diagonal <= float(maximum_width_diagonal))
)
if not valid.any():
return None
candidate_log_scores = candidate_log_scores.masked_fill(~valid, -torch.inf)
best_values = candidate_values[candidate_log_scores.argmax()].long()
return torch.cat(
(
best_values.new_tensor([grasp_start]),
best_values + coord_start,
best_values.new_tensor([grasp_end]),
)
)
def decode_grasp_rectangle(
logits,
probs,
token_ids: Dict[str, int],
keep_k=4,
start_thresh=0.7,
end_thresh=0.2,
image_size=None,
minimum_width_diagonal=1e-4,
gripper_depth_pixels=40.0,
coord_mass_threshold=1e-4,
coord_entropy_threshold=1.0,
force_frame=False,
):
"""Jointly decode center, circular angle bin, and opening width."""
del logits
coord_start = token_ids['coord_start_token_id']
coord_end = token_ids['coord_end_token_id']
rect_start = token_ids.get(
'grasp_rect_start_token_id', token_ids['box_start_token_id']
)
rect_end = token_ids.get(
'grasp_rect_end_token_id', token_ids['box_end_token_id']
)
none_token = token_ids['none_token_id']
null_token = token_ids['null_token_id']
if coord_end - coord_start != 1000:
raise ValueError(
'grasp rect decoding requires 1001 contiguous coordinate tokens'
)
if float(minimum_width_diagonal) < 0.0:
raise ValueError('minimum_width_diagonal must be non-negative')
if float(gripper_depth_pixels) <= 0.0:
raise ValueError('gripper_depth_pixels must be positive')
if not 0.0 <= float(coord_entropy_threshold) <= 1.0:
raise ValueError('coord_entropy_threshold must be in [0, 1]')
if image_size is not None:
size = torch.as_tensor(image_size).flatten()
if size.numel() != 2:
raise ValueError('image_size must be (width, height) for grasp rect')
if float(size[0].item()) <= 0 or float(size[1].item()) <= 0:
raise ValueError('image_size values must be positive')
if not force_frame and probs[0, rect_start] < start_thresh:
return None
if force_frame and probs[1, none_token] >= probs[
1, coord_start : coord_end + 1
].max():
return torch.tensor(
[rect_start, none_token, rect_end, null_token, null_token, null_token],
dtype=torch.long,
device=probs.device,
)
rect_token_ids = dict(token_ids)
rect_token_ids['box_start_token_id'] = rect_start
rect_token_ids['box_end_token_id'] = rect_end
box_type = is_valid_box_frame(
probs,
rect_token_ids,
start_thresh=0.0 if force_frame else start_thresh,
end_thresh=0.0 if force_frame else end_thresh,
topk=keep_k,
)
if box_type == 'empty_box':
return torch.tensor(
[rect_start, none_token, rect_end, null_token, null_token, null_token],
dtype=torch.long,
device=probs.device,
)
if box_type == 'illegal_box':
return None
coordinate_probs = probs[1:5, coord_start : coord_end + 1]
coordinate_mass = coordinate_probs.sum(dim=-1)
if (coordinate_mass < coord_mass_threshold).any():
return None
conditional_probs = coordinate_probs / coordinate_mass.unsqueeze(-1).clamp_min(
1e-12
)
coordinate_entropy = -(
conditional_probs * conditional_probs.clamp_min(1e-12).log()
).sum(dim=-1) / torch.log(
conditional_probs.new_tensor(float(conditional_probs.shape[-1]))
)
if (coordinate_entropy > float(coord_entropy_threshold)).any():
return None
keep_k = max(1, min(int(keep_k), coordinate_probs.shape[-1]))
top_probs, top_values = coordinate_probs.topk(keep_k, dim=-1)
choice_axis = torch.arange(keep_k, device=probs.device)
combinations = torch.cartesian_prod(
choice_axis, choice_axis, choice_axis, choice_axis
)
if combinations.ndim == 1:
combinations = combinations.unsqueeze(0)
positions = torch.arange(4, device=probs.device).unsqueeze(1)
choices = combinations.transpose(0, 1)
candidate_values = top_values[positions, choices].transpose(0, 1).float()
candidate_log_scores = (
top_probs[positions, choices].clamp_min(1e-30).log().sum(dim=0)
)
width_diagonal = candidate_values[:, 3] / 1000.0
valid = width_diagonal > float(minimum_width_diagonal)
if not valid.any():
return None
candidate_log_scores = candidate_log_scores.masked_fill(~valid, -torch.inf)
best_values = candidate_values[candidate_log_scores.argmax()].long()
return torch.cat(
(
best_values.new_tensor([rect_start]),
best_values + coord_start,
best_values.new_tensor([rect_end]),
)
)
def constrain_contact_ar_token(next_token_logits, generated, token_ids):
"""Apply the dedicated contact vocabulary mask to one AR decoding slot."""
grasp_start = token_ids.get(
'grasp_start_token_id', token_ids['box_start_token_id']
)
grasp_end = token_ids.get(
'grasp_end_token_id', token_ids['box_end_token_id']
)
coord_start = token_ids['coord_start_token_id']
coord_end = token_ids['coord_end_token_id']
none_token = token_ids['none_token_id']
ref_end = token_ids['ref_end_token_id']
sequence = generated[0].tolist()
def forced(token_id, out_type):
return out_type, torch.tensor(
[token_id], dtype=generated.dtype, device=generated.device
)
if not sequence:
return None
last_grasp_start = max(
(index for index, token in enumerate(sequence) if token == grasp_start),
default=-1,
)
last_grasp_end = max(
(index for index, token in enumerate(sequence) if token == grasp_end),
default=-1,
)
if sequence[-1] == ref_end and last_grasp_start <= last_grasp_end:
return forced(grasp_start, 'continue_ar')
if last_grasp_start <= last_grasp_end:
return None
content = sequence[last_grasp_start + 1 :]
if content and content[0] == none_token:
return forced(grasp_end, 'box_end_ar')
coordinate_count = sum(
coord_start <= token <= coord_end for token in content
)
if coordinate_count >= 4:
return forced(grasp_end, 'box_end_ar')
if any(not coord_start <= token <= coord_end for token in content):
return forced(grasp_end, 'box_end_ar')
coord_logits = next_token_logits[0, 0, coord_start : coord_end + 1]
coord_token = int(coord_logits.argmax().item()) + coord_start
if not content and next_token_logits[0, 0, none_token] >= coord_logits.max():
return forced(none_token, 'coord_ar')
return forced(coord_token, 'coord_ar')
def constrain_grasp_rect_ar_token(next_token_logits, generated, token_ids):
"""Apply the dedicated grasp-rectangle vocabulary mask to one AR slot."""
rect_start = token_ids.get(
'grasp_rect_start_token_id', token_ids['box_start_token_id']
)
rect_end = token_ids.get(
'grasp_rect_end_token_id', token_ids['box_end_token_id']
)
coord_start = token_ids['coord_start_token_id']
coord_end = token_ids['coord_end_token_id']
none_token = token_ids['none_token_id']
ref_end = token_ids['ref_end_token_id']
sequence = generated[0].tolist()
def forced(token_id, out_type):
return out_type, torch.tensor(
[token_id], dtype=generated.dtype, device=generated.device
)
if not sequence:
return None
last_start = max(
(index for index, token in enumerate(sequence) if token == rect_start),
default=-1,
)
last_end = max(
(index for index, token in enumerate(sequence) if token == rect_end),
default=-1,
)
if sequence[-1] == ref_end and last_start <= last_end:
return forced(rect_start, 'continue_ar')
if last_start <= last_end:
return None
content = sequence[last_start + 1 :]
if content and content[0] == none_token:
return forced(rect_end, 'box_end_ar')
coordinate_count = sum(coord_start <= token <= coord_end for token in content)
if coordinate_count >= 4:
return forced(rect_end, 'box_end_ar')
if any(not coord_start <= token <= coord_end for token in content):
return forced(rect_end, 'box_end_ar')
coord_logits = next_token_logits[0, 0, coord_start : coord_end + 1]
coord_token = int(coord_logits.argmax().item()) + coord_start
if not content and next_token_logits[0, 0, none_token] >= coord_logits.max():
return forced(none_token, 'coord_ar')
return forced(coord_token, 'coord_ar')
def sample_tokens_ar(
logits: torch.Tensor,
generated: torch.Tensor,
token_ids: Dict[str, int],
**generate_kwargs,
):
"""
Lightweight sampling function for AR single-step sampling only.
Args:
logits: [batch_size, vocab_size] or [batch_size, 1, vocab_size]
generated: [batch_size, seq_len]
"""
# Convert to 3D for reusing repetition penalty and clipping logic
if logits.dim() == 2:
logits = logits.unsqueeze(1) # [B, 1, V]
batch_size, seq_len, vocab_size = logits.shape
assert seq_len == 1, "sample_tokens_ar only supports single-step AR sampling (seq_len == 1)"
repetition_penalty = generate_kwargs.get('repetition_penalty', 1.0)
temperature = generate_kwargs.get('temperature', 0)
top_p = generate_kwargs.get('top_p', None)
top_k = generate_kwargs.get('top_k', None)
# Apply repetition penalty only based on historically generated tokens
if repetition_penalty != 1.0:
logits = apply_repetition_penalty(logits, generated, repetition_penalty)
if temperature > 0:
logits = logits / temperature
if top_p is not None and top_p < 1:
logits = top_p_logits(logits, top_p)
if top_k is not None:
logits = top_k_logits(logits, top_k)
probs = torch.softmax(logits, dim=-1)
if temperature > 0:
try:
x0 = dists.Categorical(probs=probs).sample()
confidence = torch.gather(probs, -1, x0.unsqueeze(-1)).squeeze(-1)
except Exception:
confidence, x0 = probs.max(dim=-1)
else:
# For greedy: directly take the token with maximum probability
confidence, x0 = probs.max(dim=-1)
# Keep interface consistent with sample_tokens: return [B, 1, V] / [B, 1] shape
return probs, confidence, x0, None, None
def is_valid_box_frame(
probs,
token_ids: Dict[str, int],
start_thresh=0.6,
end_thresh=0.2,
topk=5,
):
box_start_token_id = token_ids['box_start_token_id']
box_end_token_id = token_ids['box_end_token_id']
null_token_id = token_ids['null_token_id']
im_end_token_id = token_ids['im_end_token_id']
none_token_id = token_ids['none_token_id'] # none
p_start = probs[0, box_start_token_id]
if p_start >= start_thresh:
if (probs[1, none_token_id] > 0.2 and
probs[2, box_end_token_id] > 0.2 and
probs[3, null_token_id] > 0.1 and
probs[4, null_token_id] > 0.1):
return 'empty_box'
end_target_ids = torch.tensor([box_end_token_id, null_token_id, im_end_token_id], device=probs.device)
end_score = probs[5, end_target_ids].sum()
if end_score >= end_thresh:
return 'legal_box'
return 'illegal_box'
def decode_bbox_avg(
logits,
probs,
token_ids: Dict[str, int],
keep_k=5,
start_thresh=0.7,
end_thresh=0.2,
generation_mode: str = 'hybrid',
):
"""
Decode bounding box coordinates using top-k weighted average.
Args:
logits: Logits of shape (6, vocab_size)
probs: Probability distribution of shape (6, vocab_size)
token_ids: Dictionary containing all token IDs
keep_k: Number of top-k candidate tokens to keep at each position
start_thresh: Confidence threshold for box start token
end_thresh: Confidence threshold for box end token
Returns:
Decoded bounding box coordinate list in format [box_start, x1, x2, y1, y2, box_end],
or None if decoding fails
"""
coord_start_token_id = token_ids['coord_start_token_id']
coord_end_token_id = token_ids['coord_end_token_id']
box_start_token_id = token_ids['box_start_token_id']
box_end_token_id = token_ids['box_end_token_id']
none_token_id = token_ids['none_token_id']
device = logits.device
box_type = is_valid_box_frame(
probs,
token_ids,
start_thresh=start_thresh,
end_thresh=end_thresh,
topk=keep_k
)
if box_type == 'empty_box':
# Handle the none case first
return torch.tensor([
box_start_token_id,
none_token_id,
box_end_token_id,
token_ids['null_token_id'],
token_ids['null_token_id'],
token_ids['null_token_id']
], dtype=torch.long, device=probs.device)
elif box_type == 'illegal_box':
return None
# Extract probabilities at positions 1-4 and compute Top-K for all 4 positions at once
pos_probs, pos_ids = torch.topk(probs[1:5], k=keep_k, dim=-1)
mask = (pos_ids >= coord_start_token_id) & (pos_ids <= coord_end_token_id)
has_valid = mask.any(dim=-1) # shape: [4]
if not has_valid.all():
return None # not a box, exit...
first_valid_idx = mask.long().argmax(dim=-1, keepdim=True) # [4, 1]
# Extract highest-probability valid_probs[0] and corresponding valid_ids[0]
first_valid_probs = pos_probs.gather(-1, first_valid_idx).squeeze(-1) # [4]
first_valid_ids = pos_ids.gather(-1, first_valid_idx).squeeze(-1) # [4]
if generation_mode == 'hybrid':
valid_counts = mask.sum(dim=-1) # [4]
# Compute max/min of valid ids: fill invalid positions with extreme values to avoid interfering with max/min
LARGE_NUM, SMALL_NUM = 999999, -999999
valid_ids_for_max = torch.where(mask, pos_ids, torch.tensor(SMALL_NUM, device=device))
valid_ids_for_min = torch.where(mask, pos_ids, torch.tensor(LARGE_NUM, device=device))
valid_max = valid_ids_for_max.max(dim=-1)[0]
valid_min = valid_ids_for_min.min(dim=-1)[0]
is_abnormal = (first_valid_probs < 0.9) & (valid_counts > 1) & ((valid_max - valid_min) > 60)
# is_abnormal = (first_valid_probs < 0.7) & (valid_counts > 1) & ((valid_max - valid_min) > 80)
# Normal positions take top-1 (first_valid_ids); abnormal positions are replaced with 0
final_coords = torch.where(is_abnormal, torch.tensor(0, device=pos_ids.device), first_valid_ids)
elif generation_mode == 'fast':
final_coords = first_valid_ids
start_t = torch.tensor([box_start_token_id], dtype=final_coords.dtype, device=device)
end_t = torch.tensor([box_end_token_id], dtype=final_coords.dtype, device=device)
return torch.cat([start_t, final_coords, end_t])
def decode_ref(
logits,
probs,
token_ids: Dict[str, int],
keep_k=5,
start_thresh=0.6,
):
ref_start_token_id = token_ids.get('ref_start_token_id')
coord_start_token_id = token_ids['coord_start_token_id']
coord_end_token_id = token_ids['coord_end_token_id']
device = probs.device
L = probs.size(0)
# 1. Check if the first position is [ and its probability meets start_thresh
# Note: we directly use the probability of the ref token at position 0 for the check
if probs[0, ref_start_token_id] < start_thresh:
return None
# 2. Extract Top-K probabilities and token IDs for all subsequent positions
pos_probs, pos_ids = torch.topk(probs[1:], k=keep_k, dim=-1) # shape: [L-1, keep_k]
# 3. Build mask: identify coordinate tokens (<0> ~ <1000>)
is_coord = (pos_ids >= coord_start_token_id) & (pos_ids <= coord_end_token_id)
# Invert: valid tokens are non-coordinate tokens
is_valid = ~is_coord # shape: [L-1, keep_k]
# Ensure each position has at least one non-coordinate valid token in its Top-K
has_valid = is_valid.any(dim=-1) # shape: [L-1]
if not has_valid.all():
return None
# 4. Get the highest-probability valid token
# Since topk results are sorted in descending order of probability,
# argmax returns the first index where is_valid is True, i.e., the index of the most probable valid token
first_valid_idx = is_valid.long().argmax(dim=-1, keepdim=True) # shape: [L-1, 1]
# Extract the final token IDs
final_text_ids = pos_ids.gather(-1, first_valid_idx).squeeze(-1) # shape: [L-1]
start_t = torch.tensor([ref_start_token_id], dtype=final_text_ids.dtype, device=device)
return torch.cat([start_t, final_text_ids])
def handle_pattern(
x0,
token_ids: Dict[str, int],
generation_mode: str = 'hybrid',
geometry_type: str = 'bbox',
):
"""
Args:
x0: Token ID list of length 6
token_ids: Dictionary containing all token IDs
"""
null_token_id = token_ids['null_token_id']
im_end_token_id = token_ids['im_end_token_id']
if geometry_type == 'contact':
box_start_token_id = token_ids.get(
'grasp_start_token_id', token_ids['box_start_token_id']
)
box_end_token_id = token_ids.get(
'grasp_end_token_id', token_ids['box_end_token_id']
)
elif geometry_type == 'grasp_rect':
box_start_token_id = token_ids.get(
'grasp_rect_start_token_id', token_ids['box_start_token_id']
)
box_end_token_id = token_ids.get(
'grasp_rect_end_token_id', token_ids['box_end_token_id']
)
else:
box_start_token_id = token_ids['box_start_token_id']
box_end_token_id = token_ids['box_end_token_id']
none_token_id = token_ids['none_token_id']
coord_start_token_id = token_ids['coord_start_token_id']
coord_end_token_id = token_ids['coord_end_token_id']
ref_end_token_id = token_ids['ref_end_token_id']
x0 = x0.tolist()
if x0[0] == null_token_id:
return {
"type": "im_end",
"tokens": [im_end_token_id],
"need_switch_to_ar": False,
"is_terminal": True,
}
elif x0[0] == im_end_token_id:
return {
"type": "im_end",
"tokens": [im_end_token_id],
"need_switch_to_ar": False,
"is_terminal": True,
}
elif x0[:2] == [box_start_token_id, none_token_id]:
return {
"type": "empty_box",
"tokens": [box_start_token_id, none_token_id, box_end_token_id],
"need_switch_to_ar": False,
"is_terminal": False,
}
elif x0[0] == box_start_token_id:
coord_ix = 1
for coord in x0[1:5]:
if coord_start_token_id <= coord <= coord_end_token_id:
coord_ix += 1
else:
break
# Four-coordinate bbox or contact pair, selected by the caller's task.
if coord_ix == 5 and x0[5] == box_end_token_id:
pattern_type = {
'contact': 'contact_box',
'grasp_rect': 'grasp_rect_box',
}.get(geometry_type, 'coord_box')
return {
"type": pattern_type,
"tokens": x0,
"need_switch_to_ar": False,
"is_terminal": False,
}
# Two-coordinate pointing:
# Convention: the first two coordinates are valid coord tokens, the third token is box_end.
# Remaining positions (if any) are not part of the pattern; truncate at box_end.
elif (
geometry_type not in ('contact', 'grasp_rect')
and coord_ix == 3
and x0[3] == box_end_token_id
):
return {
"type": "point_box",
"tokens": x0[:4],
"need_switch_to_ar": False,
"is_terminal": False,
}
else:
if generation_mode == 'fast':
if geometry_type == 'contact':
return {
"type": "contact_decode_error",
"tokens": [box_start_token_id, box_end_token_id],
"need_switch_to_ar": False,
"is_terminal": True,
}
if geometry_type == 'grasp_rect':
return {
"type": "grasp_rect_decode_error",
"tokens": [box_start_token_id, box_end_token_id],
"need_switch_to_ar": False,
"is_terminal": True,
}
# fast mode: treat as coord_box, stay in MTP
return {
"type": "coord_box",
"tokens": x0,
"need_switch_to_ar": False,
"is_terminal": False,
}
else:
# hybrid mode: error_box, switch to AR
return {
"type": "error_box",
"tokens": x0[:coord_ix],
"need_switch_to_ar": True,
"is_terminal": False,
}
else:
for i, token in enumerate(x0):
if token == null_token_id:
x0 = x0[:i]
break
if len(x0) >= 2 and x0[-1] == x0[-2] == ref_end_token_id:
x0 = x0[:-1]
return {
"type": "ref_object",
"tokens": x0,
"need_switch_to_ar": False,
"is_terminal": False,
}
]