""" 评估入口脚本 使用方式: # 在测试集上评估 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()