lijn14
创建工程
c1a46f7
Raw
History Blame
2.12 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. 根据策略选择解码函数
"""
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]