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)