"""Shared grasp task-token adapter operations for training and inference.""" from __future__ import annotations import torch import torch.nn.functional as F def apply_grasp_task_output_delta( hidden_states: torch.Tensor, logits: torch.Tensor, task_token_ids: torch.Tensor, output_delta: torch.Tensor, ) -> torch.Tensor: """Apply the compact grasp-token output adapter to every logit row.""" if hidden_states.shape[:-1] != logits.shape[:-1]: raise ValueError( "hidden-state/logit leading dimensions differ: " f"hidden={tuple(hidden_states.shape)}, logits={tuple(logits.shape)}" ) if output_delta.ndim != 2 or output_delta.shape[1] != hidden_states.shape[-1]: raise ValueError( "grasp output adapter has incompatible shape: " f"delta={tuple(output_delta.shape)}, hidden={tuple(hidden_states.shape)}" ) task_token_ids = task_token_ids.to(device=logits.device, dtype=torch.long) if task_token_ids.numel() != output_delta.shape[0]: raise ValueError( "grasp task-token IDs and output-adapter rows differ: " f"tokens={task_token_ids.numel()}, rows={output_delta.shape[0]}" ) correction = F.linear( hidden_states.to(dtype=output_delta.dtype), output_delta ).to(dtype=logits.dtype) return logits.index_add(-1, task_token_ids, correction)