{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!pip install torch==2.3.1 --index-url https://download.pytorch.org/whl/cu121\n",
    "!pip install transformers==4.37.0\n",
    "!pip install datasets==2.21.0\n",
    "!pip install accelerate==0.21.0\n",
    "!pip install rouge==1.0.1\n",
    "!pip install tqdm==4.66.5\n",
    "!pip install jieba==0.42.1"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from transformers import AutoTokenizer\n",
    "from transformers import GPT2LMHeadModel\n",
    "from datasets import load_dataset\n",
    "from tqdm import tqdm\n",
    "import torch\n",
    "from torch.utils.tensorboard import SummaryWriter\n",
    "from rouge import Rouge\n",
    "import jieba"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class LCSTSDataset(torch.utils.data.Dataset):\n",
    "    def __init__(self, raw_data) -> None:\n",
    "        super().__init__()\n",
    "        self.data = raw_data\n",
    "        # To prevent out-of-vocabulary tokens from being transformed into [UNK]\n",
    "        self.token_replacement = [\n",
    "            [\"：\", \":\"],\n",
    "            [\"，\", \",\"],\n",
    "            [\"“\", '\"'],\n",
    "            [\"”\", '\"'],\n",
    "            [\"？\", \"?\"],\n",
    "            [\"……\", \"...\"],\n",
    "            [\"！\", \"!\"],\n",
    "        ]\n",
    "\n",
    "    def __getitem__(self, index):\n",
    "        d = self.data[index]\n",
    "        # Substitute some full-width punctuations with half-width ones\n",
    "        for k in d:\n",
    "            for tok in self.token_replacement:\n",
    "                d[k] = d[k].replace(tok[0], tok[1])\n",
    "        return d\n",
    "\n",
    "    def __len__(self):\n",
    "        return len(self.data)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# `pad_token_id`=`tokenizer.eos_token_id`:\n",
    "# For each batch, first finished sentences should have <|endoftext|> rather than [PAD] at the end.\n",
    "# Check more details from the following link.\n",
    "# https://github.com/huggingface/transformers/blob/b880508440f43f80e35a78ccd2a32f3bde91cb23/src/transformers/generation_utils.py#L1248-L1251\n",
    "\n",
    "# `max_new_tokens`: If you don’t set max_new_tokens,\n",
    "# Hugging Face will also count the input tokens!\n",
    "\n",
    "def do_evaluate(tokenizer, model, validation_loader, rouge_metric, inner_check=False):\n",
    "    pbar = tqdm(validation_loader)\n",
    "    pbar.set_description(f\"Evaluating\")\n",
    "\n",
    "    predictions = []\n",
    "    references = []\n",
    "    count = 0\n",
    "    for ground_truth, inputs in pbar:\n",
    "        output = [\n",
    "            s.split(\"[SEP]\")[1].replace(\" \", \"\").split(\"<|endoftext|>\")[0]\n",
    "            for s in tokenizer.batch_decode(\n",
    "                model.generate(\n",
    "                    **inputs,\n",
    "                    max_new_tokens=200, # Maximum number of tokens to generate\n",
    "                    pad_token_id=tokenizer.eos_token_id,\n",
    "                )\n",
    "            )\n",
    "        ]\n",
    "        targets = [\n",
    "            s.split(\"[SEP]\")[1].replace(\" \", \"\").replace(\"<|endoftext|>\", \"\")\n",
    "            for s in tokenizer.batch_decode(ground_truth[\"input_ids\"])\n",
    "        ]\n",
    "        assert len(output) == len(targets)\n",
    "        output = [\" \"] if output == [\"\"] else output\n",
    "        # We use jieba to perform word-level evaluations with ROUGE\n",
    "        predictions.extend([\" \".join(jieba.lcut(o)) for o in output])\n",
    "        references.extend([\" \".join(jieba.lcut(t)) for t in targets])\n",
    "        count += 1\n",
    "        if count > 100 and inner_check:\n",
    "            break # During training, we only evaluate the first 100 examples.\n",
    "\n",
    "    score = rouge_metric.get_scores(predictions, references, avg=True)\n",
    "    if inner_check:\n",
    "        print(\"Validation using 100 examples: \", score)\n",
    "    else:\n",
    "        print(score)\n",
    "\n",
    "    return score, predictions, references"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def collate_fn(batch):\n",
    "    complete_text = [\n",
    "        f\"[CLS]{example['text']}[SEP]{example['summary']}<|endoftext|>\"\n",
    "        for example in batch\n",
    "    ]\n",
    "    complete_text = tokenizer.batch_encode_plus(\n",
    "        complete_text,\n",
    "        padding=True,\n",
    "        truncation=True,\n",
    "        return_tensors=\"pt\",\n",
    "        add_special_tokens=False,\n",
    "    )\n",
    "    # Set label padding tokens to -100 for loss masking\n",
    "    labels = torch.where(\n",
    "        condition=complete_text.input_ids != tokenizer.pad_token_id,\n",
    "        input=complete_text.input_ids,\n",
    "        other=-100,\n",
    "    )\n",
    "    complete_text[\"labels\"] = labels\n",
    "    complete_text = {k: complete_text[k].to(device) for k in complete_text}\n",
    "\n",
    "    infer_text = [example[\"text\"] for example in batch]\n",
    "    infer_text = tokenizer.batch_encode_plus(\n",
    "        infer_text,\n",
    "        padding=True,\n",
    "        truncation=True,\n",
    "        return_tensors=\"pt\",\n",
    "    )\n",
    "    infer_text = {k: infer_text[k].to(device) for k in infer_text}\n",
    "    return complete_text, infer_text"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TRAIN_BATCH_SIZE = 32\n",
    "VAL_BATCH_SIZE = 1 # During evaluation, we don't pad the input.\n",
    "NUM_EPOCHS = 3\n",
    "LR = 1e-5\n",
    "SAVED_DIR = \"saved_models\"\n",
    "model_name = \"uer/gpt2-chinese-cluecorpussmall\"\n",
    "\n",
    "# TensorBoard writer\n",
    "writer = SummaryWriter(f\"runs/{SAVED_DIR}_test_bs{VAL_BATCH_SIZE}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "tokenizer = AutoTokenizer.from_pretrained(\n",
    "    model_name,\n",
    "    padding_side=\"left\", # Use left padding for GPT2\n",
    ")\n",
    "# You can set your device id instead of cuda:0\n",
    "device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Sometimes checking the Hugging Face dataset is slow,\n",
    "# it will be faster if we transform the dataset object into a list using .to_list(). \n",
    "\n",
    "raw_train = load_dataset(\n",
    "    \"hugcyp/LCSTS\", split=\"train\", cache_dir=\"./cache/\"\n",
    ").to_list()\n",
    "raw_val = load_dataset(\n",
    "    \"hugcyp/LCSTS\", split=\"validation\", cache_dir=\"./cache/\"\n",
    ").to_list()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "train_set = LCSTSDataset(raw_train)\n",
    "val_set = LCSTSDataset(raw_val)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "train_loader = torch.utils.data.DataLoader(\n",
    "    train_set,\n",
    "    batch_size=TRAIN_BATCH_SIZE,\n",
    "    shuffle=True,\n",
    "    collate_fn=collate_fn,\n",
    ")\n",
    "val_loader = torch.utils.data.DataLoader(\n",
    "    val_set,\n",
    "    batch_size=VAL_BATCH_SIZE,\n",
    "    shuffle=False,\n",
    "    collate_fn=collate_fn,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# `resize_token_embeddings`:\n",
    "# Increasing the size will add newly initialized vectors at the end. \n",
    "\n",
    "model = GPT2LMHeadModel.from_pretrained(model_name)\n",
    "tokenizer.add_special_tokens({\"eos_token\": \"<|endoftext|>\"}) # Add a new eos token\n",
    "model.resize_token_embeddings(len(tokenizer))\n",
    "model = model.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Set up the optimizer and the evaluation metric\n",
    "optimizer = torch.optim.AdamW(model.parameters(), lr=LR)\n",
    "rouge_metric = Rouge()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "step_i = 0\n",
    "for epoch in range(NUM_EPOCHS):\n",
    "    pbar = tqdm(train_loader)\n",
    "    pbar.set_description(f\"Training epoch [{epoch+1}/{NUM_EPOCHS}]\")\n",
    "    for inputs, _ in pbar:\n",
    "        optimizer.zero_grad()\n",
    "        loss = model(**inputs).loss\n",
    "        loss.backward()\n",
    "        optimizer.step()\n",
    "        pbar.set_postfix(loss=loss.item())\n",
    "        # Log the loss to TensorBoard\n",
    "        writer.add_scalar(\"Loss/train\", loss.item(), step_i)\n",
    "\n",
    "        if step_i % 1000 == 0 and step_i != 0: # Evaluate every 1000 steps\n",
    "            score, pres, refs = do_evaluate(\n",
    "                tokenizer=tokenizer,\n",
    "                model=model,\n",
    "                validation_loader=val_loader,\n",
    "                rouge_metric=rouge_metric,\n",
    "                inner_check=True,\n",
    "            )\n",
    "            print(f\"Rouge scores on step{step_i} of epoch {epoch}:\", score)\n",
    "            print(\"Predictions:\", pres[:5]) # Check the first 5 predictions\n",
    "            print(\"References:\", refs[:5])  # Check the first 5 references\n",
    "            writer.add_scalar(\"Rouge-1/val\", score[\"rouge-1\"][\"f\"], step_i)\n",
    "            writer.add_scalar(\"Rouge-2/val\", score[\"rouge-2\"][\"f\"], step_i)\n",
    "\n",
    "        step_i += 1\n",
    "    score, pres, refs = do_evaluate(\n",
    "        tokenizer=tokenizer,\n",
    "        model=model,\n",
    "        validation_loader=val_loader,\n",
    "        rouge_metric=rouge_metric,\n",
    "    )\n",
    "    torch.save(model, f\"{SAVED_DIR}/ep{epoch}.ckpt\")\n",
    "\n",
    "tokenizer.save_pretrained(f\"{SAVED_DIR}/tokenizer\")"
   ]
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
