File size: 2,120 Bytes
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
"""
评估器模块 — 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]