UCAS-EasyTranslate / scripts /evaluate.py
jiaoruotong's picture
[Person D] Implement evaluation module: metrics, decoding, evaluator, scripts
ef0a52e verified
Raw
History Blame
3.75 kB
"""
评估入口脚本
使用方式:
# 在测试集上评估
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()