"""Autoregressive decoding using forward calls, including batched cached beams.""" from typing import Callable import torch from torch import Tensor ModelForward = Callable[[Tensor], Tensor] def _check(input_ids, max_new_tokens): """Validate the prompt tensor shape and nonnegative generation budget.""" if input_ids.ndim != 2 or input_ids.shape[1] == 0 or max_new_tokens < 0: raise ValueError( "Expected nonempty (batch, sequence) prompts and nonnegative length" ) def _forward(model, ids, mask, cache): # Hugging Face models expose config; simple callable test models need only ids. """Return next-token logits and an optional cache using only model forward calls.""" if hasattr(model, "config") and hasattr(model.config, "model_type"): new_ids = ids if cache is None else ids[:, -1:] positions = mask.long().cumsum(-1).sub(1).clamp_min(0)[:, -new_ids.shape[1] :] result = model.forward( input_ids=new_ids, attention_mask=mask, position_ids=positions, past_key_values=cache, use_cache=True, ) return result.logits[:, -1].float(), getattr(result, "past_key_values", None) result = ( model.forward(ids, attention_mask=mask) if hasattr(model, "forward") else model(ids) ) logits = result.logits if hasattr(result, "logits") else result return logits[:, -1].float(), None def _reorder(cache, indices): """Reorder cached beam states according to their selected parent indices.""" if cache is None: return None if hasattr(cache, "reorder_cache"): cache.reorder_cache(indices) return cache return tuple(tuple(t.index_select(0, indices) for t in layer) for layer in cache) @torch.no_grad() def _sample( model, input_ids, max_new_tokens, mode, k=None, p=None, temperature=1.0, eos_token_id=None, attention_mask=None, ): """Append tokens using greedy, top-k or nucleus selection. Maintain attention masks and optional cached states. Finished sequences append only EOS while other sequences continue within the token budget. """ _check(input_ids, max_new_tokens) if temperature <= 0: raise ValueError("temperature must be positive") ids = input_ids.clone() mask = torch.ones_like(ids) if attention_mask is None else attention_mask.clone() finished = torch.zeros(ids.shape[0], device=ids.device, dtype=torch.bool) cache = None for _ in range(max_new_tokens): logits, cache = _forward(model, ids, mask, cache) if mode == "greedy": token = logits.argmax(-1) else: logits = logits / temperature if mode == "top_k": values, indices = logits.topk(min(k, logits.shape[-1]), dim=-1) choice = torch.multinomial(values.softmax(-1), 1) token = indices.gather(1, choice).squeeze(1) else: values, indices = logits.sort(descending=True, dim=-1, stable=True) probs = values.softmax(-1) remove = probs.cumsum(-1) - probs >= p values = values.masked_fill(remove, -torch.inf) choice = torch.multinomial(values.softmax(-1), 1) token = indices.gather(1, choice).squeeze(1) if eos_token_id is not None: token = torch.where(finished, eos_token_id, token) finished |= token.eq(eos_token_id) ids = torch.cat((ids, token[:, None]), -1) mask = torch.cat((mask, torch.ones_like(token[:, None])), -1) if bool(finished.all()): break return ids def greedy_decode(model_forward, input_ids, max_new_tokens, **kwargs): """Append the highest-logit token at each step until EOS or the token limit.""" return _sample(model_forward, input_ids, max_new_tokens, "greedy", **kwargs) def top_k_decode(model_forward, input_ids, k, max_new_tokens, **kwargs): """Sample among the highest k logits; k=1 uses deterministic greedy decoding.""" if k < 1: raise ValueError("k must be positive") if k == 1: return greedy_decode(model_forward, input_ids, max_new_tokens, **kwargs) return _sample(model_forward, input_ids, max_new_tokens, "top_k", k=k, **kwargs) def top_p_decode(model_forward, input_ids, p, max_new_tokens, **kwargs): """Sample from the smallest prefix reaching the requested mass.""" if not 0 < p <= 1: raise ValueError("p must be in (0, 1]") return _sample(model_forward, input_ids, max_new_tokens, "top_p", p=p, **kwargs) @torch.no_grad() def beam_search( model_forward, input_ids, width, max_new_tokens, eos_token_id=None, length_penalty=0.0, attention_mask=None, ): """Return the best cumulative-log-probability beam for each prompt. Keep the requested number of hypotheses and reorder caches after pruning. Completed beams keep their scores; optional length normalization is used only to choose the final hypothesis. Width one matches greedy decoding. """ _check(input_ids, max_new_tokens) if width < 1 or length_penalty < 0: raise ValueError("width must be positive and length_penalty nonnegative") if width == 1: return greedy_decode( model_forward, input_ids, max_new_tokens, eos_token_id=eos_token_id, attention_mask=attention_mask, ) if max_new_tokens == 0: return input_ids.clone() b, prompt_length = input_ids.shape ids = input_ids.repeat_interleave(width, 0) mask = ( torch.ones_like(input_ids) if attention_mask is None else attention_mask ).repeat_interleave(width, 0) scores = torch.full((b, width), -torch.inf, device=ids.device) scores[:, 0] = 0.0 finished = torch.zeros(b, width, device=ids.device, dtype=torch.bool) lengths = torch.zeros(b, width, device=ids.device, dtype=torch.long) cache = None for _ in range(max_new_tokens): logits, cache = _forward(model_forward, ids, mask, cache) vocab = logits.shape[-1] if width > vocab: raise ValueError("beam width exceeds vocabulary size") log_probs = logits.log_softmax(-1).reshape(b, width, vocab) if eos_token_id is not None: log_probs = log_probs.masked_fill(finished[..., None], -torch.inf) log_probs[:, :, eos_token_id] = torch.where( finished, 0.0, log_probs[:, :, eos_token_id] ) candidates = scores[..., None] + log_probs scores, flat_indices = candidates.reshape(b, -1).topk(width, -1) parent, token = flat_indices // vocab, flat_indices % vocab global_parent = ( parent + torch.arange(b, device=ids.device)[:, None] * width ).reshape(-1) old_finished = finished.gather(1, parent) lengths = lengths.gather(1, parent) + (~old_finished).long() finished = ( old_finished if eos_token_id is None else old_finished | token.eq(eos_token_id) ) ids = torch.cat((ids.index_select(0, global_parent), token.reshape(-1, 1)), -1) mask = torch.cat( ( mask.index_select(0, global_parent), torch.ones(b * width, 1, dtype=mask.dtype, device=ids.device), ), -1, ) cache = _reorder(cache, global_parent) if bool(finished.all()): break normalized = scores / lengths.clamp_min(1).float().pow(length_penalty) best = normalized.argmax(-1) + torch.arange(b, device=ids.device) * width return ids.index_select(0, best)