Views
No views yet
1context = "This is a long text based of multiple concatenated paragraphs."
2question = "My question about something mentioned inside the context."
3
4model_inputs = tokenizer([f"question: {question} context: {context}"], max_length=512, padding=True, truncation=True)
5input_ids = torch.tensor(model_inputs['input_ids']).to(device)
6attention_mask = torch.tensor(model_inputs['attention_mask']).to(device)
7with torch.no_grad():
8 sample_output = model.generate(input_ids[:1], max_length=85)
9 sample_output_text = tokenizer.decode(sample_output[0], skip_special_tokens=True)
10 input_text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
11 print(f"Sample Input:\n \"{input_text}\"\n\n")
12 print(f"Model Output: \"{sample_output_text}\"")1msg = f"""
2Answer the following question using the content provided in the context.
3Do not answer questions where the answer isn't inside the context.
4
5
6Question: {sample['question']}
7Context: {sample['context']}
8"""1import torch
2from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
3
4# Base model (e.g., T5-large)
5# https://huggingface.co/collections/google/flan-t5-release-65005c39e3201fff885e22fb
6model_name = 'google/flan-t5-small'
7model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
8tokenizer = AutoTokenizer.from_pretrained(model_name)
9
10# Move only the student model to GPU if available
11device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
12model = model.to(device)1from datasets import load_dataset
2
3# Load dataset
4ds = load_dataset('philipp-zettl/long-qa')
5
6# Split the dataset into training and validation
7train_dataset = ds['train']
8validation_dataset = ds['test']1def preprocess_batch(batch, tokenizer, max_input_length=512, max_output_length=128):
2 questions = batch['question']
3 contexts = batch['context']
4 answers = batch['answer']
5
6 inputs = [f"question: {q} context: {c}" for q, c in zip(questions, contexts)]
7 model_inputs = tokenizer(inputs, max_length=max_input_length, padding=True, truncation=True)
8
9 labels = tokenizer(answers, max_length=max_output_length, padding=True, truncation=True)
10 model_inputs['labels'] = labels['input_ids']
11
12 return model_inputs
13
14# Tokenize the dataset
15train_dataset = train_dataset.map(lambda batch: preprocess_batch(batch, tokenizer), batched=True)
16validation_dataset = validation_dataset.map(lambda batch: preprocess_batch(batch, tokenizer), batched=True)
17
18# Set format for PyTorch
19train_dataset.set_format(type='torch', columns=['input_ids', 'attention_mask', 'labels'])
20validation_dataset.set_format(type='torch', columns=['input_ids', 'attention_mask', 'labels'])1from tqdm import tqdm
2from transformers import AdamW, DataCollatorForSeq2Seq
3from torch.utils.data import DataLoader
4from torch.utils.tensorboard import SummaryWriter
5
6torch.cuda.empty_cache()
7
8model_name = 'google/flan-t5-small'
9model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
10tokenizer = AutoTokenizer.from_pretrained(model_name)
11
12# Move only the student model to GPU if available
13device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
14model = model.to(device)
15
16# Training parameters
17epochs = 50
18learning_rate = 3e-5
19temperature = 2.0
20batch_size = 8
21optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)
22
23# Create a data collator for padding and batching
24data_collator = DataCollatorForSeq2Seq(tokenizer=tokenizer, model=model)
25
26# Create DataLoaders with the data collator
27train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, collate_fn=data_collator)
28validation_dataloader = DataLoader(validation_dataset, batch_size=batch_size, collate_fn=data_collator)
29
30writer = SummaryWriter(comment='t5-small-long-qa')
31
32# Store losses and learning rates
33train_losses = []
34val_losses = []
35learning_rates = []
36
37print("Starting training...")
38
39# Training loop
40for epoch in range(epochs):
41 model.train()
42 total_loss = 0
43 print(f"Epoch {epoch+1}/{epochs}")
44
45 progress_bar = tqdm(train_dataloader, desc="Training", leave=False)
46
47 for step, batch in enumerate(progress_bar):
48 # Move student inputs to GPU
49 input_ids = batch['input_ids'].to(device)
50 attention_mask = batch['attention_mask'].to(device)
51 labels = batch['labels'].to(device)
52
53 # Teacher forward pass on CPU
54 outputs = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
55 logits = outputs.logits
56
57 # Calculate losses
58 loss = outputs.loss # Cross-entropy loss
59 writer.add_scalar("Loss/train", loss, epoch * len(train_dataloader) + step)
60
61 # Backpropagation
62 optimizer.zero_grad()
63 loss.backward()
64 optimizer.step()
65
66 total_loss += loss.item()
67
68 # Verbose logging
69 if step % len(train_dataloader)//10 == 1 or step == len(train_dataloader) - 1:
70 progress_bar.set_postfix({
71 'step': step,
72 'loss': loss.item(),
73 })
74
75 # Generate a sample output from the student model
76 model.eval()
77 with torch.no_grad():
78 sample_output = model.generate(input_ids[:1], max_length=50)
79 sample_output_text = tokenizer.decode(sample_output[0], skip_special_tokens=True)
80 input_text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
81 writer.add_text(f"Sample Input", input_text, step)
82 writer.add_text(f"Sample Output", sample_output_text, step)
83 model.train()
84
85
86 avg_train_loss = total_loss / len(train_dataloader)
87 train_losses.append(avg_train_loss)
88 learning_rates.append(optimizer.param_groups[0]['lr'])
89
90 # Validation step
91 model.eval()
92 total_val_loss = 0
93 with torch.no_grad():
94 for batch in validation_dataloader:
95 input_ids = batch['input_ids'].to(device)
96 attention_mask = batch['attention_mask'].to(device)
97 labels = batch['labels'].to(device)
98
99 outputs = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
100 val_loss = outputs.loss
101 total_val_loss += val_loss.item()
102
103 avg_val_loss = total_val_loss / len(validation_dataloader)
104 val_losses.append(avg_val_loss)
105
106 writer.add_scalar("AVG Loss/train", avg_train_loss, epoch)
107 writer.add_scalar("AVG Loss/val", avg_val_loss, epoch)
108
109 print(f"Epoch {epoch+1} completed. Avg Train Loss: {avg_train_loss:.4f}, Avg Val Loss: {avg_val_loss:.4f}")
110
111
112print("Training complete.")
113writer.close()