jiaoruotong's picture
[Person D] Implement evaluation module: metrics, decoding, evaluator, scripts
ef0a52e verified
Raw
History Blame
6.11 kB
"""
评估器模块 — 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]