Aluode/PerceptionLabPortable
0
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 