#listing out vocab for random token masking vocab = tokenizer.get_vocab() valid_token_ids = list(vocab.values()) def mask_batch(batch_tokens, clone=True): if clone: batch_tokens = torch.clone(batch_tokens) # Define the percentage of tokens to potentially mask replace_percentage = 0.15 # Define tokens that should not be replaced excluded_tokens = {0, 100, 101, 102, 103} # Create a mask to identify tokens that are eligible for replacement eligible_mask = ~torch.isin(batch_tokens, torch.tensor(list(excluded_tokens)).to(device)) # Count the number of eligible tokens num_eligible_tokens = eligible_mask.sum().item() # Calculate the number of tokens to potentially mask num_tokens_to_mask = int(num_eligible_tokens * replace_percentage) # Create a random permutation of eligible token indices eligible_indices = eligible_mask.nonzero(as_tuple=True) random_indices = torch.randperm(num_eligible_tokens)[:num_tokens_to_mask] # Create a probability distribution for replacement replacement_probs = torch.tensor([0.8, 0.1, 0.1]) # Probabilities for [103, random token, leave unchanged] replacement_choices = torch.multinomial(replacement_probs, num_tokens_to_mask, replacement=True) # Vector to store if a token was masked (0: not masked, 1: masked) masked_indicator = torch.zeros_like(batch_tokens, dtype=torch.int32) # Apply replacements based on sampled choices for i, idx in enumerate(random_indices): row = eligible_indices[0][idx] col = eligible_indices[1][idx] #replacing with [MASK] if replacement_choices[i] == 0: batch_tokens[row, col] = 103 masked_indicator[row, col] = 1 #replacing with random token elif replacement_choices[i] == 1: batch_tokens[row, col] = random.choice(valid_token_ids) masked_indicator[row, col] = 1 #not replacing at all elif replacement_choices[i] == 2: masked_indicator[row, col] = 1 return batch_tokens, masked_indicator batch_tokens, masked_indicator = mask_batch(sequence_tokens_batches[0]) batch_tokens