Spico/Mirror
4
1import random2import re3from collections import defaultdict4from typing import Iterable, Iterator, List, MutableSet, Optional, Tuple, TypeVar, Union5 6import torch7import torch.nn.functional as F8from rex.data.collate_fn import GeneralCollateFn9from rex.data.transforms.base import CachedTransformBase, CachedTransformOneBase10from rex.metrics import calc_p_r_f1_from_tp_fp_fn11from rex.utils.io import load_json12from rex.utils.iteration import windowed_queue_iter13from rex.utils.logging import logger14from transformers import AutoTokenizer15from transformers.models.bert.tokenization_bert_fast import BertTokenizerFast16from transformers.models.deberta_v2.tokenization_deberta_v2_fast import (17 DebertaV2TokenizerFast,18)19from transformers.tokenization_utils_base import BatchEncoding20 21from src.utils import (22 decode_nnw_nsw_thw_mat,23 decode_nnw_thw_mat,24 encode_nnw_nsw_thw_mat,25 encode_nnw_thw_mat,26)27 28Filled = TypeVar("Filled")29 30 31class PaddingMixin:32 max_seq_len: int33 34 def pad_seq(self, batch_seqs: Iterable[Filled], fill: Filled) -> Iterable[Filled]:35 max_len = max(len(seq) for seq in batch_seqs)36 assert max_len <= self.max_seq_len37 for i in range(len(batch_seqs)):38 batch_seqs[i] = batch_seqs[i] + [fill] * (max_len - len(batch_seqs[i]))39 return batch_seqs40 41 def pad_mat(42 self, mats: List[torch.Tensor], fill: Union[int, float]43 ) -> List[torch.Tensor]:44 max_len = max(mat.shape[0] for mat in mats)45 assert max_len <= self.max_seq_len46 for i in range(len(mats)):47 num_add = max_len - mats[i].shape[0]48 mats[i] = F.pad(49 mats[i], (0, 0, 0, num_add, 0, num_add), mode="constant", value=fill50 )51 return mats52 53 54class PointerTransformMixin:55 tokenizer: BertTokenizerFast56 max_seq_len: int57 space_token: str = "[unused1]"58 59 def build_ins(60 self,61 query_tokens: list[str],62 context_tokens: list[str],63 answer_indexes: list[list[int]],64 add_context_tokens: list[str] = None,65 ) -> Tuple:66 # -2: cls and sep67 reserved_seq_len = self.max_seq_len - 3 - len(query_tokens)68 # reserve at least 20 tokens69 if reserved_seq_len < 20:70 raise ValueError(71 f"Query {query_tokens} too long: {len(query_tokens)} "72 f"while max seq len is {self.max_seq_len}"73 )74 75 input_tokens = [self.tokenizer.cls_token]76 input_tokens += query_tokens77 input_tokens += [self.tokenizer.sep_token]78 offset = len(input_tokens)79 input_tokens += context_tokens[:reserved_seq_len]80 available_token_range = range(81 offset, offset + len(context_tokens[:reserved_seq_len])82 )83 input_tokens += [self.tokenizer.sep_token]84 85 add_context_len = 086 max_add_context_len = self.max_seq_len - len(input_tokens) - 187 add_context_flag = False88 if add_context_tokens and len(add_context_tokens) > 0:89 add_context_flag = True90 add_context_len = len(add_context_tokens[:max_add_context_len])91 input_tokens += add_context_tokens[:max_add_context_len]92 input_tokens += [self.tokenizer.sep_token]93 new_tokens = []94 for t in input_tokens:95 if len(t.strip()) > 0:96 new_tokens.append(t)97 else:98 new_tokens.append(self.space_token)99 input_tokens = new_tokens100 input_ids = self.tokenizer.convert_tokens_to_ids(input_tokens)101 102 mask = [1]103 mask += [2] * len(query_tokens)104 mask += [3]105 mask += [4] * len(context_tokens[:reserved_seq_len])106 mask += [5]107 if add_context_flag:108 mask += [6] * add_context_len109 mask += [7]110 assert len(mask) == len(input_ids) <= self.max_seq_len111 112 available_spans = [tuple(i + offset for i in index) for index in answer_indexes]113 available_spans = list(114 filter(115 lambda index: all(i in available_token_range for i in index),116 available_spans,117 )118 )119 120 token_len = len(input_ids)121 pad_len = self.max_seq_len - token_len122 input_tokens += pad_len * [self.tokenizer.pad_token]123 input_ids += pad_len * [self.tokenizer.pad_token_id]124 mask += pad_len * [0]125 126 return input_tokens, input_ids, mask, offset, available_spans127 128 def update_labels(self, data: dict) -> dict:129 bs = len(data["input_ids"])130 seq_len = self.max_seq_len131 labels = torch.zeros((bs, 2, seq_len, seq_len))132 for i, batch_spans in enumerate(data["available_spans"]):133 # offset = data["offset"][i]134 # pad_len = data["mask"].count(0)135 # token_len = seq_len - pad_len136 for span in batch_spans:137 if len(span) == 1:138 labels[i, :, span[0], span[0]] = 1139 else:140 for s, e in windowed_queue_iter(span, 2, 1, drop_last=True):141 labels[i, 0, s, e] = 1142 labels[i, 1, span[-1], span[0]] = 1143 # labels[i, :, 0:offset, :] = -100144 # labels[i, :, :, 0:offset] = -100145 # labels[i, :, :, token_len:] = -100146 # labels[i, :, token_len:, :] = -100147 data["labels"] = labels148 return data149 150 def update_consecutive_span_labels(self, data: dict) -> dict:151 bs = len(data["input_ids"])152 seq_len = self.max_seq_len153 labels = torch.zeros((bs, 1, seq_len, seq_len))154 for i, batch_spans in enumerate(data["available_spans"]):155 for span in batch_spans:156 assert span == tuple(sorted(set(span)))157 if len(span) == 1:158 labels[i, 0, span[0], span[0]] = 1159 else:160 labels[i, 0, span[0], span[-1]] = 1161 data["labels"] = labels162 return data163 164 165class CachedPointerTaggingTransform(CachedTransformBase, PointerTransformMixin):166 def __init__(167 self,168 max_seq_len: int,169 plm_dir: str,170 ent_type2query_filepath: str,171 mode: str = "w2",172 negative_sample_prob: float = 1.0,173 ) -> None:174 super().__init__()175 176 self.max_seq_len: int = max_seq_len177 self.tokenizer: BertTokenizerFast = BertTokenizerFast.from_pretrained(plm_dir)178 self.ent_type2query: dict = load_json(ent_type2query_filepath)179 self.negative_sample_prob = negative_sample_prob180 181 self.collate_fn: GeneralCollateFn = GeneralCollateFn(182 {183 "input_ids": torch.long,184 "mask": torch.long,185 "labels": torch.long,186 },187 guessing=False,188 missing_key_as_null=True,189 )190 if mode == "w2":191 self.collate_fn.update_before_tensorify = self.update_labels192 elif mode == "cons":193 self.collate_fn.update_before_tensorify = (194 self.update_consecutive_span_labels195 )196 else:197 raise ValueError(f"Mode: {mode} not recognizable")198 199 def transform(200 self,201 transform_loader: Iterator,202 dataset_name: str = None,203 **kwargs,204 ) -> Iterable:205 final_data = []206 # tp = fp = fn = 0207 for data in transform_loader:208 ent_type2ents = defaultdict(set)209 for ent in data["ents"]:210 ent_type2ents[ent["type"]].add(tuple(ent["index"]))211 for ent_type in self.ent_type2query:212 gold_ents = ent_type2ents[ent_type]213 if (214 len(gold_ents) < 1215 and dataset_name == "train"216 and random.random() > self.negative_sample_prob217 ):218 # skip negative samples219 continue220 # res = self.build_ins(ent_type, data["tokens"], gold_ents)221 query = self.ent_type2query[ent_type]222 query_tokens = self.tokenizer.tokenize(query)223 try:224 res = self.build_ins(query_tokens, data["tokens"], gold_ents)225 except (ValueError, AssertionError):226 continue227 input_tokens, input_ids, mask, offset, available_spans = res228 ins = {229 "id": data.get("id", str(len(final_data))),230 "ent_type": ent_type,231 "gold_ents": gold_ents,232 "raw_tokens": data["tokens"],233 "input_tokens": input_tokens,234 "input_ids": input_ids,235 "mask": mask,236 "offset": offset,237 "available_spans": available_spans,238 # labels are dynamically padded in collate fn239 "labels": None,240 # "labels": labels.tolist(),241 }242 final_data.append(ins)243 244 # # upper bound analysis245 # pred_spans = set(decode_nnw_thw_mat(labels.unsqueeze(0))[0])246 # g_ents = set(available_spans)247 # tp += len(g_ents & pred_spans)248 # fp += len(pred_spans - g_ents)249 # fn += len(g_ents - pred_spans)250 251 # # upper bound results252 # measures = calc_p_r_f1_from_tp_fp_fn(tp, fp, fn)253 # logger.info(f"Upper Bound: {measures}")254 255 return final_data256 257 def predict_transform(self, texts: List[str]):258 dataset = []259 for text_id, text in enumerate(texts):260 data_id = f"Prediction#{text_id}"261 tokens = self.tokenizer.tokenize(text)262 dataset.append(263 {264 "id": data_id,265 "tokens": tokens,266 "ents": [],267 }268 )269 final_data = self(dataset, disable_pbar=True)270 return final_data271 272 273class CachedPointerMRCTransform(CachedTransformBase, PointerTransformMixin):274 def __init__(275 self,276 max_seq_len: int,277 plm_dir: str,278 mode: str = "w2",279 ) -> None:280 super().__init__()281 282 self.max_seq_len: int = max_seq_len283 self.tokenizer: BertTokenizerFast = BertTokenizerFast.from_pretrained(plm_dir)284 285 self.collate_fn: GeneralCollateFn = GeneralCollateFn(286 {287 "input_ids": torch.long,288 "mask": torch.long,289 "labels": torch.long,290 },291 guessing=False,292 missing_key_as_null=True,293 )294 295 if mode == "w2":296 self.collate_fn.update_before_tensorify = self.update_labels297 elif mode == "cons":298 self.collate_fn.update_before_tensorify = (299 self.update_consecutive_span_labels300 )301 else:302 raise ValueError(f"Mode: {mode} not recognizable")303 304 def transform(305 self,306 transform_loader: Iterator,307 dataset_name: str = None,308 **kwargs,309 ) -> Iterable:310 final_data = []311 for data in transform_loader:312 try:313 res = self.build_ins(314 data["query_tokens"],315 data["context_tokens"],316 data["answer_index"],317 data.get("background_tokens"),318 )319 except (ValueError, AssertionError):320 continue321 input_tokens, input_ids, mask, offset, available_spans = res322 ins = {323 "id": data.get("id", str(len(final_data))),324 "gold_spans": sorted(set(tuple(x) for x in data["answer_index"])),325 "raw_tokens": data["context_tokens"],326 "input_tokens": input_tokens,327 "input_ids": input_ids,328 "mask": mask,329 "offset": offset,330 "available_spans": available_spans,331 "labels": None,332 }333 final_data.append(ins)334 335 return final_data336 337 def predict_transform(self, data: list[dict]):338 """339 Args:340 data: a list of dict with query, context, and background strings341 """342 dataset = []343 for idx, ins in enumerate(data):344 idx = f"Prediction#{idx}"345 dataset.append(346 {347 "id": idx,348 "query_tokens": list(ins["query"]),349 "context_tokens": list(ins["context"]),350 "background_tokens": list(ins.get("background")),351 "answer_index": [],352 }353 )354 final_data = self(dataset, disable_pbar=True, num_samples=0)355 return final_data356 357 358class CachedLabelPointerTransform(CachedTransformOneBase):359 """Transform for label-token linking for skip consecutive spans"""360 361 def __init__(362 self,363 max_seq_len: int,364 plm_dir: str,365 mode: str = "w2",366 label_span: str = "tag",367 include_instructions: bool = True,368 **kwargs,369 ) -> None:370 super().__init__()371 372 self.max_seq_len: int = max_seq_len373 self.mode = mode374 self.label_span = label_span375 self.include_instructions = include_instructions376 377 self.tokenizer: DebertaV2TokenizerFast = DebertaV2TokenizerFast.from_pretrained(378 plm_dir379 )380 self.lc_token = "[LC]"381 self.lm_token = "[LM]"382 self.lr_token = "[LR]"383 self.i_token = "[I]"384 self.tl_token = "[TL]"385 self.tp_token = "[TP]"386 self.b_token = "[B]"387 num_added = self.tokenizer.add_tokens(388 [389 self.lc_token,390 self.lm_token,391 self.lr_token,392 self.i_token,393 self.tl_token,394 self.tp_token,395 self.b_token,396 ]397 )398 assert num_added == 7399 400 self.collate_fn: GeneralCollateFn = GeneralCollateFn(401 {402 "input_ids": torch.long,403 "mask": torch.long,404 "labels": torch.long,405 "spans": None,406 },407 guessing=False,408 missing_key_as_null=True,409 # only for pre-training410 discard_missing=False,411 )412 413 self.collate_fn.update_before_tensorify = self.skip_consecutive_span_labels414 415 def transform(self, instance: dict, **kwargs):416 # input417 tokens = [self.tokenizer.cls_token]418 mask = [1]419 label_map = {"lc": {}, "lm": {}, "lr": {}}420 # (2, 3): {"type": "lc", "task": "cls/ent/rel/event/hyper_rel/discontinuous_ent", "string": ""}421 span_to_label = {}422 423 def _update_seq(424 label: str,425 label_type: str,426 task: str = "",427 label_mask: int = 4,428 content_mask: int = 5,429 ):430 if label not in label_map[label_type]:431 label_token_map = {432 "lc": self.lc_token,433 "lm": self.lm_token,434 "lr": self.lr_token,435 }436 label_tag_start_idx = len(tokens)437 tokens.append(label_token_map[label_type])438 mask.append(label_mask)439 label_tag_end_idx = len(tokens) - 1 # exact end position440 label_tokens = self.tokenizer(label, add_special_tokens=False).tokens()441 label_content_start_idx = len(tokens)442 tokens.extend(label_tokens)443 mask.extend([content_mask] * len(label_tokens))444 label_content_end_idx = len(tokens) - 1 # exact end position445 446 if self.label_span == "tag":447 start_idx = label_tag_start_idx448 end_idx = label_tag_end_idx449 elif self.label_span == "content":450 start_idx = label_content_start_idx451 end_idx = label_content_end_idx452 else:453 raise ValueError(f"label_span={self.label_span} is not supported")454 455 if end_idx == start_idx:456 label_map[label_type][label] = (start_idx,)457 else:458 label_map[label_type][label] = (start_idx, end_idx)459 span_to_label[label_map[label_type][label]] = {460 "type": label_type,461 "task": task,462 "string": label,463 }464 return label_map[label_type][label]465 466 if self.include_instructions:467 instruction = instance.get("instruction")468 if not instruction:469 logger.warning(470 "include_instructions=True, while the instruction is empty!"471 )472 else:473 instruction = ""474 if instruction:475 tokens.append(self.i_token)476 mask.append(2)477 instruction_tokens = self.tokenizer(478 instruction, add_special_tokens=False479 ).tokens()480 tokens.extend(instruction_tokens)481 mask.extend([3] * len(instruction_tokens))482 types = instance["schema"].get("cls")483 if types:484 for t in types:485 _update_seq(t, "lc", task="cls")486 mention_types = instance["schema"].get("ent")487 if mention_types:488 for mt in mention_types:489 _update_seq(mt, "lm", task="ent")490 discon_ent_types = instance["schema"].get("discontinuous_ent")491 if discon_ent_types:492 for mt in discon_ent_types:493 _update_seq(mt, "lm", task="discontinuous_ent")494 rel_types = instance["schema"].get("rel")495 if rel_types:496 for rt in rel_types:497 _update_seq(rt, "lr", task="rel")498 hyper_rel_schema = instance["schema"].get("hyper_rel")499 if hyper_rel_schema:500 for rel, qualifiers in hyper_rel_schema.items():501 _update_seq(rel, "lr", task="hyper_rel")502 for qualifier in qualifiers:503 _update_seq(qualifier, "lr", task="hyper_rel")504 event_schema = instance["schema"].get("event")505 if event_schema:506 for event_type, roles in event_schema.items():507 _update_seq(event_type, "lm", task="event")508 for role in roles:509 _update_seq(role, "lr", task="event")510 511 text = instance.get("text")512 if text:513 text_tokenized = self.tokenizer(514 text, return_offsets_mapping=True, add_special_tokens=False515 )516 if any(val for val in label_map.values()):517 text_label_token = self.tl_token518 else:519 text_label_token = self.tp_token520 tokens.append(text_label_token)521 mask.append(6)522 remain_token_len = self.max_seq_len - 1 - len(tokens)523 if remain_token_len < 5 and kwargs.get("dataset_name", "train") == "train":524 return None525 text_off = len(tokens)526 text_tokens = text_tokenized.tokens()[:remain_token_len]527 tokens.extend(text_tokens)528 mask.extend([7] * len(text_tokens))529 else:530 text_tokenized = None531 532 bg = instance.get("bg")533 if bg:534 bg_tokenized = self.tokenizer(535 bg, return_offsets_mapping=True, add_special_tokens=False536 )537 tokens.append(self.b_token)538 mask.append(8)539 remain_token_len = self.max_seq_len - 1 - len(tokens)540 if remain_token_len < 5 and kwargs.get("dataset_name", "train") == "train":541 return None542 bg_tokens = bg_tokenized.tokens()[:remain_token_len]543 tokens.extend(bg_tokens)544 mask.extend([9] * len(bg_tokens))545 else:546 bg_tokenized = None547 548 tokens.append(self.tokenizer.sep_token)549 mask.append(10)550 551 # labels552 # spans: [[(ent_type start, ent_type end + 1), (ent s, ent e + 1)]]553 spans = [] # one span may have many parts554 if "cls" in instance["ans"]:555 for t in instance["ans"]["cls"]:556 part = label_map["lc"][t]557 spans.append([part])558 if "ent" in instance["ans"]:559 for ent in instance["ans"]["ent"]:560 label_part = label_map["lm"][ent["type"]]561 position_seq = self.char_to_token_span(562 ent["span"], text_tokenized, text_off563 )564 spans.append([label_part, position_seq])565 if "discontinuous_ent" in instance["ans"]:566 for ent in instance["ans"]["discontinuous_ent"]:567 label_part = label_map["lm"][ent["type"]]568 ent_span = [label_part]569 for part in ent["span"]:570 position_seq = self.char_to_token_span(571 part, text_tokenized, text_off572 )573 ent_span.append(position_seq)574 spans.append(ent_span)575 if "rel" in instance["ans"]:576 for rel in instance["ans"]["rel"]:577 label_part = label_map["lr"][rel["relation"]]578 head_position_seq = self.char_to_token_span(579 rel["head"]["span"], text_tokenized, text_off580 )581 tail_position_seq = self.char_to_token_span(582 rel["tail"]["span"], text_tokenized, text_off583 )584 spans.append([label_part, head_position_seq, tail_position_seq])585 if "hyper_rel" in instance["ans"]:586 for rel in instance["ans"]["hyper_rel"]:587 label_part = label_map["lr"][rel["relation"]]588 head_position_seq = self.char_to_token_span(589 rel["head"]["span"], text_tokenized, text_off590 )591 tail_position_seq = self.char_to_token_span(592 rel["tail"]["span"], text_tokenized, text_off593 )594 # rel_span = [label_part, head_position_seq, tail_position_seq]595 for q in rel["qualifiers"]:596 q_label_part = label_map["lr"][q["label"]]597 q_position_seq = self.char_to_token_span(598 q["span"], text_tokenized, text_off599 )600 spans.append(601 [602 label_part,603 head_position_seq,604 tail_position_seq,605 q_label_part,606 q_position_seq,607 ]608 )609 if "event" in instance["ans"]:610 for event in instance["ans"]["event"]:611 event_type_label_part = label_map["lm"][event["event_type"]]612 trigger_position_seq = self.char_to_token_span(613 event["trigger"]["span"], text_tokenized, text_off614 )615 trigger_part = [event_type_label_part, trigger_position_seq]616 spans.append(trigger_part)617 for arg in event["args"]:618 role_label_part = label_map["lr"][arg["role"]]619 arg_position_seq = self.char_to_token_span(620 arg["span"], text_tokenized, text_off621 )622 arg_part = [role_label_part, trigger_position_seq, arg_position_seq]623 spans.append(arg_part)624 if "span" in instance["ans"]:625 # Extractive-QA or Extractive-MRC tasks626 for span in instance["ans"]["span"]:627 span_position_seq = self.char_to_token_span(628 span["span"], text_tokenized, text_off629 )630 spans.append([span_position_seq])631 632 if self.mode == "w2":633 new_spans = []634 for parts in spans:635 new_parts = []636 for part in parts:637 new_parts.append(tuple(range(part[0], part[-1] + 1)))638 new_spans.append(new_parts)639 spans = new_spans640 elif self.mode == "span":641 spans = spans642 else:643 raise ValueError(f"mode={self.mode} is not supported")644 645 ins = {646 "raw": instance,647 "tokens": tokens,648 "input_ids": self.tokenizer.convert_tokens_to_ids(tokens),649 "mask": mask,650 "spans": spans,651 "label_map": label_map,652 "span_to_label": span_to_label,653 "labels": None, # labels are calculated dynamically in collate_fn654 }655 return ins656 657 def char_to_token_span(658 self, span: list[int], tokenized: BatchEncoding, offset: int = 0659 ) -> list[int]:660 token_s = tokenized.char_to_token(span[0])661 token_e = tokenized.char_to_token(span[1] - 1)662 if token_e == token_s:663 position_seq = (offset + token_s,)664 else:665 position_seq = (offset + token_s, offset + token_e)666 return position_seq667 668 def skip_consecutive_span_labels(self, data: dict) -> dict:669 bs = len(data["input_ids"])670 max_seq_len = max(len(input_ids) for input_ids in data["input_ids"])671 batch_seq_len = min(self.max_seq_len, max_seq_len)672 for i in range(bs):673 data["input_ids"][i] = data["input_ids"][i][:batch_seq_len]674 data["mask"][i] = data["mask"][i][:batch_seq_len]675 assert len(data["input_ids"][i]) == len(data["mask"][i])676 pad_len = batch_seq_len - len(data["mask"][i])677 data["input_ids"][i] = (678 data["input_ids"][i] + [self.tokenizer.pad_token_id] * pad_len679 )680 data["mask"][i] = data["mask"][i] + [0] * pad_len681 data["labels"][i] = encode_nnw_nsw_thw_mat(data["spans"][i], batch_seq_len)682 683 # # for debugging only684 # pred_spans = decode_nnw_nsw_thw_mat(data["labels"][i].unsqueeze(0))[0]685 # sorted_gold = sorted(set(tuple(x) for x in data["spans"][i]))686 # sorted_pred = sorted(set(tuple(x) for x in pred_spans))687 # if sorted_gold != sorted_pred:688 # breakpoint()689 690 # # for pre-training only691 # del data["spans"]692 693 return data694 