# 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, }