CoolFace
Apppublic

Spico/Mirror

sourceHugging Faceapache-2.0updated 6d agoView on Hugging Face
4likes
transform.py694 linesDownload Raw Back to src
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