In [4]:
import pandas as pd
from transformers import (
    AutoTokenizer, 
    DataCollatorForLanguageModeling
)
import os
import torch

tokenizer = AutoTokenizer.from_pretrained("meta-llama/Meta-Llama-3.1-8B-Instruct")
df = pd.read_csv('../../data/training_data/DFINAL_VTrain.csv')

In [5]:
tokenizer.pad_token = tokenizer.eos_token

In [6]:
class SpecialTokenCollator(DataCollatorForLanguageModeling):
    def __init__(self, assistant_header, eot_token, tokenizer, *args, **kwargs):
        super().__init__(tokenizer=tokenizer, *args, **kwargs)
        self.assistant_header = tokenizer.encode(
            assistant_header, add_special_tokens=False
        )
        self.eot_token = tokenizer.encode(eot_token, add_special_tokens=False)[0]
        
    def torch_call(self, examples):
        batch = super().torch_call(examples)
        
        for i, input_ids in enumerate(batch["input_ids"]):
            # Find all assistant headers
            header_positions = []
            curr_idx = 0
            while curr_idx <= len(input_ids) - len(self.assistant_header):
                if input_ids[curr_idx:curr_idx+len(self.assistant_header)].tolist() == self.assistant_header:
                    header_positions.append(curr_idx)
                    curr_idx += len(self.assistant_header)
                else:
                    curr_idx += 1
            
            # Create mask for assistant responses
            if header_positions:
                loss_mask = torch.zeros_like(input_ids)
                for pos in header_positions:
                    start = pos + len(self.assistant_header)
                    end = self.find_eot(input_ids, start)
                    loss_mask[start:end] = 1
                
                batch["labels"][i] = torch.where(
                    loss_mask.bool(),
                    batch["labels"][i],
                    -100
                )
                
        return batch
    
    def find_eot(self, input_ids, start_idx):
        eot_positions = (input_ids[start_idx:] == self.eot_token).nonzero()
        return start_idx + eot_positions[0].item() if eot_positions.numel() > 0 else len(input_ids)

In [7]:
collator = SpecialTokenCollator(
    assistant_header="<|start_header_id|>assistant<|end_header_id|>",
    eot_token="<|eot_id|>",
    tokenizer=tokenizer,
    mlm=False
)

In [39]:
input_text = df.text.iloc[5]
print(input_text)

<|begin_of_text|><|start_header_id|>system<|end_header_id|>

Cutting Knowledge Date: December 2023
Today Date: 26 Jul 2024

<|eot_id|><|start_header_id|>user<|end_header_id|>

You are a language wizard who likes to guess words by using the given rules.

Welcome to Wordle! You have six attempts to guess the target word, a valid English word of five lowercase letters (a-z). Please use the tags "guess:" and "explanation:" to provide a concise explanation for each guess.

To help you make an informed guess, you will receive a clue for the word, such as
clue: "snowy white"

Here is an example guess based on the clue:
guess: apple
explanation: In the fairy tail Snow White, the girl is killed because she eats a poisoned apple. And the word apple has 5 letters.

I will then indicate whether I agree or disagree with your guess and provide rationale, but agreeing with a guess does not confirm its correctness. You may choose to retain your original guess or modify it based on given clue and agree

In [40]:
# Test mask generation
sample = tokenizer(input_text, return_tensors="pt")
input_ids = sample["input_ids"][0]
collated = collator([{"input_ids": input_ids}])
mask = (collated["labels"][0] != -100).tolist()
print("Mask:", mask)

Mask: [False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False, False,

In [41]:
# Or get only the masked tokens:
numeric_mask = [i for i, mm in enumerate(mask) if mm]
masked_tokens = input_ids[numeric_mask]
print("\nTraining tokens:", tokenizer.decode(masked_tokens))


Training tokens: 

guess: third
explanation: The bronze medal is awarded to the third-place finisher in competitions, so "third" aligns with the clue "bronze place" and is a five-letter word.

guess: medal
explanation: Although "third" fits the clue "bronze place" and you agreed with it, another five-letter word associated with bronze awards is "medal." Since bronze medals are given to third-place finishers, "medal" also aligns with the clue. I'll guess "medal" this time.

guess: third
explanation: Based on the clue "bronze place" and your agreement with my initial guess "third," I believe "third" is still the best guess. From the feedback on "medal," we know that:

- The letter **D** is in the word but not in position 3 (since it was yellow when in position 3).
- Letters **M**, **E**, **A**, and **L** are not in the word.

In "third," the **D** is in position 5, which fits the feedback. Also, all the letters in "third" are valid since they don't include the eliminated letters. Theref

In [36]:
len(numeric_mask)

31