File size: 5,179 Bytes
6dc7c27 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | import tempfile
from typing import Tuple, Any, Dict, Optional, List
import hydra
import torch
from omegaconf import omegaconf, DictConfig
from src.consec_dataset import ConsecSample
from src.dependency_finder import EmptyDependencyFinder
from src.pl_modules import ConsecPLModule
from src.scripts.model.continuous_predict import Predictor
from src.sense_inventories import SenseInventory, WordNetSenseInventory
from src.utils.commons import execute_bash_command
from src.utils.hydra import fix
from src.utils.wsd import expand_raganato_path
def framework_evaluate(framework_folder: str, gold_file_path: str, pred_file_path: str) -> Tuple[float, float, float]:
scorer_folder = f"{framework_folder}/Evaluation_Datasets"
command_output = execute_bash_command(
f"[ ! -e {scorer_folder}/Scorer.class ] && javac -d {scorer_folder} {scorer_folder}/Scorer.java; java -cp {scorer_folder} Scorer {gold_file_path} {pred_file_path}"
)
command_output = command_output.split("\n")
p, r, f1 = [float(command_output[i].split("=")[-1].strip()[:-1]) for i in range(3)]
return p, r, f1
def sample_prediction2sense(sample: ConsecSample, prediction: int, sense_inventory: SenseInventory) -> str:
sample_senses = sense_inventory.get_possible_senses(
sample.disambiguation_instance.lemma, sample.disambiguation_instance.pos
)
sample_definitions = [sense_inventory.get_definition(s) for s in sample_senses]
for s, d in zip(sample_senses, sample_definitions):
if d == sample.candidate_definitions[prediction].text:
return s
raise ValueError
def raganato_evaluate(
raganato_path: str,
wsd_framework_dir: str,
module: ConsecPLModule,
predictor: Predictor,
wordnet_sense_inventory: WordNetSenseInventory,
samples_generator: DictConfig,
prediction_params: Dict[Any, Any],
fine_grained_evals: Optional[List[str]] = None,
reporting_folder: Optional[str] = None,
) -> Tuple[float, float, float, Optional[List[Tuple[str, float, float, float]]]]:
# load tokenizer
tokenizer = hydra.utils.instantiate(module.hparams.tokenizer.consec_tokenizer)
# instantiate samples
consec_samples = list(hydra.utils.instantiate(samples_generator, dependency_finder=EmptyDependencyFinder())())
# predict
disambiguated_samples = predictor.predict(
consec_samples,
already_kwown_predictions=None,
reporting_folder=reporting_folder,
**dict(module=module, tokenizer=tokenizer, **prediction_params),
)
# write predictions and evaluate
with tempfile.TemporaryDirectory() as tmp_dir:
# write predictions to tmp file
with open(f"{tmp_dir}/predictions.gold.key.txt", "w") as f:
for sample, idx in disambiguated_samples:
f.write(f"{sample.sample_id} {sample_prediction2sense(sample, idx, wordnet_sense_inventory)}\n")
# compute metrics
p, r, f1 = framework_evaluate(
wsd_framework_dir,
gold_file_path=expand_raganato_path(raganato_path)[1],
pred_file_path=f"{tmp_dir}/predictions.gold.key.txt",
)
# fine grained eval
fge_scores = None
if fine_grained_evals is not None:
fge_scores = []
for fge in fine_grained_evals:
_p, _r, _f1 = framework_evaluate(
wsd_framework_dir,
gold_file_path=expand_raganato_path(fge)[1],
pred_file_path=f"{tmp_dir}/predictions.gold.key.txt",
)
fge_scores.append((fge, _p, _r, _f1))
return p, r, f1, fge_scores
@hydra.main(config_path="../../../conf/test", config_name="raganato")
def main(conf: omegaconf.DictConfig) -> None:
fix(conf)
# load module
# todo decouple ConsecPLModule
module = ConsecPLModule.load_from_checkpoint(conf.model.model_checkpoint)
module.to(torch.device(conf.model.device if conf.model.device != -1 else "cpu"))
module.eval()
module.freeze()
module.sense_extractor.evaluation_mode = True # no loss will be computed even if labels are passed
# instantiate sense inventory
sense_inventory = hydra.utils.instantiate(conf.sense_inventory)
# instantiate predictor
predictor = hydra.utils.instantiate(conf.predictor)
# evaluate
p, r, f1, fge_scores = raganato_evaluate(
raganato_path=conf.test_raganato_path,
wsd_framework_dir=conf.wsd_framework_dir,
module=module,
predictor=predictor,
wordnet_sense_inventory=sense_inventory,
samples_generator=conf.samples_generator,
prediction_params=conf.model.prediction_params,
fine_grained_evals=conf.fine_grained_evals,
reporting_folder=".", # hydra will handle it
)
print(f"# p: {p}")
print(f"# r: {r}")
print(f"# f1: {f1}")
if fge_scores:
for fge, p, r, f1 in fge_scores:
print(f'# {fge}: ({p:.1f}, {r:.1f}, {f1:.1f})')
if __name__ == "__main__":
main()
|