mrm8488/PromptSource
4
1import re2 3import datasets4import tensorflow as tf5 6import promptsource.utils7 8 9def feature_to_spec(feature, length=False):10 if isinstance(feature, datasets.ClassLabel):11 return tf.TensorSpec(shape=() if not length else (None if length == -1 else length,), dtype=tf.int64)12 elif isinstance(feature, datasets.Value):13 return tf.TensorSpec(14 shape=() if not length else (None if length == -1 else length,), dtype=getattr(tf.dtypes, feature.dtype)15 )16 elif hasattr(feature, "dtype") and hasattr(feature, "shape"):17 return tf.TensorSpec(shape=feature.shape, dtype=feature.dtype)18 elif isinstance(feature, datasets.Sequence):19 return feature_to_spec(feature.feature, length=feature.length)20 elif isinstance(feature, list):21 return [feature_to_spec(f, length=length) for f in feature]22 elif isinstance(feature, dict):23 return {k: feature_to_spec(v, length=length) for k, v in feature.items()}24 else:25 raise ValueError(f"Unparseable feature type {type(feature)}")26 27 28def hf_dataset_to_tf_dataset(dataset):29 return tf.data.Dataset.from_generator(30 dataset.__iter__, output_signature={k: feature_to_spec(v) for k, v in dataset.features.items()}31 )32 33 34def apply_template(dataset, template):35 def map_fn(ex):36 ex = promptsource.utils.removeHyphen(ex)37 inputs_and_targets = template.apply(ex)38 answer_choices = template.get_answer_choices_list(ex)39 if len(inputs_and_targets) == 2:40 inputs, targets = inputs_and_targets41 if targets == "":42 ex = {"inputs": inputs, "targets": "<NO LABEL>"}43 else:44 ex = {"inputs": inputs, "targets": targets}45 # When template results in an empty example, template.apply returns [""]46 # Also, if the template gets split wrong, len can be > 247 # We will filter these out later48 else:49 ex = {"inputs": "", "targets": ""}50 51 if answer_choices:52 ex["answer_choices"] = answer_choices53 54 return ex55 56 def filter_fn(ex):57 return len(ex["inputs"]) > 0 and len(ex["targets"]) > 058 59 original_columns = dataset.column_names60 dataset = dataset.map(map_fn).filter(filter_fn)61 # map keeps original columns, remove them62 return dataset.remove_columns(set(original_columns) - {"inputs", "targets", "answer_choices"})63 64 65def get_dataset_splits(dataset_name, subset_name=None):66 info = datasets.get_dataset_infos(dataset_name)67 subset_name = subset_name or list(info.keys())[0]68 return info[subset_name].splits69 70 71def task_clean(text):72 # Clean the text according to allowed characters for a task name73 return re.sub(r"[^\w\d\._]+", "_", text)74 75 76def get_task_name(dataset_name, subset_name, template_name):77 return task_clean(dataset_name + (f"_{subset_name}_" if subset_name is not None else "_") + template_name)78 