""" 评估入口脚本 使用方式: # 在测试集上评估 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: # pragma: no cover 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", ""), unk_token=getattr(hf_tokenizer, "unk_token", ""), bos_token=getattr(hf_tokenizer, "bos_token", ""), eos_token=getattr(hf_tokenizer, "eos_token", ""), ) 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()