| """
|
| 解码策略模块 — 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
|
|
|