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()