| """ |
| 评估入口脚本 |
| |
| 使用方式: |
| # 在测试集上评估 |
| 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) |
|
|
| |
| with open(args.config, "r", encoding="utf-8") as f: |
| config = yaml.safe_load(f) |
|
|
| |
| 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) |
|
|
| |
| 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() |
|
|
| |
| 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, |
| ) |
|
|
| |
| evaluator = Evaluator(model=model, tokenizer=tokenizer, config=config) |
|
|
| |
| 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) |
|
|
| |
| 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() |
|
|