File size: 6,365 Bytes
c1a46f7 d572bbd c1a46f7 d572bbd 5672463 c1a46f7 d572bbd c1a46f7 d572bbd c1a46f7 5672463 d572bbd 5672463 c1a46f7 d572bbd 5672463 c1a46f7 d572bbd | 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:
"""翻译模型评估器。"""
def __init__(self, model: nn.Module, tokenizer, config: dict):
"""初始化评估器。"""
self.model = model
self.tokenizer = tokenizer
self.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 _validate_batch(self, batch: dict) -> bool:
"""验证批处理格式是否符合预期。"""
required_keys = ["src_ids"]
for key in required_keys:
if key not in batch:
logger.error(f"Batch is missing required key: {key}")
return False
if not isinstance(batch["src_ids"], torch.Tensor):
logger.error(f"src_ids must be a torch.Tensor, got {type(batch['src_ids'])}")
return False
return True
def evaluate(
self,
dataloader: DataLoader,
src_texts: Optional[list[str]] = None,
ref_texts: Optional[list[str]] = None,
) -> dict:
"""
在给定数据上进行评估。
Args:
dataloader: 测试数据集的 DataLoader
src_texts: 源文本列表 (用于某些指标计算)
ref_texts: 参考译文列表
Returns:
dict: 包含所有评估指标的字典
Raises:
ValueError: 如果批处理格式不正确
"""
self.model.eval()
device = next(self.model.parameters()).device
hypotheses = []
with torch.no_grad():
for batch in tqdm(dataloader, desc="Evaluating"):
# 验证批处理格式
if not self._validate_batch(batch):
raise ValueError("Invalid batch format. Expected 'src_ids' key with torch.Tensor value.")
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]:
"""翻译一批文本。"""
self.model.eval()
device = next(self.model.parameters()).device
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)
with torch.no_grad():
output_ids = self.decode_fn(src_ids, src_padding_mask)
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] |