CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
glue.py644 linesDownload Raw Back to processors
1# coding=utf-82# Copyright 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.3# Copyright (c) 2018, NVIDIA CORPORATION.  All rights reserved.4#5# Licensed under the Apache License, Version 2.0 (the "License");6# you may not use this file except in compliance with the License.7# You may obtain a copy of the License at8#9#     http://www.apache.org/licenses/LICENSE-2.010#11# Unless required by applicable law or agreed to in writing, software12# distributed under the License is distributed on an "AS IS" BASIS,13# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14# See the License for the specific language governing permissions and15# limitations under the License.16"""GLUE processors and helpers"""17 18import os19import warnings20from dataclasses import asdict21from enum import Enum22from typing import Optional, Union23 24from ...tokenization_utils import PreTrainedTokenizer25from ...utils import is_tf_available, logging26from .utils import DataProcessor, InputExample, InputFeatures27 28 29if is_tf_available():30    import tensorflow as tf31 32logger = logging.get_logger(__name__)33 34DEPRECATION_WARNING = (35    "This {0} will be removed from the library soon, preprocessing should be handled with the ๐Ÿค— Datasets "36    "library. You can have a look at this example script for pointers: "37    "https://github.com/huggingface/transformers/blob/main/examples/pytorch/text-classification/run_glue.py"38)39 40 41def glue_convert_examples_to_features(42    examples: Union[list[InputExample], "tf.data.Dataset"],43    tokenizer: PreTrainedTokenizer,44    max_length: Optional[int] = None,45    task=None,46    label_list=None,47    output_mode=None,48):49    """50    Loads a data file into a list of `InputFeatures`51 52    Args:53        examples: List of `InputExamples` or `tf.data.Dataset` containing the examples.54        tokenizer: Instance of a tokenizer that will tokenize the examples55        max_length: Maximum example length. Defaults to the tokenizer's max_len56        task: GLUE task57        label_list: List of labels. Can be obtained from the processor using the `processor.get_labels()` method58        output_mode: String indicating the output mode. Either `regression` or `classification`59 60    Returns:61        If the `examples` input is a `tf.data.Dataset`, will return a `tf.data.Dataset` containing the task-specific62        features. If the input is a list of `InputExamples`, will return a list of task-specific `InputFeatures` which63        can be fed to the model.64 65    """66    warnings.warn(DEPRECATION_WARNING.format("function"), FutureWarning)67    if is_tf_available() and isinstance(examples, tf.data.Dataset):68        if task is None:69            raise ValueError("When calling glue_convert_examples_to_features from TF, the task parameter is required.")70        return _tf_glue_convert_examples_to_features(examples, tokenizer, max_length=max_length, task=task)71    return _glue_convert_examples_to_features(72        examples, tokenizer, max_length=max_length, task=task, label_list=label_list, output_mode=output_mode73    )74 75 76if is_tf_available():77 78    def _tf_glue_convert_examples_to_features(79        examples: tf.data.Dataset,80        tokenizer: PreTrainedTokenizer,81        task=str,82        max_length: Optional[int] = None,83    ) -> tf.data.Dataset:84        """85        Returns:86            A `tf.data.Dataset` containing the task-specific features.87 88        """89        processor = glue_processors[task]()90        examples = [processor.tfds_map(processor.get_example_from_tensor_dict(example)) for example in examples]91        features = glue_convert_examples_to_features(examples, tokenizer, max_length=max_length, task=task)92        label_type = tf.float32 if task == "sts-b" else tf.int6493 94        def gen():95            for ex in features:96                d = {k: v for k, v in asdict(ex).items() if v is not None}97                label = d.pop("label")98                yield (d, label)99 100        input_names = tokenizer.model_input_names101 102        return tf.data.Dataset.from_generator(103            gen,104            (dict.fromkeys(input_names, tf.int32), label_type),105            ({k: tf.TensorShape([None]) for k in input_names}, tf.TensorShape([])),106        )107 108 109def _glue_convert_examples_to_features(110    examples: list[InputExample],111    tokenizer: PreTrainedTokenizer,112    max_length: Optional[int] = None,113    task=None,114    label_list=None,115    output_mode=None,116):117    if max_length is None:118        max_length = tokenizer.model_max_length119 120    if task is not None:121        processor = glue_processors[task]()122        if label_list is None:123            label_list = processor.get_labels()124            logger.info(f"Using label list {label_list} for task {task}")125        if output_mode is None:126            output_mode = glue_output_modes[task]127            logger.info(f"Using output mode {output_mode} for task {task}")128 129    label_map = {label: i for i, label in enumerate(label_list)}130 131    def label_from_example(example: InputExample) -> Union[int, float, None]:132        if example.label is None:133            return None134        if output_mode == "classification":135            return label_map[example.label]136        elif output_mode == "regression":137            return float(example.label)138        raise KeyError(output_mode)139 140    labels = [label_from_example(example) for example in examples]141 142    batch_encoding = tokenizer(143        [(example.text_a, example.text_b) for example in examples],144        max_length=max_length,145        padding="max_length",146        truncation=True,147    )148 149    features = []150    for i in range(len(examples)):151        inputs = {k: batch_encoding[k][i] for k in batch_encoding}152 153        feature = InputFeatures(**inputs, label=labels[i])154        features.append(feature)155 156    for i, example in enumerate(examples[:5]):157        logger.info("*** Example ***")158        logger.info(f"guid: {example.guid}")159        logger.info(f"features: {features[i]}")160 161    return features162 163 164class OutputMode(Enum):165    classification = "classification"166    regression = "regression"167 168 169class MrpcProcessor(DataProcessor):170    """Processor for the MRPC data set (GLUE version)."""171 172    def __init__(self, *args, **kwargs):173        super().__init__(*args, **kwargs)174        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)175 176    def get_example_from_tensor_dict(self, tensor_dict):177        """See base class."""178        return InputExample(179            tensor_dict["idx"].numpy(),180            tensor_dict["sentence1"].numpy().decode("utf-8"),181            tensor_dict["sentence2"].numpy().decode("utf-8"),182            str(tensor_dict["label"].numpy()),183        )184 185    def get_train_examples(self, data_dir):186        """See base class."""187        logger.info(f"LOOKING AT {os.path.join(data_dir, 'train.tsv')}")188        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")189 190    def get_dev_examples(self, data_dir):191        """See base class."""192        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")193 194    def get_test_examples(self, data_dir):195        """See base class."""196        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")197 198    def get_labels(self):199        """See base class."""200        return ["0", "1"]201 202    def _create_examples(self, lines, set_type):203        """Creates examples for the training, dev and test sets."""204        examples = []205        for i, line in enumerate(lines):206            if i == 0:207                continue208            guid = f"{set_type}-{i}"209            text_a = line[3]210            text_b = line[4]211            label = None if set_type == "test" else line[0]212            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))213        return examples214 215 216class MnliProcessor(DataProcessor):217    """Processor for the MultiNLI data set (GLUE version)."""218 219    def __init__(self, *args, **kwargs):220        super().__init__(*args, **kwargs)221        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)222 223    def get_example_from_tensor_dict(self, tensor_dict):224        """See base class."""225        return InputExample(226            tensor_dict["idx"].numpy(),227            tensor_dict["premise"].numpy().decode("utf-8"),228            tensor_dict["hypothesis"].numpy().decode("utf-8"),229            str(tensor_dict["label"].numpy()),230        )231 232    def get_train_examples(self, data_dir):233        """See base class."""234        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")235 236    def get_dev_examples(self, data_dir):237        """See base class."""238        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev_matched.tsv")), "dev_matched")239 240    def get_test_examples(self, data_dir):241        """See base class."""242        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test_matched.tsv")), "test_matched")243 244    def get_labels(self):245        """See base class."""246        return ["contradiction", "entailment", "neutral"]247 248    def _create_examples(self, lines, set_type):249        """Creates examples for the training, dev and test sets."""250        examples = []251        for i, line in enumerate(lines):252            if i == 0:253                continue254            guid = f"{set_type}-{line[0]}"255            text_a = line[8]256            text_b = line[9]257            label = None if set_type.startswith("test") else line[-1]258            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))259        return examples260 261 262class MnliMismatchedProcessor(MnliProcessor):263    """Processor for the MultiNLI Mismatched data set (GLUE version)."""264 265    def __init__(self, *args, **kwargs):266        super().__init__(*args, **kwargs)267        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)268 269    def get_dev_examples(self, data_dir):270        """See base class."""271        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev_mismatched.tsv")), "dev_mismatched")272 273    def get_test_examples(self, data_dir):274        """See base class."""275        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test_mismatched.tsv")), "test_mismatched")276 277 278class ColaProcessor(DataProcessor):279    """Processor for the CoLA data set (GLUE version)."""280 281    def __init__(self, *args, **kwargs):282        super().__init__(*args, **kwargs)283        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)284 285    def get_example_from_tensor_dict(self, tensor_dict):286        """See base class."""287        return InputExample(288            tensor_dict["idx"].numpy(),289            tensor_dict["sentence"].numpy().decode("utf-8"),290            None,291            str(tensor_dict["label"].numpy()),292        )293 294    def get_train_examples(self, data_dir):295        """See base class."""296        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")297 298    def get_dev_examples(self, data_dir):299        """See base class."""300        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")301 302    def get_test_examples(self, data_dir):303        """See base class."""304        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")305 306    def get_labels(self):307        """See base class."""308        return ["0", "1"]309 310    def _create_examples(self, lines, set_type):311        """Creates examples for the training, dev and test sets."""312        test_mode = set_type == "test"313        if test_mode:314            lines = lines[1:]315        text_index = 1 if test_mode else 3316        examples = []317        for i, line in enumerate(lines):318            guid = f"{set_type}-{i}"319            text_a = line[text_index]320            label = None if test_mode else line[1]321            examples.append(InputExample(guid=guid, text_a=text_a, text_b=None, label=label))322        return examples323 324 325class Sst2Processor(DataProcessor):326    """Processor for the SST-2 data set (GLUE version)."""327 328    def __init__(self, *args, **kwargs):329        super().__init__(*args, **kwargs)330        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)331 332    def get_example_from_tensor_dict(self, tensor_dict):333        """See base class."""334        return InputExample(335            tensor_dict["idx"].numpy(),336            tensor_dict["sentence"].numpy().decode("utf-8"),337            None,338            str(tensor_dict["label"].numpy()),339        )340 341    def get_train_examples(self, data_dir):342        """See base class."""343        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")344 345    def get_dev_examples(self, data_dir):346        """See base class."""347        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")348 349    def get_test_examples(self, data_dir):350        """See base class."""351        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")352 353    def get_labels(self):354        """See base class."""355        return ["0", "1"]356 357    def _create_examples(self, lines, set_type):358        """Creates examples for the training, dev and test sets."""359        examples = []360        text_index = 1 if set_type == "test" else 0361        for i, line in enumerate(lines):362            if i == 0:363                continue364            guid = f"{set_type}-{i}"365            text_a = line[text_index]366            label = None if set_type == "test" else line[1]367            examples.append(InputExample(guid=guid, text_a=text_a, text_b=None, label=label))368        return examples369 370 371class StsbProcessor(DataProcessor):372    """Processor for the STS-B data set (GLUE version)."""373 374    def __init__(self, *args, **kwargs):375        super().__init__(*args, **kwargs)376        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)377 378    def get_example_from_tensor_dict(self, tensor_dict):379        """See base class."""380        return InputExample(381            tensor_dict["idx"].numpy(),382            tensor_dict["sentence1"].numpy().decode("utf-8"),383            tensor_dict["sentence2"].numpy().decode("utf-8"),384            str(tensor_dict["label"].numpy()),385        )386 387    def get_train_examples(self, data_dir):388        """See base class."""389        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")390 391    def get_dev_examples(self, data_dir):392        """See base class."""393        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")394 395    def get_test_examples(self, data_dir):396        """See base class."""397        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")398 399    def get_labels(self):400        """See base class."""401        return [None]402 403    def _create_examples(self, lines, set_type):404        """Creates examples for the training, dev and test sets."""405        examples = []406        for i, line in enumerate(lines):407            if i == 0:408                continue409            guid = f"{set_type}-{line[0]}"410            text_a = line[7]411            text_b = line[8]412            label = None if set_type == "test" else line[-1]413            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))414        return examples415 416 417class QqpProcessor(DataProcessor):418    """Processor for the QQP data set (GLUE version)."""419 420    def __init__(self, *args, **kwargs):421        super().__init__(*args, **kwargs)422        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)423 424    def get_example_from_tensor_dict(self, tensor_dict):425        """See base class."""426        return InputExample(427            tensor_dict["idx"].numpy(),428            tensor_dict["question1"].numpy().decode("utf-8"),429            tensor_dict["question2"].numpy().decode("utf-8"),430            str(tensor_dict["label"].numpy()),431        )432 433    def get_train_examples(self, data_dir):434        """See base class."""435        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")436 437    def get_dev_examples(self, data_dir):438        """See base class."""439        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")440 441    def get_test_examples(self, data_dir):442        """See base class."""443        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")444 445    def get_labels(self):446        """See base class."""447        return ["0", "1"]448 449    def _create_examples(self, lines, set_type):450        """Creates examples for the training, dev and test sets."""451        test_mode = set_type == "test"452        q1_index = 1 if test_mode else 3453        q2_index = 2 if test_mode else 4454        examples = []455        for i, line in enumerate(lines):456            if i == 0:457                continue458            guid = f"{set_type}-{line[0]}"459            try:460                text_a = line[q1_index]461                text_b = line[q2_index]462                label = None if test_mode else line[5]463            except IndexError:464                continue465            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))466        return examples467 468 469class QnliProcessor(DataProcessor):470    """Processor for the QNLI data set (GLUE version)."""471 472    def __init__(self, *args, **kwargs):473        super().__init__(*args, **kwargs)474        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)475 476    def get_example_from_tensor_dict(self, tensor_dict):477        """See base class."""478        return InputExample(479            tensor_dict["idx"].numpy(),480            tensor_dict["question"].numpy().decode("utf-8"),481            tensor_dict["sentence"].numpy().decode("utf-8"),482            str(tensor_dict["label"].numpy()),483        )484 485    def get_train_examples(self, data_dir):486        """See base class."""487        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")488 489    def get_dev_examples(self, data_dir):490        """See base class."""491        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")492 493    def get_test_examples(self, data_dir):494        """See base class."""495        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")496 497    def get_labels(self):498        """See base class."""499        return ["entailment", "not_entailment"]500 501    def _create_examples(self, lines, set_type):502        """Creates examples for the training, dev and test sets."""503        examples = []504        for i, line in enumerate(lines):505            if i == 0:506                continue507            guid = f"{set_type}-{line[0]}"508            text_a = line[1]509            text_b = line[2]510            label = None if set_type == "test" else line[-1]511            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))512        return examples513 514 515class RteProcessor(DataProcessor):516    """Processor for the RTE data set (GLUE version)."""517 518    def __init__(self, *args, **kwargs):519        super().__init__(*args, **kwargs)520        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)521 522    def get_example_from_tensor_dict(self, tensor_dict):523        """See base class."""524        return InputExample(525            tensor_dict["idx"].numpy(),526            tensor_dict["sentence1"].numpy().decode("utf-8"),527            tensor_dict["sentence2"].numpy().decode("utf-8"),528            str(tensor_dict["label"].numpy()),529        )530 531    def get_train_examples(self, data_dir):532        """See base class."""533        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")534 535    def get_dev_examples(self, data_dir):536        """See base class."""537        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")538 539    def get_test_examples(self, data_dir):540        """See base class."""541        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")542 543    def get_labels(self):544        """See base class."""545        return ["entailment", "not_entailment"]546 547    def _create_examples(self, lines, set_type):548        """Creates examples for the training, dev and test sets."""549        examples = []550        for i, line in enumerate(lines):551            if i == 0:552                continue553            guid = f"{set_type}-{line[0]}"554            text_a = line[1]555            text_b = line[2]556            label = None if set_type == "test" else line[-1]557            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))558        return examples559 560 561class WnliProcessor(DataProcessor):562    """Processor for the WNLI data set (GLUE version)."""563 564    def __init__(self, *args, **kwargs):565        super().__init__(*args, **kwargs)566        warnings.warn(DEPRECATION_WARNING.format("processor"), FutureWarning)567 568    def get_example_from_tensor_dict(self, tensor_dict):569        """See base class."""570        return InputExample(571            tensor_dict["idx"].numpy(),572            tensor_dict["sentence1"].numpy().decode("utf-8"),573            tensor_dict["sentence2"].numpy().decode("utf-8"),574            str(tensor_dict["label"].numpy()),575        )576 577    def get_train_examples(self, data_dir):578        """See base class."""579        return self._create_examples(self._read_tsv(os.path.join(data_dir, "train.tsv")), "train")580 581    def get_dev_examples(self, data_dir):582        """See base class."""583        return self._create_examples(self._read_tsv(os.path.join(data_dir, "dev.tsv")), "dev")584 585    def get_test_examples(self, data_dir):586        """See base class."""587        return self._create_examples(self._read_tsv(os.path.join(data_dir, "test.tsv")), "test")588 589    def get_labels(self):590        """See base class."""591        return ["0", "1"]592 593    def _create_examples(self, lines, set_type):594        """Creates examples for the training, dev and test sets."""595        examples = []596        for i, line in enumerate(lines):597            if i == 0:598                continue599            guid = f"{set_type}-{line[0]}"600            text_a = line[1]601            text_b = line[2]602            label = None if set_type == "test" else line[-1]603            examples.append(InputExample(guid=guid, text_a=text_a, text_b=text_b, label=label))604        return examples605 606 607glue_tasks_num_labels = {608    "cola": 2,609    "mnli": 3,610    "mrpc": 2,611    "sst-2": 2,612    "sts-b": 1,613    "qqp": 2,614    "qnli": 2,615    "rte": 2,616    "wnli": 2,617}618 619glue_processors = {620    "cola": ColaProcessor,621    "mnli": MnliProcessor,622    "mnli-mm": MnliMismatchedProcessor,623    "mrpc": MrpcProcessor,624    "sst-2": Sst2Processor,625    "sts-b": StsbProcessor,626    "qqp": QqpProcessor,627    "qnli": QnliProcessor,628    "rte": RteProcessor,629    "wnli": WnliProcessor,630}631 632glue_output_modes = {633    "cola": "classification",634    "mnli": "classification",635    "mnli-mm": "classification",636    "mrpc": "classification",637    "sst-2": "classification",638    "sts-b": "regression",639    "qqp": "classification",640    "qnli": "classification",641    "rte": "classification",642    "wnli": "classification",643}644 
Aluode/PerceptionLabPortable ยท CoolFace