lijn14
创建工程
c1a46f7
Raw
History Blame
4 kB
"""
解码策略模块 — 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, # [B, S]
src_padding_mask: torch.BoolTensor, # [B, S]
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, # [B, S]
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")