Team Ai
Apppublic

ceyda/ExplaiNER

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
data.py229 linesDownload Raw Back to src
1from functools import partial2 3import pandas as pd4import streamlit as st5import torch6from datasets import Dataset, DatasetDict, load_dataset  # type: ignore7from torch.nn.functional import cross_entropy8from transformers import DataCollatorForTokenClassification  # type: ignore9 10from src.utils import device, tokenizer_hash_funcs11 12 13@st.cache(allow_output_mutation=True)14def get_data(15    ds_name: str, config_name: str, split_name: str, split_sample_size: int, randomize_sample: bool16) -> Dataset:17    """Loads a Dataset from the HuggingFace hub (if not already loaded).18 19    Uses `datasets.load_dataset` to load the dataset (see its documentation for additional details).20 21    Args:22        ds_name (str): Path or name of the dataset.23        config_name (str): Name of the dataset configuration.24        split_name (str): Which split of the data to load.25        split_sample_size (int): The number of examples to load from the split.26 27    Returns:28        Dataset: A Dataset object.29    """30    ds: DatasetDict = load_dataset(ds_name, name=config_name, use_auth_token=True).shuffle(31        seed=0 if randomize_sample else None32    )  # type: ignore33    split = ds[split_name].select(range(split_sample_size))34    return split35 36 37@st.cache(38    allow_output_mutation=True,39    hash_funcs=tokenizer_hash_funcs,40)41def get_collator(tokenizer) -> DataCollatorForTokenClassification:42    """Returns a DataCollator that will dynamically pad the inputs received, as well as the labels.43 44    Args:45        tokenizer ([PreTrainedTokenizer] or [PreTrainedTokenizerFast]): The tokenizer used for encoding the data.46 47    Returns:48        DataCollatorForTokenClassification: The DataCollatorForTokenClassification object.49    """50    return DataCollatorForTokenClassification(tokenizer)51 52 53def create_word_ids_from_input_ids(tokenizer, input_ids: list[int]) -> list[int]:54    """Takes a list of input_ids and return corresponding word_ids55 56    Args:57        tokenizer: The tokenizer that was used to obtain the input ids.58        input_ids (list[int]): List of token ids.59 60    Returns:61        list[int]: Word ids corresponding to the input ids.62    """63    word_ids = []64    wid = -165    tokens = [tokenizer.convert_ids_to_tokens(i) for i in input_ids]66 67    for i, tok in enumerate(tokens):68        if tok in tokenizer.all_special_tokens:69            word_ids.append(-1)70            continue71 72        if not tokens[i - 1].endswith("@@") and tokens[i - 1] != "<unk>":73            wid += 174 75        word_ids.append(wid)76 77    assert len(word_ids) == len(input_ids)78    return word_ids79 80 81def tokenize(batch, tokenizer) -> dict:82    """Tokenizes a batch of examples.83 84    Args:85        batch: The examples to tokenize86        tokenizer: The tokenizer to use87 88    Returns:89        dict: The tokenized batch90    """91    tokenized_inputs = tokenizer(batch["tokens"], truncation=True, is_split_into_words=True)92    labels = []93    wids = []94 95    for idx, label in enumerate(batch["ner_tags"]):96        try:97            word_ids = tokenized_inputs.word_ids(batch_index=idx)98        except ValueError:99            word_ids = create_word_ids_from_input_ids(100                tokenizer, tokenized_inputs["input_ids"][idx]101            )102        previous_word_idx = None103        label_ids = []104        for word_idx in word_ids:105            if word_idx == -1 or word_idx is None or word_idx == previous_word_idx:106                label_ids.append(-100)107            else:108                label_ids.append(label[word_idx])109            previous_word_idx = word_idx110        wids.append(word_ids)111        labels.append(label_ids)112    tokenized_inputs["word_ids"] = wids113    tokenized_inputs["labels"] = labels114    return tokenized_inputs115 116 117def stringify_ner_tags(batch: dict, tags) -> dict:118    """Stringifies a dataset batch's NER tags."""119    return {"ner_tags_str": [tags.int2str(idx) for idx in batch["ner_tags"]]}120 121 122def encode_dataset(split: Dataset, tokenizer):123    """Encodes a dataset split.124 125    Args:126        split (Dataset): A Dataset object.127        tokenizer: A PreTrainedTokenizer object.128 129    Returns:130        Dataset: A Dataset object with the encoded inputs.131    """132 133    tags = split.features["ner_tags"].feature134    split = split.map(partial(stringify_ner_tags, tags=tags), batched=True)135    remove_columns = split.column_names136    ids = split["id"]137    split = split.map(138        partial(tokenize, tokenizer=tokenizer),139        batched=True,140        remove_columns=remove_columns,141    )142    word_ids = [[id if id is not None else -1 for id in wids] for wids in split["word_ids"]]143    return split.remove_columns(["word_ids"]), word_ids, ids144 145 146def forward_pass_with_label(batch, model, collator, num_classes: int) -> dict:147    """Runs the forward pass for a batch of examples.148 149    Args:150        batch: The batch to process151        model: The model to process the batch with152        collator: A data collator153        num_classes (int): Number of classes154 155    Returns:156        dict: a dictionary containing `losses`, `preds` and `hidden_states`157    """158 159    # Convert dict of lists to list of dicts suitable for data collator160    features = [dict(zip(batch, t)) for t in zip(*batch.values())]161 162    # Pad inputs and labels and put all tensors on device163    batch = collator(features)164    input_ids = batch["input_ids"].to(device)165    attention_mask = batch["attention_mask"].to(device)166    labels = batch["labels"].to(device)167 168    with torch.no_grad():169        # Pass data through model170        output = model(input_ids, attention_mask, output_hidden_states=True)171        # logit.size: [batch_size, sequence_length, classes]172 173        # Predict class with largest logit value on classes axis174        preds = torch.argmax(output.logits, axis=-1).cpu().numpy()  # type: ignore175 176        # Calculate loss per token after flattening batch dimension with view177        loss = cross_entropy(178            output.logits.view(-1, num_classes), labels.view(-1), reduction="none"179        )180 181        # Unflatten batch dimension and convert to numpy array182        loss = loss.view(len(input_ids), -1).cpu().numpy()183        hidden_states = output.hidden_states[-1].cpu().numpy()184 185        # logits = output.logits.view(len(input_ids), -1).cpu().numpy()186 187    return {"losses": loss, "preds": preds, "hidden_states": hidden_states}188 189 190def predict(split_encoded: Dataset, model, tokenizer, collator, tags) -> pd.DataFrame:191    """Generates predictions for a given dataset split and returns the results as a dataframe.192 193    Args:194        split_encoded (Dataset): The dataset to process195        model: The model to process the dataset with196        tokenizer: The tokenizer to process the dataset with197        collator: The data collator to use198        tags: The tags used in the dataset199 200    Returns:201        pd.DataFrame: A dataframe containing token-level predictions.202    """203 204    split_encoded = split_encoded.map(205        partial(206            forward_pass_with_label,207            model=model,208            collator=collator,209            num_classes=tags.num_classes,210        ),211        batched=True,212        batch_size=8,213    )214    df: pd.DataFrame = split_encoded.to_pandas()  # type: ignore215 216    df["tokens"] = df["input_ids"].apply(217        lambda x: tokenizer.convert_ids_to_tokens(x)  # type: ignore218    )219    df["labels"] = df["labels"].apply(220        lambda x: ["IGN" if i == -100 else tags.int2str(int(i)) for i in x]221    )222    df["preds"] = df["preds"].apply(lambda x: [model.config.id2label[i] for i in x])223    df["preds"] = df.apply(lambda x: x["preds"][: len(x["input_ids"])], axis=1)224    df["losses"] = df.apply(lambda x: x["losses"][: len(x["input_ids"])], axis=1)225    df["hidden_states"] = df.apply(lambda x: x["hidden_states"][: len(x["input_ids"])], axis=1)226    df["total_loss"] = df["losses"].apply(sum)227 228    return df229