{ "cells": [ { "cell_type": "code", "execution_count": 7, "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import pandas as pd\n", "import torch\n", "from torch.utils.data import Dataset, DataLoader\n", "from datasets import load_dataset\n", "from transformers import AutoTokenizer, AutoConfig, SwitchTransformersForConditionalGeneration\n", "device = \"cuda\" if torch.cuda.is_available() else \"cpu\" " ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Building custom DataLoader" ] }, { "cell_type": "code", "execution_count": 8, "metadata": {}, "outputs": [], "source": [ "\n", "\n", "class CustomDataLoader(Dataset):\n", " def __init__(self, dataframe, tokenizer, source_len, summ_len) -> None:\n", " self.tokenizer = tokenizer\n", " self.data = dataframe\n", " self.source_len = source_len\n", " self.summ_len = summ_len\n", " self.context = self.data['document']\n", " self.summaries = self.data['summaries']\n", "\n", " def __len__(self):\n", " return len(self.context)\n", "\n", " def __getitem__(self, index) :\n", " context = self.context[index]\n", " summary = self.summaries[index]\n", "\n", " source = self.tokenizer.batch_encode_plus([context], max_length = self.source_len, pad_to_max_length=True, return_tensors='pt')\n", " target = self.tokenizer.batch_encode_plus([summary], max_length = self.summ_len, pad_to_max_length=True, return_tensors='pt')\n", "\n", " source_ids = source['input_ids'].squeeze()\n", " source_mask = source['attention_mask'].squeeze()\n", " target_ids = target['input_ids'].squeeze()\n", " target_mask = target['attention_mask'].squeeze()\n", "\n", " return {\n", " \"source_ids\" : source_ids.to(dtype = torch.long),\n", " \"source_mask\": source_mask.to(dtype=torch.long),\n", " \"target_ids\": target_ids.to(dtype=torch.long),\n", " \"target_mask\": target_mask.to(dtype=torch.long)\n", " }" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Training Loop" ] }, { "cell_type": "code", "execution_count": 9, "metadata": {}, "outputs": [], "source": [ "def train(epoch, tokenizer, model, loader, optimizer):\n", " model.train()\n", " for _, data in enumerate(loader, 0):\n", " labels = data['target_ids'].to(device, dtype=torch.long)\n", " labels = model._shift_right(labels)\n", "\n", " labels = labels.masked_fill_(labels == 0, -100)\n", " ids = data['source_ids'].to(device, dtype=torch.long)\n", " mask = data['source_mask'].to(device, dtype=torch.long)\n", " decoder_input_ids = torch.zeros_like(labels).long()\n", "\n", " outputs = model(input_ids = ids, attention_mask = mask, labels=labels, output_router_logits=True, return_dict = True)\n", " loss = outputs[0]\n", "\n", " if _ % 10 == 0:\n", " print(f\"Training Loss: {loss.item()}\")\n", "\n", " if (_+1) % 2000 == 0:\n", " break\n", "\n", " optimizer.zero_grad()\n", " loss.backward()\n", " optimizer.step()" ] }, { "cell_type": "code", "execution_count": 10, "metadata": {}, "outputs": [], "source": [ "from numpy import dtype\n", "\n", "\n", "def validate(epoch, tokenizer, model, device, loader):\n", " model.eval()\n", " predictions = []\n", " actuals = []\n", " with torch.no_grad():\n", " for _, data in enumerate(loader, 0):\n", " y = data['target_ids'].to(device, dtype=torch.long)\n", " ids = data['source_ids'].to(device, dtype=torch.long)\n", " mask = data['source_mask'].to(device, dtype= torch.long)\n", "\n", " generated_ids = model.generate(\n", " input_ids = ids,\n", " attention_mask = mask,\n", " max_length = 150,\n", " num_beams = 2,\n", " repition_penalty=2.5,\n", " length_penalty=1.0,\n", " early_stopping = True\n", " )\n", "\n", " preds = [tokenizer.decode(g, skip_special_tokens=True, clean_up_tokenization_spaces=True) for g in generated_ids]\n", " target = [tokenizer.decode(t, skip_special_tokens=True, claen_up_tokenization_spaces=True) for t in y]\n", " if _ %100 == 0:\n", " print(f\"Completed {_}\")\n", " break\n", "\n", " predictions.extend(preds)\n", " actuals.extend(target)\n", "\n", " return predictions , actuals" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "def main():\n", " train_batch_size = 16\n", " val_batch_size = 16\n", " epochs = 2\n", " val_epochs = 1\n", " learning_rate = 1e-4\n", " max_len = 256\n", " summary_len = 256\n", "\n", " tokenizer = AutoTokenizer.from_pretrained(\"google/switch-base-8\")\n", "\n", " dataset = load_dataset(\"xsum\")\n", " def preprend(example):\n", " return {\"document\": [\"summarization: \" + x for x in example['document']]}\n", " encoded_dataset = dataset.map(preprend,batched=True)\n", "\n", " train_dataset = encoded_dataset['train'] # type: ignore\n", " val_dataset = encoded_dataset['validation'] # type: ignore\n", "\n", " training_set = CustomDataLoader(train_dataset, tokenizer, max_len, summary_len)\n", " val_set = CustomDataLoader(val_dataset, tokenizer, max_len, summary_len)\n", "\n", " train_params = {\n", " 'batch_size' : train_batch_size,\n", " 'shuffle' : True,\n", " 'num_workers': 0\n", " }\n", "\n", " val_params = {\n", " \"batch_size\" : val_batch_size,\n", " \"shuffle\": False,\n", " \"num_workers\": 0\n", " }\n", "\n", " train_loader = DataLoader(training_set,**train_params )\n", " val_loader = DataLoader(val_set, **val_params)\n", "\n", " model = SwitchTransformersForConditionalGeneration.from_pretrained(\"google/switch-base-8\", torch_dtype= torch.bfloat16)\n", " model = model.to(device) # type: ignore\n", "\n", " optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) # type: ignore\n", " print(\"initializing Fine-Tuning \")\n", "\n", " for epoch in range(epochs):\n", " train(epoch, tokenizer, model, train_loader, optimizer)\n", "\n", " print(\"Validating on the fine-tuned model\")\n", " for epoch in range(val_epochs):\n", " predictions , actuals = validate(epoch, tokenizer, model, device, val_loader)\n", " final_df = pd.DataFrame({\"Generated Text\" :predictions, \"Actual Text\":actuals})\n", " final_df.to_csv(\"./output/predictions.csv\")\n", " print(\"Output Files generated\")\n", "\n", " model.save_pretrained(\"switch-transformer\")\n", " tokenizer.save_pretrained(\"switch-transformer-tokenizer\")\n", " return model, tokenizer\n", "\n", "\n", "\n", "trained_model, tokenizer = main()" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "model = SwitchTransformersForConditionalGeneration.from_pretrained(\"switch-transformer\", torch_dtype=torch.bfloat16)\n", "tokenizer = AutoTokenizer.from_pretrained(\"switch-transformer-tokenizer\")" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.12.3" } }, "nbformat": 4, "nbformat_minor": 2 }