CoolFace
Apppublic

Spico/Mirror

sourceHugging Faceapache-2.0updated 7d agoView on Hugging Face
4likes
metric.py556 linesDownload Raw Back to src
1from collections import defaultdict2from typing import Tuple3 4from rex.metrics import calc_p_r_f1_from_tp_fp_fn, safe_division5from rex.metrics.base import MetricBase6from rex.metrics.tagging import tagging_prf17from rex.utils.batch import decompose_batch_into_instances8from rex.utils.iteration import windowed_queue_iter9from rex.utils.random import generate_random_string_with_datetime10from sklearn.metrics import accuracy_score, matthews_corrcoef11 12 13class MrcNERMetric(MetricBase):14    def get_instances_from_batch(self, raw_batch: dict, out_batch: dict) -> Tuple:15        gold_instances = []16        pred_instances = []17 18        batch_gold = decompose_batch_into_instances(raw_batch)19        assert len(batch_gold) == len(out_batch["pred"])20 21        for i, gold in enumerate(batch_gold):22            gold_instances.append(23                {24                    "id": gold["id"],25                    "ents": {(gold["ent_type"], gent) for gent in gold["gold_ents"]},26                }27            )28            pred_instances.append(29                {30                    "id": gold["id"],31                    "ents": {(gold["ent_type"], pent) for pent in out_batch["pred"][i]},32                }33            )34 35        return gold_instances, pred_instances36 37    def calculate_scores(self, golds: list, preds: list) -> dict:38        id2gold = defaultdict(set)39        id2pred = defaultdict(set)40        # aggregate all ents with diff queries before evaluating41        for gold in golds:42            id2gold[gold["id"]].update(gold["ents"])43        for pred in preds:44            id2pred[pred["id"]].update(pred["ents"])45        assert len(id2gold) == len(id2pred)46 47        gold_ents = []48        pred_ents = []49        for _id in id2gold:50            gold_ents.append(id2gold[_id])51            pred_ents.append(id2pred[_id])52 53        return tagging_prf1(gold_ents, pred_ents, type_idx=0)54 55 56class MrcSpanMetric(MetricBase):57    def get_instances_from_batch(self, raw_batch: dict, out_batch: dict) -> Tuple:58        gold_instances = []59        pred_instances = []60 61        batch_gold = decompose_batch_into_instances(raw_batch)62        assert len(batch_gold) == len(out_batch["pred"])63 64        for i, gold in enumerate(batch_gold):65            gold_instances.append(66                {67                    "id": gold["id"],68                    "spans": set(tuple(span) for span in gold["gold_spans"]),69                }70            )71            pred_instances.append(72                {73                    "id": gold["id"],74                    "spans": set(out_batch["pred"][i]),75                }76            )77 78        return gold_instances, pred_instances79 80    def calculate_scores(self, golds: list, preds: list) -> dict:81        id2gold = defaultdict(set)82        id2pred = defaultdict(set)83        # aggregate all ents with diff queries before evaluating84        for gold in golds:85            id2gold[gold["id"]].update(gold["spans"])86        for pred in preds:87            id2pred[pred["id"]].update(pred["spans"])88        assert len(id2gold) == len(id2pred)89 90        gold_spans = []91        pred_spans = []92        for _id in id2gold:93            gold_spans.append(id2gold[_id])94            pred_spans.append(id2pred[_id])95 96        return tagging_prf1(gold_spans, pred_spans, type_idx=None)97 98 99def calc_char_event(golds, preds):100    """101    Calculate char-level event argument scores102 103    References:104        - https://aistudio.baidu.com/aistudio/competition/detail/46/0/submit-result105 106    Args:107        golds: a list of gold answers (a list of `event_list`), len=#data,108            format is a list of `event_list`109        preds: a list of pred answers, len=#data110    """111 112    def _match_arg_char_f1(gold_arg, pred_args):113        gtype, grole, gstring = gold_arg114        gchars = set(gstring)115        garg_len = len(gchars)116        cands = []117        for parg in pred_args:118            if parg[0] == gtype and parg[1] == grole:119                pchars = set(str(parg[-1]))120                parg_len = len(pchars)121                pmatch = len(pchars & gchars)122                p = safe_division(pmatch, parg_len)123                r = safe_division(pmatch, garg_len)124                f1 = safe_division(2 * p * r, p + r)125                cands.append(f1)126        if len(cands) > 0:127            f1 = sorted(cands)[-1]128            return f1129        else:130            return 0.0131 132    pscore = num_gargs = num_pargs = 0133    for _golds, _preds in zip(golds, preds):134        # _golds and _preds pair in one data instance135        gold_args = []136        pred_args = []137        for gold in _golds:138            for arg in gold.get("arguments", []):139                gold_args.append(140                    (gold.get("event_type"), arg.get("role"), arg.get("argument"))141                )142        for pred in _preds:143            for arg in pred.get("arguments", []):144                pred_args.append(145                    (pred.get("event_type"), arg.get("role"), arg.get("argument"))146                )147 148        num_gargs += len(gold_args)149        num_pargs += len(pred_args)150        for gold_arg in gold_args:151            pscore += _match_arg_char_f1(gold_arg, pred_args)152 153    p = safe_division(pscore, num_pargs)154    r = safe_division(pscore, num_gargs)155    f1 = safe_division(2 * p * r, p + r)156    return {157        "p": p,158        "r": r,159        "f1": f1,160        "pscore": pscore,161        "num_pargs": num_pargs,162        "num_gargs": num_gargs,163    }164 165 166def calc_trigger_identification_metrics(golds, preds):167    tp = fp = fn = 0168    for _golds, _preds in zip(golds, preds):169        gold_triggers = {gold["trigger"] for gold in _golds}170        pred_triggers = {pred["trigger"] for pred in _preds}171        tp += len(gold_triggers & pred_triggers)172        fp += len(pred_triggers - gold_triggers)173        fn += len(gold_triggers - pred_triggers)174    metrics = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)175    return metrics176 177 178def calc_trigger_classification_metrics(golds, preds):179    tp = fp = fn = 0180    for _golds, _preds in zip(golds, preds):181        gold_tgg_cls = {(gold["trigger"], gold["event_type"]) for gold in _golds}182        pred_tgg_cls = {(pred["trigger"], pred["event_type"]) for pred in _preds}183        tp += len(gold_tgg_cls & pred_tgg_cls)184        fp += len(pred_tgg_cls - gold_tgg_cls)185        fn += len(gold_tgg_cls - pred_tgg_cls)186    metrics = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)187    return metrics188 189 190def calc_arg_identification_metrics(golds, preds):191    """Calculate argument identification metrics192 193    Notice:194        An entity could take different roles in an event,195            so the base number must be calculated by196            (arg, event type, pos, role)197    """198    tp = fp = fn = 0199    for _golds, _preds in zip(golds, preds):200        gold_args = set()201        pred_args = set()202        for gold in _golds:203            _args = {204                (arg["role"], arg["argument"], gold["event_type"])205                for arg in gold["arguments"]206            }207            gold_args.update(_args)208        for pred in _preds:209            _args = {210                (arg["role"], arg["argument"], pred["event_type"])211                for arg in pred["arguments"]212            }213            pred_args.update(_args)214        # logic derived from OneIE215        _tp = 0216        _tp_fp = len(pred_args)217        _tp_fn = len(gold_args)218        _gold_args_wo_role = {_ga[1:] for _ga in gold_args}219        for pred_arg in pred_args:220            if pred_arg[1:] in _gold_args_wo_role:221                _tp += 1222        tp += _tp223        fp += _tp_fp - _tp224        fn += _tp_fn - _tp225    metrics = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)226    return metrics227 228 229def calc_arg_classification_metrics(golds, preds):230    tp = fp = fn = 0231    for _golds, _preds in zip(golds, preds):232        gold_arg_cls = set()233        pred_arg_cls = set()234        for gold in _golds:235            _args = {236                (arg["argument"], arg["role"], gold["event_type"])237                for arg in gold["arguments"]238            }239            gold_arg_cls.update(_args)240        for pred in _preds:241            _args = {242                (arg["argument"], arg["role"], pred["event_type"])243                for arg in pred["arguments"]244            }245            pred_arg_cls.update(_args)246        tp += len(gold_arg_cls & pred_arg_cls)247        fp += len(pred_arg_cls - gold_arg_cls)248        fn += len(gold_arg_cls - pred_arg_cls)249    metrics = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)250    return metrics251 252 253def calc_ent(golds, preds):254    """255    Args:256        golds, preds: [(type, index list), ...]257    """258    res = tagging_prf1(golds, preds, type_idx=0)259    return res260 261 262def calc_rel(golds, preds):263    gold_ents = []264    pred_ents = []265    for gold, pred in zip(golds, preds):266        gold_ins_ents = []267        for t in gold:268            gold_ins_ents.extend(t[1:])269        gold_ents.append(gold_ins_ents)270        pred_ins_ents = []271        for t in pred:272            pred_ins_ents.extend(t[1:])273        pred_ents.append(pred_ins_ents)274 275    metrics = {276        "ent": tagging_prf1(gold_ents, pred_ents, type_idx=None),277        "rel": tagging_prf1(golds, preds, type_idx=None),278    }279    return metrics280 281 282def calc_cls(golds, preds):283    metrics = {284        "mcc": -1,285        "acc": -1,286        "mf1": tagging_prf1(golds, preds, type_idx=None),287    }288    y_true = []289    y_pred = []290    for gold, pred in zip(golds, preds):291        y_true.append(" ".join(sorted(gold)))292        y_pred.append(" ".join(sorted(pred)))293    if y_true and y_pred:294        metrics["acc"] = accuracy_score(y_true, y_pred)295    else:296        metrics["acc"] = 0.0297    metrics["mcc"] = matthews_corrcoef(y_true, y_pred)298    return metrics299 300 301def calc_span(golds, preds, mode="span"):302    def _get_tokens(spans: list[tuple[tuple[int]]]) -> list[int]:303        tokens = []304        for span in spans:305            for part in span:306                _toks = []307                if len(part) == 1:308                    _toks = [part[0]]309                elif len(part) > 1:310                    if mode == "w2":311                        _toks = [*part]312                    elif mode == "span":313                        _toks = [*range(part[0], part[1] + 1)]314                    else:315                        raise ValueError316                tokens.extend(_toks)317        return tokens318 319    metrics = {320        "em": -1,321        "f1": None,322    }323    acc_num = 0324    tp = fp = fn = 0325    for gold, pred in zip(golds, preds):326        if gold == pred:327            acc_num += 1328        gold_tokens = _get_tokens(gold)329        pred_tokens = _get_tokens(pred)330        tp += len(set(gold_tokens) & set(pred_tokens))331        fp += len(set(pred_tokens) - set(gold_tokens))332        fn += len(set(gold_tokens) - set(pred_tokens))333    if len(golds) > 0:334        metrics["em"] = acc_num / len(golds)335    else:336        metrics["em"] = 0.0337    metrics["f1"] = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)338    return metrics339 340 341class MultiPartSpanMetric(MetricBase):342    def _encode_span_to_label_dict(self, span_to_label: dict) -> list:343        span_to_label_list = []344        for key, val in span_to_label.items():345            span_to_label_list.append({"key": key, "val": val})346        return span_to_label_list347 348    def _decode_span_to_label(self, span_to_label_list: list) -> dict:349        span_to_label = {}350        for content in span_to_label_list:351            span_to_label[tuple(content["key"])] = content["val"]352        return span_to_label353 354    def get_instances_from_batch(self, raw_batch: dict, out_batch: dict) -> Tuple:355        gold_instances = []356        pred_instances = []357 358        batch_gold = decompose_batch_into_instances(raw_batch)359        assert len(batch_gold) == len(out_batch["pred"])360 361        for i, gold in enumerate(batch_gold):362            ins_id = gold["raw"].get("id", generate_random_string_with_datetime())363            # encode to list to make the span_to_label dict json-serializable364            # where the original dict key is a tuple365            span_to_label_list = self._encode_span_to_label_dict(gold["span_to_label"])366            gold["span_to_label"] = span_to_label_list367            gold_instances.append(368                {369                    "id": ins_id,370                    "span_to_label_list": span_to_label_list,371                    "raw_gold_content": gold,372                    "spans": set(373                        tuple(multi_part_span) for multi_part_span in gold["spans"]374                    ),375                }376            )377            pred_instances.append(378                {379                    "id": ins_id,380                    "spans": set(381                        tuple(multi_part_span)382                        for multi_part_span in out_batch["pred"][i]383                    ),384                }385            )386 387        return gold_instances, pred_instances388 389    def calculate_scores(self, golds: list, preds: list) -> dict:390        # for general purpose evaluation391        general_gold_spans, general_pred_spans = [], []392        # cls task393        gold_cls_list, pred_cls_list = [], []394        # ent task395        gold_ent_list, pred_ent_list = [], []396        # rel task397        gold_rel_list, pred_rel_list = [], []398        # event task399        gold_event_list, pred_event_list = [], []400        # span task401        gold_span_list, pred_span_list = [], []402        # discon ent task403        gold_discon_ent_list, pred_discon_ent_list = [], []404        # hyper rel task405        gold_hyper_rel_list, pred_hyper_rel_list = [], []406 407        for gold, pred in zip(golds, preds):408            general_gold_spans.append(gold["spans"])409            general_pred_spans.append(pred["spans"])410            span_to_label = self._decode_span_to_label(gold["span_to_label_list"])411            gold_clses, pred_clses = [], []412            gold_ents, pred_ents = [], []413            gold_rels, pred_rels = [], []414            gold_trigger_to_event = defaultdict(415                lambda: {"event_type": "", "arguments": []}416            )417            pred_trigger_to_event = defaultdict(418                lambda: {"event_type": "", "arguments": []}419            )420            gold_events, pred_events = [], []421            gold_spans, pred_spans = [], []422            gold_discon_ents, pred_discon_ents = [], []423            gold_hyper_rels, pred_hyper_rels = [], []424 425            raw_schema = gold["raw_gold_content"]["raw"]["schema"]426            for span in gold["spans"]:427                if span[0] in span_to_label:428                    label = span_to_label[span[0]]429                    if label["task"] == "cls" and len(span) == 1:430                        gold_clses.append(label["string"])431                    elif label["task"] == "ent" and len(span) == 2:432                        gold_ents.append((label["string"], *span[1:]))433                    elif label["task"] == "rel" and len(span) == 3:434                        gold_rels.append((label["string"], *span[1:]))435                    elif label["task"] == "event":436                        if label["type"] == "lm" and len(span) == 2:437                            gold_trigger_to_event[span[1]]["event_type"] = label["string"]  # fmt: skip438                        elif label["type"] == "lr" and len(span) == 3:439                            gold_trigger_to_event[span[1]]["arguments"].append(440                                {"argument": span[2], "role": label["string"]}441                            )442                    elif label["task"] == "discontinuous_ent" and len(span) > 1:443                        gold_discon_ents.append((label["string"], *span[1:]))444                    elif label["task"] == "hyper_rel" and len(span) == 5 and span[3] in span_to_label:  # fmt: skip445                        q_label = span_to_label[span[3]]446                        gold_hyper_rels.append((label["string"], span[1], span[2], q_label["string"], span[4]))  # fmt: skip447                else:448                    # span task has no labels449                    gold_spans.append(tuple(span))450            for trigger, item in gold_trigger_to_event.items():451                legal_roles = raw_schema["event"][item["event_type"]]452                gold_events.append(453                    {454                        "trigger": trigger,455                        "event_type": item["event_type"],456                        "arguments": [457                            arg458                            for arg in filter(459                                lambda arg: arg["role"] in legal_roles,460                                item["arguments"],461                            )462                        ],463                    }464                )465 466            for span in pred["spans"]:467                if span[0] in span_to_label:468                    label = span_to_label[span[0]]469                    if label["task"] == "cls" and len(span) == 1:470                        pred_clses.append(label["string"])471                    elif label["task"] == "ent" and len(span) == 2:472                        pred_ents.append((label["string"], *span[1:]))473                    elif label["task"] == "rel" and len(span) == 3:474                        pred_rels.append((label["string"], *span[1:]))475                    elif label["task"] == "event":476                        if label["type"] == "lm" and len(span) == 2:477                            pred_trigger_to_event[span[1]]["event_type"] = label["string"]  # fmt: skip478                        elif label["type"] == "lr" and len(span) == 3:479                            pred_trigger_to_event[span[1]]["arguments"].append(480                                {"argument": span[2], "role": label["string"]}481                            )482                    elif label["task"] == "discontinuous_ent" and len(span) > 1:483                        pred_discon_ents.append((label["string"], *span[1:]))484                    elif label["task"] == "hyper_rel" and len(span) == 5 and span[3] in span_to_label:  # fmt: skip485                        q_label = span_to_label[span[3]]486                        pred_hyper_rels.append((label["string"], span[1], span[2], q_label["string"], span[4]))  # fmt: skip487                else:488                    # span task has no labels489                    pred_spans.append(tuple(span))490            for trigger, item in pred_trigger_to_event.items():491                if item["event_type"] not in raw_schema["event"]:492                    continue493                legal_roles = raw_schema["event"][item["event_type"]]494                pred_events.append(495                    {496                        "trigger": trigger,497                        "event_type": item["event_type"],498                        "arguments": [499                            arg500                            for arg in filter(501                                lambda arg: arg["role"] in legal_roles,502                                item["arguments"],503                            )504                        ],505                    }506                )507 508            gold_cls_list.append(gold_clses)509            pred_cls_list.append(pred_clses)510            gold_ent_list.append(gold_ents)511            pred_ent_list.append(pred_ents)512            gold_rel_list.append(gold_rels)513            pred_rel_list.append(pred_rels)514            gold_event_list.append(gold_events)515            pred_event_list.append(pred_events)516            gold_span_list.append(gold_spans)517            pred_span_list.append(pred_spans)518            gold_discon_ent_list.append(gold_discon_ents)519            pred_discon_ent_list.append(pred_discon_ents)520            gold_hyper_rel_list.append(gold_hyper_rels)521            pred_hyper_rel_list.append(pred_hyper_rels)522 523        metrics = {524            "general_spans": tagging_prf1(525                general_gold_spans, general_pred_spans, type_idx=None526            ),527            "cls": calc_cls(gold_cls_list, pred_cls_list),528            "ent": calc_ent(gold_ent_list, pred_ent_list),529            "rel": calc_rel(gold_rel_list, pred_rel_list),530            "event": {531                "trigger_id": calc_trigger_identification_metrics(532                    gold_event_list, pred_event_list533                ),534                "trigger_cls": calc_trigger_classification_metrics(535                    gold_event_list, pred_event_list536                ),537                "arg_id": calc_arg_identification_metrics(538                    gold_event_list, pred_event_list539                ),540                "arg_cls": calc_arg_classification_metrics(541                    gold_event_list, pred_event_list542                ),543                "char_event": calc_char_event(gold_event_list, pred_event_list),544            },545            "discontinuous_ent": tagging_prf1(546                gold_discon_ent_list, pred_discon_ent_list, type_idx=None547            ),548            "hyper_rel": tagging_prf1(549                gold_hyper_rel_list, pred_hyper_rel_list, type_idx=None550            ),551            # "span": tagging_prf1(gold_span_list, pred_span_list, type_idx=None),552            "span": calc_span(gold_span_list, pred_span_list),553        }554 555        return metrics556