Download generate_utils.py from charlesH777/grasp-anything-9040: direct link, hf CLI and curl.
- Browser
- Download file 37 kB
-
https://huggingface.co/charlesH777/grasp-anything-9040/resolve/main/generate_utils.py
- Command line
-
hf download hf://charlesH777/grasp-anything-9040/generate_utils.py
-
curl -L -o generate_utils.py https://huggingface.co/charlesH777/grasp-anything-9040/resolve/main/generate_utils.py
37 kB
| # 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 <box>none</box> 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 <ref> 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: <box><x><y></box> | |
| # 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, | |
| } | |