| """
|
| 评估入口脚本
|
|
|
| 使用方式:
|
| # 在测试集上评估
|
| 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
|
|
|
| try:
|
| from omegaconf import OmegaConf
|
| except ImportError:
|
| OmegaConf = None
|
| import yaml
|
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
|
|
|
| import torch
|
| from torch.utils.data import DataLoader
|
|
|
| from easytranslate.data.collator import TranslationCollator
|
| from easytranslate.data.dataset import (
|
| TranslationDataset,
|
| load_custom_dataset,
|
| load_opus_dataset,
|
| load_wmt_dataset,
|
| )
|
| from easytranslate.data.tokenizer import TokenizerWrapper, build_tokenizer
|
| from easytranslate.evaluation.evaluator import Evaluator
|
| from easytranslate.model import TransformerTranslationModel
|
| from easytranslate.model.finetune import load_pretrained_model
|
|
|
|
|
| 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 _get_config(config, *keys, default=None):
|
| value = config
|
| for key in keys:
|
| if isinstance(value, dict):
|
| value = value.get(key, default)
|
| else:
|
| value = getattr(value, key, default)
|
| if value is default:
|
| break
|
| return value
|
|
|
|
|
| def _load_config(path, cli_overrides=None):
|
| if OmegaConf is not None:
|
| config = OmegaConf.load(path)
|
| if cli_overrides:
|
| config = OmegaConf.merge(config, OmegaConf.from_cli(cli_overrides))
|
| return config
|
|
|
| with open(path, "r", encoding="utf-8") as fin:
|
| config = yaml.safe_load(fin)
|
| if cli_overrides:
|
| print("Warning: OmegaConf is not installed; CLI overrides are ignored.")
|
| return config
|
|
|
|
|
| def load_test_split(config):
|
| dataset_name = _get_config(config, "data", "dataset_name")
|
|
|
| if dataset_name == "wmt":
|
| try:
|
| dataset = load_wmt_dataset(
|
| year=_get_config(config, "data", "wmt", "year"),
|
| language_pair=_get_config(config, "data", "wmt", "language_pair"),
|
| split="test",
|
| )
|
| except Exception:
|
| dataset = load_wmt_dataset(
|
| year=_get_config(config, "data", "wmt", "year"),
|
| language_pair=_get_config(config, "data", "wmt", "language_pair"),
|
| split="validation",
|
| )
|
| return list(dataset["src"]), list(dataset["tgt"])
|
|
|
| if dataset_name == "opus":
|
| try:
|
| dataset = load_opus_dataset(
|
| subset=_get_config(config, "data", "opus", "subset"),
|
| split="test",
|
| )
|
| except Exception:
|
| dataset = load_opus_dataset(
|
| subset=_get_config(config, "data", "opus", "subset"),
|
| split="validation",
|
| )
|
| return list(dataset["src"]), list(dataset["tgt"])
|
|
|
| if dataset_name == "custom":
|
| data = load_custom_dataset(
|
| train_src=_get_config(config, "data", "custom", "train_src"),
|
| train_tgt=_get_config(config, "data", "custom", "train_tgt"),
|
| val_src=_get_config(config, "data", "custom", "val_src"),
|
| val_tgt=_get_config(config, "data", "custom", "val_tgt"),
|
| test_src=_get_config(config, "data", "custom", "test_src"),
|
| test_tgt=_get_config(config, "data", "custom", "test_tgt"),
|
| preprocessing_config=_get_config(config, "data", "preprocessing"),
|
| )
|
| if "test" not in data:
|
| raise ValueError("Custom dataset missing test split")
|
| return data["test"]["src"], data["test"]["tgt"]
|
|
|
| raise ValueError(f"Unsupported dataset_name: {dataset_name}")
|
|
|
|
|
| def _get_tokenizer_train_texts(config, allow_auto: bool = False) -> list[str] | None:
|
| dataset_name = _get_config(config, "data", "dataset_name")
|
| if dataset_name == "custom":
|
| custom = _get_config(config, "data", "custom") or {}
|
| data = load_custom_dataset(
|
| train_src=custom.get("train_src"),
|
| train_tgt=custom.get("train_tgt"),
|
| val_src=custom.get("val_src"),
|
| val_tgt=custom.get("val_tgt"),
|
| test_src=custom.get("test_src"),
|
| test_tgt=custom.get("test_tgt"),
|
| preprocessing_config=_get_config(config, "data", "preprocessing"),
|
| )
|
| return list(data["train"]["src"]) + list(data["train"]["tgt"])
|
|
|
| if not allow_auto:
|
| return None
|
|
|
| if dataset_name == "wmt":
|
| try:
|
| dataset = load_wmt_dataset(
|
| year=_get_config(config, "data", "wmt", "year"),
|
| language_pair=_get_config(config, "data", "wmt", "language_pair"),
|
| split="train",
|
| )
|
| except Exception:
|
| dataset = load_wmt_dataset(
|
| year=_get_config(config, "data", "wmt", "year"),
|
| language_pair=_get_config(config, "data", "wmt", "language_pair"),
|
| split="validation",
|
| )
|
| return list(dataset["src"]) + list(dataset["tgt"])
|
|
|
| if dataset_name == "opus":
|
| try:
|
| dataset = load_opus_dataset(
|
| subset=_get_config(config, "data", "opus", "subset"),
|
| split="train",
|
| )
|
| except Exception:
|
| dataset = load_opus_dataset(
|
| subset=_get_config(config, "data", "opus", "subset"),
|
| split="validation",
|
| )
|
| return list(dataset["src"]) + list(dataset["tgt"])
|
|
|
| return None
|
|
|
|
|
| def build_model_and_tokenizer(config, device):
|
| model_type = _get_config(config, "model", "type")
|
| if model_type == "transformer_scratch":
|
| tokenizer_config = _get_config(config, "tokenizer") or {}
|
| tokenizer_path = tokenizer_config.get("path") or tokenizer_config.get("tokenizer_path")
|
| tokenizer_type = tokenizer_config.get("type", "bpe")
|
| auto_train = bool(tokenizer_config.get("auto_train", False))
|
| if tokenizer_type in {"bpe", "sentencepiece"} and not tokenizer_path:
|
| train_texts = _get_tokenizer_train_texts(config, allow_auto=auto_train)
|
| if train_texts is None:
|
| raise ValueError(
|
| "BPE tokenizer requires tokenizer.path or a local custom dataset with train texts. "
|
| "Automatic WMT/OPUS download is disabled by default. "
|
| "Set tokenizer.auto_train=true to enable it, or provide tokenizer.path/pretrained tokenizer."
|
| )
|
| tokenizer = build_tokenizer(tokenizer_config, train_texts=train_texts)
|
| else:
|
| tokenizer = build_tokenizer(tokenizer_config)
|
| model = TransformerTranslationModel(
|
| src_vocab_size=tokenizer.vocab_size,
|
| tgt_vocab_size=tokenizer.vocab_size,
|
| d_model=_get_config(config, "model", "transformer", "d_model"),
|
| nhead=_get_config(config, "model", "transformer", "nhead"),
|
| num_encoder_layers=_get_config(config, "model", "transformer", "num_encoder_layers"),
|
| num_decoder_layers=_get_config(config, "model", "transformer", "num_decoder_layers"),
|
| dim_feedforward=_get_config(config, "model", "transformer", "dim_feedforward"),
|
| dropout=_get_config(config, "model", "transformer", "dropout"),
|
| activation=_get_config(config, "model", "transformer", "activation"),
|
| max_seq_len=_get_config(config, "model", "transformer", "max_seq_len"),
|
| use_flash_attention=_get_config(config, "model", "transformer", "use_flash_attention"),
|
| use_rotary_embedding=_get_config(config, "model", "transformer", "use_rotary_embedding"),
|
| pre_norm=_get_config(config, "model", "transformer", "pre_norm"),
|
| pad_id=tokenizer.pad_token_id,
|
| )
|
| return model.to(device), tokenizer
|
|
|
| model, hf_tokenizer = load_pretrained_model(
|
| config.model.pretrained.model_name,
|
| config.model.pretrained.src_lang,
|
| config.model.pretrained.tgt_lang,
|
| device=str(device),
|
| )
|
| tokenizer = TokenizerWrapper(
|
| hf_tokenizer,
|
| pad_token=getattr(hf_tokenizer, "pad_token", "<pad>"),
|
| unk_token=getattr(hf_tokenizer, "unk_token", "<unk>"),
|
| bos_token=getattr(hf_tokenizer, "bos_token", "<s>"),
|
| eos_token=getattr(hf_tokenizer, "eos_token", "</s>"),
|
| )
|
| return model, tokenizer
|
|
|
|
|
| def load_checkpoint(model, checkpoint_path, device):
|
| checkpoint = torch.load(checkpoint_path, map_location=device)
|
| if isinstance(checkpoint, dict):
|
| if "model_state_dict" in checkpoint:
|
| model.load_state_dict(checkpoint["model_state_dict"])
|
| elif "state_dict" in checkpoint:
|
| model.load_state_dict(checkpoint["state_dict"])
|
| else:
|
| try:
|
| model.load_state_dict(checkpoint)
|
| except Exception as exc:
|
| raise ValueError("Checkpoint does not contain a valid model state dict") from exc
|
| else:
|
| raise ValueError("Unsupported checkpoint format")
|
| return model
|
|
|
|
|
| def main():
|
| args, cli_overrides = parse_args()
|
|
|
| print("=" * 60)
|
| print(" EasyTranslate - Evaluation")
|
| print("=" * 60)
|
|
|
| config = _load_config(args.config, cli_overrides)
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| model, tokenizer = build_model_and_tokenizer(config, device)
|
| model = load_checkpoint(model, args.checkpoint, device)
|
|
|
| src_texts, tgt_texts = load_test_split(config)
|
|
|
| dataset = TranslationDataset(
|
| src_texts=src_texts,
|
| tgt_texts=tgt_texts,
|
| tokenizer=tokenizer,
|
| max_src_len=_get_config(config, "data", "preprocessing", "max_src_len"),
|
| max_tgt_len=_get_config(config, "data", "preprocessing", "max_tgt_len"),
|
| )
|
| collator = TranslationCollator(pad_token_id=tokenizer.pad_token_id)
|
| dataloader = DataLoader(
|
| dataset,
|
| batch_size=config.data.dataloader.batch_size,
|
| shuffle=False,
|
| num_workers=int(config.data.dataloader.num_workers),
|
| pin_memory=bool(config.data.dataloader.pin_memory),
|
| collate_fn=collator,
|
| )
|
|
|
| evaluator = Evaluator(model, tokenizer, config)
|
| results = evaluator.evaluate(dataloader, src_texts=src_texts, ref_texts=tgt_texts)
|
|
|
| output_path = Path(args.output)
|
| output_path.parent.mkdir(parents=True, exist_ok=True)
|
| with output_path.open("w", encoding="utf-8") as fout:
|
| json.dump(results, fout, ensure_ascii=False, indent=2)
|
|
|
| print(json.dumps(results, ensure_ascii=False, indent=2))
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|