File size: 3,753 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
ef0a52e
c1a46f7
 
 
ef0a52e
 
 
c1a46f7
 
ef0a52e
 
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ef0a52e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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
"""
评估入口脚本

使用方式:
    # 在测试集上评估
    python scripts/evaluate.py --config configs/default_config.yaml --checkpoint checkpoints/best_model.pt

    # 指定解码策略
    python scripts/evaluate.py --checkpoint checkpoints/best_model.pt evaluation.decoding.strategy=beam_search evaluation.decoding.beam_size=10
"""

import argparse
import json
import sys
from pathlib import Path

import torch
import yaml

sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

from easytranslate.evaluation.evaluator import Evaluator


def parse_args():
    parser = argparse.ArgumentParser(description="EasyTranslate Evaluation")
    parser.add_argument("--config", type=str, default="configs/default_config.yaml")
    parser.add_argument("--checkpoint", type=str, required=True, help="模型检查点路径")
    parser.add_argument("--output", type=str, default="outputs/evaluation_results.json", help="结果保存路径")
    args, unknown = parser.parse_known_args()
    return args, unknown


def main():
    """
    评估主流程。

    TODO [Person D]:
    1. 加载配置和检查点
    2. 重建模型并加载权重
    3. 加载测试数据
    4. 构建 Evaluator
    5. 运行评估
    6. 打印和保存结果
    """
    args, cli_overrides = parse_args()

    print("=" * 60)
    print("  EasyTranslate - Evaluation")
    print("=" * 60)

    # 1. 加载配置
    with open(args.config, "r", encoding="utf-8") as f:
        config = yaml.safe_load(f)

    # 应用命令行覆盖 (格式: key.subkey=value)
    for override in cli_overrides:
        if "=" in override:
            key, value = override.split("=", 1)
            keys = key.split(".")
            d = config
            for k in keys[:-1]:
                d = d.setdefault(k, {})
            d[keys[-1]] = yaml.safe_load(value)

    # 2. 加载检查点并重建模型
    checkpoint = torch.load(args.checkpoint, map_location="cpu")
    model_config = config.get("model", {})

    from easytranslate.model.transformer import Transformer
    model = Transformer(model_config)
    model.load_state_dict(checkpoint["model_state_dict"])

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)
    model.eval()

    # 3. 加载 tokenizer 和测试数据
    from easytranslate.data.tokenizer import build_tokenizer
    from easytranslate.data.dataset import TranslationDataset
    from torch.utils.data import DataLoader

    tokenizer = build_tokenizer(config.get("data", {}).get("tokenizer", {}))

    test_dataset = TranslationDataset(
        config=config.get("data", {}),
        tokenizer=tokenizer,
        split="test",
    )
    test_loader = DataLoader(
        test_dataset,
        batch_size=config.get("evaluation", {}).get("batch_size", 32),
        shuffle=False,
    )

    # 4. 构建 Evaluator
    evaluator = Evaluator(model=model, tokenizer=tokenizer, config=config)

    # 5. 运行评估
    src_texts = [sample["src"] for sample in test_dataset.raw_data]
    ref_texts = [sample["tgt"] for sample in test_dataset.raw_data]
    results = evaluator.evaluate(test_loader, src_texts=src_texts, ref_texts=ref_texts)

    # 6. 打印和保存结果
    print("\n" + "=" * 60)
    print("  Evaluation Results")
    print("=" * 60)
    for metric, score in results.items():
        if not isinstance(score, list):
            print(f"  {metric:>10s}: {score:.4f}")

    output_path = Path(args.output)
    output_path.parent.mkdir(parents=True, exist_ok=True)
    with open(output_path, "w", encoding="utf-8") as f:
        json.dump(results, f, indent=2, ensure_ascii=False)
    print(f"\n  Results saved to: {output_path}")


if __name__ == "__main__":
    main()