| """ |
| 评估器模块 — 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. 根据策略选择解码函数 |
| """ |
| raise NotImplementedError("TODO: Person D 实现 Evaluator.__init__") |
|
|
| 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 |
| """ |
| raise NotImplementedError("TODO: Person D 实现 evaluate") |
|
|
| def translate(self, texts: list[str]) -> list[str]: |
| """ |
| 翻译一批文本。 |
| |
| TODO [Person D]: |
| 1. tokenize 输入文本 |
| 2. 调用解码函数生成翻译 |
| 3. 解码为文本 |
| 4. 返回翻译结果列表 |
| """ |
| raise NotImplementedError("TODO: Person D 实现 translate") |
|
|
| def translate_single(self, text: str) -> str: |
| """翻译单条文本。""" |
| return self.translate([text])[0] |
|
|