product_emb / train.py
aakarsh-yadav-tcgls
update train file
1a360e3
Raw
History Blame Contribute Delete
3.77 kB
import os
import json
import pickle
import shutil
import torch
from flask import Flask, send_file
from transformers import AutoTokenizer, AutoModelForCausalLM, Trainer, TrainingArguments, DataCollatorWithPadding
from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions
from torch import nn
from datasets import Dataset
class MultiHeadLatentGPT2(AutoModelForCausalLM):
"""Custom GPT-2 model with Multi-Head Latent representations."""
def __init__(self, config):
super().__init__(config)
self.num_latent_heads = 4 # Define the number of latent heads
self.latent_heads = nn.ModuleList([nn.Linear(config.hidden_size, config.hidden_size) for _ in range(self.num_latent_heads)])
def forward(self, input_ids, attention_mask=None, labels=None):
"""Forward pass with multiple latent heads."""
outputs = super().forward(input_ids, attention_mask=attention_mask, labels=labels)
# Extract hidden states from GPT-2
hidden_states = outputs.hidden_states[-1]
# Process with latent heads
latent_outputs = [head(hidden_states) for head in self.latent_heads]
combined_latent_output = torch.mean(torch.stack(latent_outputs), dim=0)
# Compute loss
loss = None
if labels is not None:
loss_fct = nn.CrossEntropyLoss()
loss = loss_fct(combined_latent_output.view(-1, self.config.vocab_size), labels.view(-1))
return CausalLMOutputWithCrossAttentions(
loss=loss,
logits=combined_latent_output,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions
)
def load_dialogues(file_path):
try:
with open(file_path, 'r', encoding='utf-8') as file:
return json.load(file)
except (FileNotFoundError, json.JSONDecodeError):
return []
def train_model():
"""Train a fine-tuned GPT-2 model using Multi-Head Latent representations."""
dialogues = load_dialogues("whatsapp_chat/dialogues.json")
dataset = Dataset.from_list(dialogues)
tokenizer = AutoTokenizer.from_pretrained("gpt2")
model = MultiHeadLatentGPT2.from_pretrained("gpt2")
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
def tokenize_function(example):
prompt = example["input_text"][:900] + tokenizer.eos_token
completion = example["response"][:900] + tokenizer.eos_token
tokenized = tokenizer(prompt + completion, truncation=True, padding="max_length", max_length=1000)
tokenized["labels"] = tokenized["input_ids"].copy()
return tokenized
tokenized_dataset = dataset.map(tokenize_function, batched=False)
data_collator = DataCollatorWithPadding(tokenizer=tokenizer)
training_args = TrainingArguments(
output_dir="./results",
evaluation_strategy="epoch",
per_device_train_batch_size=4,
num_train_epochs=3,
logging_dir="./logs",
save_total_limit=2
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_dataset,
eval_dataset=tokenized_dataset,
data_collator=data_collator
)
trainer.train()
model.save_pretrained("./fine_tuned_chat_model")
tokenizer.save_pretrained("./fine_tuned_chat_model")
shutil.make_archive("model", 'zip', '.', "fine_tuned_chat_model")
print("Fine-tuned Multi-Head Latent model trained and saved as model.zip")
app = Flask(__name__)
@app.route('/download_model', methods=['GET'])
def download_model():
return send_file("model.zip", as_attachment=True)
if __name__ == "__main__":
train_model()
app.run(host='0.0.0.0', port=5000)