UCAS-EasyTranslate / scripts /evaluate.py
jiaoruotong's picture
Add Streamlit preview frontend and normalize line endings
86fe6bc verified
Raw
History Blame
11.5 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
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", "<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()