Seeds / src /04_evaluate.py
BeaMorais's picture
Update src/04_evaluate.py
e243b72 verified
Raw History Blame Contribute Delete
10.4 kB
import json
import os
import subprocess
import sys
import evaluate
import torch
from datasets import Dataset
from transformers import (
AutoTokenizer,
AutoModelForSeq2SeqLM
)
# ============================================================
# CONFIGURAÇÕES
# ============================================================
MODEL_PATH = "models/seeds-byt5-pilot-final"
TEST_FILE = "data/test.jsonl"
OUTPUT_DIR = "outputs"
# Para a avaliação definitiva mudar para None.
MAX_TEST_EXAMPLES = 100
MAX_SOURCE_LENGTH = 128
MAX_TARGET_LENGTH = 128
BATCH_SIZE = 1
# ============================================================
# CABEÇALHO
# ============================================================
print("=" * 70)
print("SEEDS - AVALIAÇÃO DO MODELO")
print("English → Macushi")
print("=" * 70)
print()
print(
f"Modelo: {MODEL_PATH}"
)
print(
f"Arquivo de teste: {TEST_FILE}"
)
print(
f"Máximo de exemplos: "
f"{MAX_TEST_EXAMPLES}"
)
print(
f"Comprimento source: "
f"{MAX_SOURCE_LENGTH}"
)
print(
f"Comprimento target: "
f"{MAX_TARGET_LENGTH}"
)
print()
# ============================================================
# GARANTIR DATASET
# ============================================================
def test_file_exists():
return (
os.path.exists(TEST_FILE)
and os.path.getsize(TEST_FILE) > 0
)
def prepare_dataset_if_needed():
if test_file_exists():
print("✅ Dataset de teste encontrado.")
print()
return
print(
"⚠️ Dataset de teste não encontrado."
)
print()
print(
"🔄 Executando automaticamente:"
)
print(
" src/02_prepare_dataset.py"
)
print()
subprocess.run(
[
sys.executable,
"src/02_prepare_dataset.py"
],
check=True
)
if not test_file_exists():
raise RuntimeError(
"O arquivo de teste não foi "
"criado após a execução de "
"02_prepare_dataset.py."
)
print()
print(
"✅ Dataset de teste preparado."
)
print()
prepare_dataset_if_needed()
# ============================================================
# GARANTIR MODELO
# ============================================================
def model_exists():
required_files = [
"config.json"
]
if not os.path.isdir(MODEL_PATH):
return False
for filename in required_files:
if not os.path.exists(
os.path.join(
MODEL_PATH,
filename
)
):
return False
return True
def train_model_if_needed():
if model_exists():
print(
"✅ Modelo encontrado."
)
print()
return
print(
"⚠️ Modelo treinado não encontrado."
)
print()
print(
"🔄 Executando automaticamente:"
)
print(
" src/03_train.py"
)
print()
subprocess.run(
[
sys.executable,
"src/03_train.py"
],
check=True
)
if not model_exists():
raise RuntimeError(
"O modelo não foi encontrado "
"após a execução de "
"03_train.py."
)
print()
print(
"✅ Modelo preparado para avaliação."
)
print()
train_model_if_needed()
# ============================================================
# CARREGAR DATASET
# ============================================================
print("Carregando dataset de teste...")
with open(
TEST_FILE,
"r",
encoding="utf-8"
) as f:
test_data = [
json.loads(linha)
for linha in f
if linha.strip()
]
print(
f"Dataset completo: "
f"{len(test_data)} exemplos"
)
if MAX_TEST_EXAMPLES is not None:
test_data = test_data[
:MAX_TEST_EXAMPLES
]
print(
f"Dataset utilizado: "
f"{len(test_data)} exemplos"
)
print()
# ============================================================
# DATASET HUGGING FACE
# ============================================================
dataset = Dataset.from_list(
test_data
)
print("Colunas disponíveis:")
print(dataset.column_names)
print()
# ============================================================
# TOKENIZER
# ============================================================
print("Carregando tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_PATH
)
print("Tokenizer carregado.")
print()
# ============================================================
# MODELO
# ============================================================
print("Carregando modelo...")
model = AutoModelForSeq2SeqLM.from_pretrained(
MODEL_PATH
)
model.config.use_cache = True
model.eval()
device = torch.device("cpu")
model.to(device)
print(
f"Modelo carregado em: {device}"
)
print()
# ============================================================
# PREPARAR TEXTOS
# ============================================================
sources = []
references = []
for exemplo in test_data:
source = exemplo["source"]
target = exemplo["target"]
prompt = (
f"translate English to Macushi: "
f"{source}"
)
sources.append(prompt)
references.append(target)
# ============================================================
# MÉTRICAS
# ============================================================
print("Carregando métricas...")
bleu = evaluate.load(
"sacrebleu"
)
chrf = evaluate.load(
"chrf"
)
print("Métricas carregadas.")
print()
# ============================================================
# GERAÇÃO
# ============================================================
print("=" * 70)
print("GERANDO TRADUÇÕES")
print("=" * 70)
print()
predictions = []
for i in range(
0,
len(sources),
BATCH_SIZE
):
batch_sources = sources[
i:i + BATCH_SIZE
]
inputs = tokenizer(
batch_sources,
max_length=MAX_SOURCE_LENGTH,
truncation=True,
padding=True,
return_tensors="pt"
)
inputs = {
key: value.to(device)
for key, value in inputs.items()
}
with torch.no_grad():
generated = model.generate(
**inputs,
max_length=MAX_TARGET_LENGTH,
num_beams=4,
early_stopping=True
)
decoded = tokenizer.batch_decode(
generated,
skip_special_tokens=True
)
predictions.extend(
decoded
)
if (
(i + BATCH_SIZE) % 10 == 0
or
i + BATCH_SIZE >= len(sources)
):
processados = min(
i + BATCH_SIZE,
len(sources)
)
print(
f"Traduções geradas: "
f"{processados}/{len(sources)}"
)
print()
print("Geração concluída.")
print()
# ============================================================
# NORMALIZAÇÃO
# ============================================================
predictions = [
prediction.strip()
for prediction in predictions
]
references = [
reference.strip()
for reference in references
]
# ============================================================
# BLEU
# ============================================================
print("Calculando BLEU...")
bleu_result = bleu.compute(
predictions=predictions,
references=[
[reference]
for reference in references
]
)
bleu_score = bleu_result["score"]
print(
f"BLEU: {bleu_score:.4f}"
)
print()
# ============================================================
# chrF
# ============================================================
print("Calculando chrF...")
chrf_result = chrf.compute(
predictions=predictions,
references=references
)
chrf_score = chrf_result["score"]
print(
f"chrF: {chrf_score:.4f}"
)
print()
# ============================================================
# EXEMPLOS
# ============================================================
print("=" * 70)
print("EXEMPLOS DE TRADUÇÃO")
print("=" * 70)
num_examples = min(
10,
len(test_data)
)
for i in range(num_examples):
print()
print(
f"EXEMPLO {i + 1}"
)
print("-" * 70)
print("English:")
print(
test_data[i]["source"]
)
print()
print(
"Macushi - referência:"
)
print(
references[i]
)
print()
print(
"Macushi - modelo:"
)
print(
predictions[i]
)
# ============================================================
# SALVAR RESULTADOS
# ============================================================
os.makedirs(
OUTPUT_DIR,
exist_ok=True
)
results = {
"model": MODEL_PATH,
"test_examples": len(
test_data
),
"max_source_length": (
MAX_SOURCE_LENGTH
),
"max_target_length": (
MAX_TARGET_LENGTH
),
"bleu": bleu_score,
"chrf": chrf_score,
"examples": []
}
for i in range(
len(test_data)
):
results["examples"].append(
{
"book": test_data[i].get(
"book"
),
"chapter": test_data[i].get(
"chapter"
),
"verse": test_data[i].get(
"verse"
),
"source": test_data[i][
"source"
],
"reference": references[i],
"prediction": predictions[i]
}
)
output_file = os.path.join(
OUTPUT_DIR,
"pilot_evaluation.json"
)
with open(
output_file,
"w",
encoding="utf-8"
) as f:
json.dump(
results,
f,
ensure_ascii=False,
indent=2
)
# ============================================================
# RESULTADO
# ============================================================
print()
print("=" * 70)
print("RESULTADO DA AVALIAÇÃO")
print("=" * 70)
print()
print(
f"Exemplos avaliados: "
f"{len(test_data)}"
)
print(
f"BLEU: "
f"{bleu_score:.4f}"
)
print(
f"chrF: "
f"{chrf_score:.4f}"
)
print()
print(
"Relatório salvo em:"
)
print(
output_file
)
print()
print("=" * 70)
print("AVALIAÇÃO FINALIZADA")
print("=" * 70)