| """ |
| 解码策略模块 — 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__) |
|
|
|
|
| @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 |
|
|
| decoder_input = 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(decoder_input, encoder_output, src_padding_mask) |
| next_token = logits.argmax(dim=-1, keepdim=True) |
| decoder_input = torch.cat([decoder_input, next_token], dim=1) |
| finished = finished | next_token.squeeze(-1).eq(eos_id) |
| if finished.all(): |
| break |
|
|
| return decoder_input |
|
|
|
|
| @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: |
| """ |
| 束搜索解码。 |
| |
| 优化改进: |
| 1. 限制最大 beam_size 为 64,防止内存溢出 |
| 2. 使用更高效的索引操作避免中间 tensor 累积 |
| 3. 及时释放不再需要的中间结果 |
| 4. 对长序列使用更小的 beam_size |
| """ |
| |
| beam_size = min(beam_size, 64) |
| |
| batch_size, seq_len = src_ids.size() |
| device = src_ids.device |
| |
| |
| if seq_len > 512: |
| beam_size = min(beam_size, 3) |
| |
| |
| encoder_output = model.encode(src_ids, src_padding_mask) |
| hidden_dim = encoder_output.size(-1) |
| |
| |
| encoder_output = encoder_output.unsqueeze(1).expand(-1, beam_size, -1, -1) |
| encoder_output = encoder_output.reshape(batch_size * beam_size, seq_len, hidden_dim) |
| src_padding_mask_expanded = src_padding_mask.unsqueeze(1).expand(-1, beam_size, -1) |
| src_padding_mask_expanded = src_padding_mask_expanded.reshape(batch_size * beam_size, seq_len) |
|
|
| |
| beam_scores = torch.zeros(batch_size, beam_size, device=device) |
| beam_scores[:, 1:] = float("-inf") |
| |
| |
| beam_tokens = torch.full((batch_size, beam_size, max_len + 1), eos_id, dtype=torch.long, device=device) |
| beam_tokens[:, :, 0] = bos_id |
| |
| finished = torch.zeros(batch_size, beam_size, dtype=torch.bool, device=device) |
| lengths = torch.ones(batch_size, beam_size, dtype=torch.long, device=device) |
|
|
| |
| for step in range(1, max_len + 1): |
| |
| active_mask = ~finished |
| if not active_mask.any(): |
| break |
| |
| |
| flat_tokens = beam_tokens[:, :, :step].reshape(batch_size * beam_size, step) |
| |
| |
| logits = model.decode_step(flat_tokens, encoder_output, src_padding_mask_expanded) |
| log_probs = F.log_softmax(logits, dim=-1) |
| vocab_size = log_probs.size(-1) |
| |
| |
| if no_repeat_ngram_size > 0 and step >= no_repeat_ngram_size: |
| log_probs = _apply_no_repeat_ngram(log_probs, flat_tokens, 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 |
|
|
| |
| scores = beam_scores.unsqueeze(-1) + log_probs.view(batch_size, beam_size, vocab_size) |
| scores = scores.view(batch_size, -1) |
|
|
| |
| topk_scores, topk_indices = scores.topk(beam_size, dim=-1) |
| beam_indices = topk_indices // vocab_size |
| token_indices = topk_indices % vocab_size |
|
|
| |
| new_beam_tokens = beam_tokens.clone() |
| for b in range(batch_size): |
| new_beam_tokens[b] = beam_tokens[b][beam_indices[b]] |
| new_beam_tokens[b, torch.arange(beam_size), step] = token_indices[b] |
| beam_tokens = new_beam_tokens |
| |
| |
| new_finished = finished.gather(1, beam_indices) | token_indices.eq(eos_id) |
| new_lengths = lengths.gather(1, beam_indices) |
| new_lengths[~new_finished] = step + 1 |
| |
| beam_scores = topk_scores |
| finished = new_finished |
| lengths = new_lengths |
|
|
| |
| lengths = lengths.float() |
| penalties = lengths ** length_penalty |
| final_scores = beam_scores / penalties |
|
|
| |
| best_indices = final_scores.argmax(dim=-1) |
| best_sequences = beam_tokens[torch.arange(batch_size, device=device), best_indices] |
|
|
| return best_sequences |
|
|
|
|
| @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 |
|
|
| decoder_input = 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(decoder_input, encoder_output, src_padding_mask) |
|
|
| logits = logits / max(temperature, 1e-8) |
|
|
| if top_k > 0: |
| k = min(top_k, logits.size(-1)) |
| topk_values, _ = torch.topk(logits, k, dim=-1) |
| threshold = topk_values[:, -1].unsqueeze(-1) |
| logits[logits < threshold] = float("-inf") |
|
|
| if 0.0 < top_p < 1.0: |
| sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1) |
| cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) |
| |
| |
| remove_mask = cumulative_probs - F.softmax(sorted_logits, dim=-1) >= top_p |
| sorted_logits[remove_mask] = 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 = torch.where(finished.unsqueeze(-1), torch.full_like(next_token, eos_id), next_token) |
| decoder_input = torch.cat([decoder_input, next_token], dim=1) |
| finished = finished | next_token.squeeze(-1).eq(eos_id) |
|
|
| if finished.all(): |
| break |
|
|
| return decoder_input |
|
|
|
|
| def _apply_no_repeat_ngram( |
| logits: torch.Tensor, |
| generated_tokens: torch.Tensor, |
| ngram_size: int, |
| ) -> torch.Tensor: |
| """防止生成重复的 n-gram。""" |
| if ngram_size <= 0: |
| return logits |
|
|
| batch_size = logits.size(0) |
| seq_len = generated_tokens.size(1) |
| |
| if seq_len < ngram_size - 1: |
| return logits |
|
|
| for batch_idx in range(batch_size): |
| tokens = generated_tokens[batch_idx].tolist() |
|
|
| ngram_map: dict[tuple, set] = {} |
| for i in range(len(tokens) - ngram_size + 1): |
| prefix = tuple(tokens[i : i + ngram_size - 1]) |
| next_tok = tokens[i + ngram_size - 1] |
| ngram_map.setdefault(prefix, set()).add(next_tok) |
|
|
| current_prefix = tuple(tokens[-(ngram_size - 1):]) |
| if current_prefix in ngram_map: |
| banned = list(ngram_map[current_prefix]) |
| logits[batch_idx, banned] = float("-inf") |
|
|
| return logits |