Download decoding.py from vanshnawander/assignment2-moe-shared: direct link, hf CLI and curl.
- Browser
- Download file 7.77 kB
-
https://huggingface.co/vanshnawander/assignment2-moe-shared/resolve/main/decoding.py
- Command line
-
hf download hf://vanshnawander/assignment2-moe-shared/decoding.py
-
curl -L -o decoding.py https://huggingface.co/vanshnawander/assignment2-moe-shared/resolve/main/decoding.py
7.77 kB
| """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) | |
| 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) | |
| 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) | |