Spaces:
Paused
Paused
| 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__) | |
| 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) | |