Masking specific token in each input sentence during Masked language modelling

Viewed 405

I have a dataset with 2 columns: token, sentence. For example:

{'token':'shrouded', 'sentence':'A mist shrouded the sun'}

I want to fine-tune one of the Huggingface Transformers model on a Masked Language Modelling task. (For now I am using distilroberta-base as per this tutorial)

Now, instead of random masking, I am trying to specifically mask the token in the sentence while training. For eg. A mist [MASK] the sun and then get the model to predict the token shrouded.

Now I understand that in random masking we can simply use DataCollatorForLanguageModeling and feed it into the Trainer. However, in this use case, masking will have to be done at the pre-processing stage. I can't figure out how to do that.

Here is the code so far:

...

datasets = load_dataset('csv', data_files=['word_sentence_1.csv'])

model_checkpoint = "distilroberta-base"

def tokenize_function(examples):
    return tokenizer(examples["sentence"])


tokenizer = AutoTokenizer.from_pretrained(model_checkpoint, use_fast=True)
tokenized_datasets = datasets.map(tokenize_function, batched=True, num_proc=4)

model = AutoModelForMaskedLM.from_pretrained(model_checkpoint)

model_name = model_checkpoint.split("/")[-1]
training_args = TrainingArguments(
    f"{model_name}-word_sentence_1_1",
    evaluation_strategy = "epoch",
    learning_rate=2e-5,
    weight_decay=0.01,
    push_to_hub=False,
)

##### Need to remove this and add logic of static masking ####
from transformers import DataCollatorForLanguageModeling
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm_probability=0.15)


trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_datasets,
    data_collator=data_collator,
)

trainer.train()
1 Answers

I ended solving this problem myself:

The idea is to manually do the masking and at the same time provide 'labels' for the masks. All I had to do was to make some changes in the tokenize_function and remove the data_collator.

MASK_TOKEN = tokenizer.convert_ids_to_tokens(tokenizer.mask_token_id)
MASK_TOKEN_ID = tokenizer.mask_token_id

def tokenize_function(examples):

    usage_arr = [ examples['sentence'][i].replace(examples['word'][i], MASK_TOKEN)  for i in range(len(examples['word']))]
    tokenized_data = tokenizer(usage_arr, padding="max_length", truncation=True)

    label_arr_list = []

    for i in range(len(usage_arr)):

        label_arr = [-100] * len(tokenized_data.input_ids[i])
        if MASK_TOKEN_ID in tokenized_data.input_ids[i]:
            label_arr[tokenized_data.input_ids[i].index(MASK_TOKEN_ID)] = tokenizer.convert_tokens_to_ids(examples['word'][i])
        label_arr_list.append(label_arr)


    tokenized_data['labels'] = label_arr_list

    return tokenized_data
Related