""" 解码策略模块 — Person D 负责实现 功能要求: 1. greedy_decode: 贪心解码 2. beam_search_decode: 束搜索解码 3. sample_decode: 采样解码 (temperature, top-k, top-p) 技术要点: - Beam Search 是翻译任务最常用的解码策略 - 需要高效处理批量解码 - 支持长度惩罚 (length penalty) 和重复惩罚 (no_repeat_ngram) - 对于预训练模型,可以直接使用 model.generate() """ from __future__ import annotations import logging from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F logger = logging.getLogger(__name__) def _apply_no_repeat_ngram( logits: torch.Tensor, generated_tokens: torch.Tensor, ngram_size: int, ) -> torch.Tensor: """ 防止生成重复的 n-gram。 """ if ngram_size <= 1: return logits batch_size, vocab_size = logits.size() for batch_idx in range(batch_size): tokens = generated_tokens[batch_idx].tolist() if len(tokens) < ngram_size - 1: continue banned_tokens: set[int] = set() ngram_map: dict[tuple[int, ...], set[int]] = {} for i in range(len(tokens) - ngram_size + 1): prefix = tuple(tokens[i : i + ngram_size - 1]) next_token = tokens[i + ngram_size - 1] ngram_map.setdefault(prefix, set()).add(next_token) prefix = tuple(tokens[-(ngram_size - 1) :]) if prefix in ngram_map: banned_tokens = ngram_map[prefix] logits[batch_idx, list(banned_tokens)] = float("-inf") return logits @torch.no_grad() def greedy_decode( model: nn.Module, src_ids: torch.Tensor, src_padding_mask: torch.BoolTensor, bos_id: int, eos_id: int, max_len: int = 256, ) -> torch.Tensor: """ 贪心解码。 """ encoder_output = model.encode(src_ids, src_padding_mask) batch_size = src_ids.size(0) device = src_ids.device generated = torch.full((batch_size, 1), bos_id, dtype=torch.long, device=device) finished = torch.zeros(batch_size, dtype=torch.bool, device=device) for _ in range(max_len): logits = model.decode_step(generated, encoder_output, src_padding_mask) next_token = logits.argmax(dim=-1, keepdim=True) generated = torch.cat([generated, next_token], dim=1) finished = finished | next_token.squeeze(-1).eq(eos_id) if finished.all(): break return generated @torch.no_grad() def beam_search_decode( model: nn.Module, src_ids: torch.Tensor, src_padding_mask: torch.BoolTensor, bos_id: int, eos_id: int, beam_size: int = 5, max_len: int = 256, length_penalty: float = 1.0, no_repeat_ngram_size: int = 0, ) -> torch.Tensor: """ 束搜索解码。 """ batch_size, seq_len = src_ids.size() device = src_ids.device encoder_output = model.encode(src_ids, src_padding_mask) encoder_output = encoder_output.unsqueeze(1).expand(batch_size, beam_size, -1, -1) encoder_output = encoder_output.reshape(batch_size * beam_size, seq_len, -1) src_padding_mask = src_padding_mask.unsqueeze(1).expand(batch_size, beam_size, seq_len) src_padding_mask = src_padding_mask.reshape(batch_size * beam_size, seq_len) beam_scores = torch.full((batch_size, beam_size), float("-inf"), device=device) beam_scores[:, 0] = 0.0 generated = torch.full((batch_size, beam_size, 1), bos_id, dtype=torch.long, device=device) finished = torch.zeros((batch_size, beam_size), dtype=torch.bool, device=device) for _ in range(max_len): flat_generated = generated.view(batch_size * beam_size, -1) logits = model.decode_step(flat_generated, encoder_output, src_padding_mask) log_probs = F.log_softmax(logits, dim=-1) if no_repeat_ngram_size > 0: log_probs = _apply_no_repeat_ngram(log_probs, flat_generated, no_repeat_ngram_size) finished_flat = finished.view(batch_size * beam_size) if finished_flat.any(): log_probs[finished_flat] = float("-inf") log_probs[finished_flat, eos_id] = 0.0 vocab_size = log_probs.size(-1) scores = beam_scores.unsqueeze(-1) + log_probs.view(batch_size, beam_size, vocab_size) scores_flat = scores.view(batch_size, -1) topk_scores, topk_indices = scores_flat.topk(beam_size, dim=-1) beam_indices = topk_indices // vocab_size token_indices = topk_indices % vocab_size next_generated = [] next_finished = [] for batch_idx in range(batch_size): selected_beams = beam_indices[batch_idx] selected_tokens = token_indices[batch_idx] next_seq = generated[batch_idx, selected_beams] next_seq = torch.cat([next_seq, selected_tokens.unsqueeze(-1)], dim=-1) next_generated.append(next_seq) next_finished.append( finished[batch_idx, selected_beams] | selected_tokens.eq(eos_id) ) generated = torch.stack(next_generated, dim=0) finished = torch.stack(next_finished, dim=0) beam_scores = topk_scores if finished.all(): break length = generated.size(1) penalty = float(length) ** float(length_penalty) final_scores = beam_scores / penalty best_indices = final_scores.argmax(dim=-1) output = generated[torch.arange(batch_size, device=device), best_indices] return output @torch.no_grad() def sample_decode( model: nn.Module, src_ids: torch.Tensor, src_padding_mask: torch.BoolTensor, bos_id: int, eos_id: int, max_len: int = 256, temperature: float = 1.0, top_k: int = 0, top_p: float = 1.0, ) -> torch.Tensor: """ 采样解码 (支持 temperature, top-k, top-p/nucleus sampling)。 """ encoder_output = model.encode(src_ids, src_padding_mask) batch_size = src_ids.size(0) device = src_ids.device generated = torch.full((batch_size, 1), bos_id, dtype=torch.long, device=device) finished = torch.zeros(batch_size, dtype=torch.bool, device=device) for _ in range(max_len): logits = model.decode_step(generated, encoder_output, src_padding_mask) logits = logits / max(temperature, 1e-8) if top_k > 0: top_k = min(top_k, logits.size(-1)) values, indices = torch.topk(logits, top_k, dim=-1) mask = torch.full_like(logits, float("-inf")) logits = mask.scatter(-1, indices, values) if 0.0 < top_p < 1.0: sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1) probs = F.softmax(sorted_logits, dim=-1) cumulative_probs = torch.cumsum(probs, dim=-1) cutoff = cumulative_probs > top_p cutoff[:, 1:] = cutoff[:, :-1].clone() cutoff[:, 0] = False sorted_logits[cutoff] = float("-inf") logits = torch.zeros_like(logits).scatter(-1, sorted_indices, sorted_logits) probs = F.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) next_token = next_token.clamp(min=0) next_token = torch.where(finished.unsqueeze(-1), torch.full_like(next_token, eos_id), next_token) generated = torch.cat([generated, next_token], dim=1) finished = finished | next_token.squeeze(-1).eq(eos_id) if finished.all(): break return generated