| """ |
| 解码策略模块 — 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: |
| """ |
| 贪心解码。 |
| |
| TODO [Person D]: 实现以下逻辑: |
| 1. encoder_output = model.encode(src_ids, src_padding_mask) |
| 2. 初始化 decoder input: [B, 1] 全为 bos_id |
| 3. for step in range(max_len): |
| a. logits = model.decode_step(decoder_input, encoder_output, src_padding_mask) |
| b. next_token = logits.argmax(dim=-1) |
| c. decoder_input = concat(decoder_input, next_token) |
| d. 如果所有序列都生成了 eos_id,则提前终止 |
| 4. 返回生成的 token ids [B, T] |
| """ |
| raise NotImplementedError("TODO: Person D 实现 greedy_decode") |
|
|
|
|
| @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: |
| """ |
| 束搜索解码。 |
| |
| TODO [Person D]: 实现以下逻辑: |
| 1. encoder_output = model.encode(src_ids, src_padding_mask) |
| 2. 将 encoder_output 扩展为 beam_size 份: [B*beam, S, D] |
| 3. 初始化 beam: |
| - beam_scores: [B, beam_size] 初始为 0 |
| - beam_tokens: [B, beam_size, 1] 初始为 bos_id |
| 4. for step in range(max_len): |
| a. 对每个 beam 计算 logits |
| b. log_probs = log_softmax(logits) |
| c. (可选) 应用 no_repeat_ngram 约束 |
| d. scores = beam_scores + log_probs |
| e. 选择 top-k candidates (k = beam_size) |
| f. 更新 beam_tokens 和 beam_scores |
| g. 将已完成的 beam 移到 finished pool |
| 5. 对 finished beams 应用 length_penalty: |
| score = score / (length ^ length_penalty) |
| 6. 选择得分最高的序列 |
| 7. 返回最佳翻译 [B, T] |
| |
| 这是翻译任务最关键的解码算法,请仔细实现。 |
| """ |
| raise NotImplementedError("TODO: Person D 实现 beam_search_decode") |
|
|
|
|
| @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)。 |
| |
| TODO [Person D]: 实现以下逻辑: |
| 1. 与贪心解码类似,但每步采样而非取 argmax |
| 2. 应用 temperature: logits = logits / temperature |
| 3. 应用 top-k: 只保留概率最高的 k 个 token |
| 4. 应用 top-p (nucleus): 只保留累积概率达到 p 的 token |
| 5. 从过滤后的分布中采样: torch.multinomial |
| """ |
| raise NotImplementedError("TODO: Person D 实现 sample_decode") |
|
|
|
|
| def _apply_no_repeat_ngram( |
| logits: torch.Tensor, |
| generated_tokens: torch.Tensor, |
| ngram_size: int, |
| ) -> torch.Tensor: |
| """ |
| 防止生成重复的 n-gram。 |
| |
| TODO [Person D]: |
| 1. 从 generated_tokens 中提取所有已出现的 (ngram_size-1)-gram |
| 2. 对于每个可能导致重复 ngram 的 next token,将其 logits 设为 -inf |
| """ |
| raise NotImplementedError("TODO: Person D 实现 _apply_no_repeat_ngram") |
|
|