Spico/Mirror
4
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 