ceyda/ExplaiNER
1
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 