Team Ai
Apppublic

mrm8488/PromptSource

sourceHugging Faceupdated 5y agoView on Hugging Face
4likes
utils.py78 linesDownload Raw Back to seqio_tasks
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