jiaoruotong's picture
[Person D] Implement evaluation module: metrics, decoding, evaluator, scripts
ef0a52e verified
Raw
History Blame
9.94 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]
"""
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, # [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]
这是翻译任务最关键的解码算法,请仔细实现。
"""
batch_size, seq_len = src_ids.size()
device = src_ids.device
# 1. Encode
encoder_output = encoder_output = model.encode(src_ids, src_padding_mask)
# 2. Expand encoder output for beam search: [B*beam, S, D]
encoder_output = encoder_output.unsqueeze(1).expand(-1, beam_size, -1, -1)
encoder_output = encoder_output.reshape(batch_size * beam_size, seq_len, -1)
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)
# 3. Initialize beams
beam_scores = torch.zeros(batch_size, beam_size, device=device)
beam_scores[:, 1:] = float("-inf") # Only first beam is active initially
beam_tokens = 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)
# 4. Iterative decoding
for _ in range(max_len):
flat_tokens = beam_tokens.view(batch_size * beam_size, -1)
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)
# (c) Apply no_repeat_ngram constraint
if no_repeat_ngram_size > 0:
log_probs = _apply_no_repeat_ngram(log_probs, flat_tokens, no_repeat_ngram_size)
# Mask finished beams
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
# (d) Compute scores
scores = beam_scores.unsqueeze(-1) + log_probs.view(batch_size, beam_size, vocab_size)
scores = scores.view(batch_size, -1) # [B, beam * vocab]
# (e) Select top-k
topk_scores, topk_indices = scores.topk(beam_size, dim=-1)
beam_indices = topk_indices // vocab_size
token_indices = topk_indices % vocab_size
# (f) Update beam tokens and scores
new_tokens = []
new_finished = []
for b in range(batch_size):
prev_seqs = beam_tokens[b][beam_indices[b]]
next_tokens = token_indices[b].unsqueeze(-1)
new_tokens.append(torch.cat([prev_seqs, next_tokens], dim=-1))
new_finished.append(finished[b][beam_indices[b]] | token_indices[b].eq(eos_id))
beam_tokens = torch.stack(new_tokens, dim=0)
finished = torch.stack(new_finished, dim=0)
beam_scores = topk_scores
if finished.all():
break
# 5. Apply length penalty
lengths = beam_tokens.size(-1) - 1 # Exclude BOS
penalties = lengths ** length_penalty
final_scores = beam_scores / penalties
# 6. Select best beam for each batch
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)。
TODO [Person D]: 实现以下逻辑:
1. 与贪心解码类似,但每步采样而非取 argmax
2. 应用 temperature: logits = logits / temperature
3. 应用 top-k: 只保留概率最高的 k 个 token
4. 应用 top-p (nucleus): 只保留累积概率达到 p 的 token
5. 从过滤后的分布中采样: torch.multinomial
"""
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)
# 2. Apply temperature
logits = logits / max(temperature, 1e-8)
# 3. Apply top-k filtering
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")
# 4. Apply top-p (nucleus) filtering
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)
mask = cumulative_probs - F.softmax(sorted_logits, dim=-1) >= top_p
sorted_logits[mask] = float("-inf")
logits = sorted_logits.scatter(1, sorted_indices.argsort(1), sorted_logits)
# 5. Sample from filtered distribution
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。
TODO [Person D]:
1. 从 generated_tokens 中提取所有已出现的 (ngram_size-1)-gram
2. 对于每个可能导致重复 ngram 的 next token,将其 logits 设为 -inf
"""
if ngram_size <= 0:
return logits
batch_size = logits.size(0)
for batch_idx in range(batch_size):
tokens = generated_tokens[batch_idx].tolist()
if len(tokens) < ngram_size - 1:
continue
# Build map of (n-1)-gram prefix -> set of next tokens that appeared
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)
# Check current prefix and ban tokens that would create repeated n-grams
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