File size: 6,105 Bytes
c1a46f7 ef0a52e c1a46f7 ef0a52e c1a46f7 ef0a52e c1a46f7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """
评估器模块 — Person D 负责实现
功能要求:
将解码和评估指标整合为统一的评估接口。
使用方法:
evaluator = Evaluator(model, tokenizer, config)
results = evaluator.evaluate(test_loader)
"""
from __future__ import annotations
import logging
from typing import Optional
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from tqdm import tqdm
from easytranslate.evaluation.metrics import compute_all_metrics
from easytranslate.evaluation.decoding import greedy_decode, beam_search_decode, sample_decode
logger = logging.getLogger(__name__)
class Evaluator:
"""
翻译模型评估器。
TODO [Person D]: 实现以下方法。
"""
def __init__(self, model: nn.Module, tokenizer, config: dict):
"""
TODO [Person D]:
1. 保存 model, tokenizer, config
2. 从 config 读取解码策略和评估指标配置
3. 根据策略选择解码函数
"""
self.model = model
self.tokenizer = tokenizer
self.config = config
# 从 config 读取评估和解码配置
eval_config = config.get("evaluation", {})
decoding_config = eval_config.get("decoding", {})
self.strategy = decoding_config.get("strategy", "beam_search")
self.max_len = decoding_config.get("max_decode_len", 256)
self.beam_size = decoding_config.get("beam_size", 5)
self.length_penalty = decoding_config.get("length_penalty", 1.0)
self.no_repeat_ngram_size = decoding_config.get("no_repeat_ngram_size", 0)
sampling_config = decoding_config.get("sampling", {})
self.temperature = sampling_config.get("temperature", 1.0)
self.top_k = sampling_config.get("top_k", 0)
self.top_p = sampling_config.get("top_p", 1.0)
self.metrics = eval_config.get("metrics", ["bleu", "comet", "chrf", "ter"])
self.bos_id = tokenizer.bos_token_id
self.eos_id = tokenizer.eos_token_id
self.pad_id = tokenizer.pad_token_id
# 根据策略选择解码函数
self.decode_fn = self._get_decode_fn()
def _get_decode_fn(self):
"""根据策略选择解码函数。"""
if self.strategy == "greedy":
return self._greedy
elif self.strategy == "sampling":
return self._sample
else:
return self._beam_search
def _greedy(self, src_ids, src_padding_mask):
return greedy_decode(
self.model, src_ids, src_padding_mask,
self.bos_id, self.eos_id, max_len=self.max_len,
)
def _beam_search(self, src_ids, src_padding_mask):
return beam_search_decode(
self.model, src_ids, src_padding_mask,
self.bos_id, self.eos_id,
beam_size=self.beam_size, max_len=self.max_len,
length_penalty=self.length_penalty,
no_repeat_ngram_size=self.no_repeat_ngram_size,
)
def _sample(self, src_ids, src_padding_mask):
return sample_decode(
self.model, src_ids, src_padding_mask,
self.bos_id, self.eos_id,
max_len=self.max_len, temperature=self.temperature,
top_k=self.top_k, top_p=self.top_p,
)
def evaluate(
self,
dataloader: DataLoader,
src_texts: Optional[list[str]] = None,
ref_texts: Optional[list[str]] = None,
) -> dict:
"""
在给定数据上进行评估。
TODO [Person D]: 实现以下逻辑:
1. model.eval()
2. 遍历 dataloader,使用选定的解码策略生成翻译
3. 将生成的 token ids 解码为文本
4. 调用 compute_all_metrics 计算指标
5. 返回评估结果 dict
"""
self.model.eval()
device = next(self.model.parameters()).device
hypotheses = []
with torch.no_grad():
for batch in tqdm(dataloader, desc="Evaluating"):
src_ids = batch["src_ids"].to(device)
src_padding_mask = batch.get("src_padding_mask")
if src_padding_mask is None:
src_padding_mask = src_ids.eq(self.pad_id)
else:
src_padding_mask = src_padding_mask.to(device)
output_ids = self.decode_fn(src_ids, src_padding_mask)
for i in range(output_ids.size(0)):
text = self.tokenizer.decode(
output_ids[i].tolist(), skip_special_tokens=True
)
hypotheses.append(text)
results = compute_all_metrics(
sources=src_texts,
hypotheses=hypotheses,
references=ref_texts,
metrics=self.metrics,
)
return results
def translate(self, texts: list[str]) -> list[str]:
"""
翻译一批文本。
TODO [Person D]:
1. tokenize 输入文本
2. 调用解码函数生成翻译
3. 解码为文本
4. 返回翻译结果列表
"""
self.model.eval()
device = next(self.model.parameters()).device
# 1. Tokenize
encoded = [self.tokenizer.encode(t, add_special_tokens=True) for t in texts]
max_len_src = max(len(ids) for ids in encoded)
src_ids = torch.full((len(texts), max_len_src), self.pad_id, dtype=torch.long, device=device)
for i, ids in enumerate(encoded):
src_ids[i, :len(ids)] = torch.tensor(ids, dtype=torch.long)
src_padding_mask = src_ids.eq(self.pad_id)
# 2. Decode
with torch.no_grad():
output_ids = self.decode_fn(src_ids, src_padding_mask)
# 3. Convert to text
translations = []
for i in range(output_ids.size(0)):
text = self.tokenizer.decode(output_ids[i].tolist(), skip_special_tokens=True)
translations.append(text)
return translations
def translate_single(self, text: str) -> str:
"""翻译单条文本。"""
return self.translate([text])[0]
|