{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os.path\nimport re  # regular expressions\nimport random\nimport copy\nfrom tqdm import tqdm  # progress bars\nfrom collections import namedtuple\nfrom collections import defaultdict\nfrom itertools import pairwise\nfrom time import time\n# multi-threading for reading input files faster:\nfrom threading import Lock\nfrom concurrent.futures import ThreadPoolExecutor\n# sklearn and ML stuff:\nfrom sklearn.metrics import accuracy_score, balanced_accuracy_score\n# viz:\nimport matplotlib.pyplot as plt\n# Hugging Faces and PyTorch\nfrom transformers import AutoModelForSequenceClassification, TrainingArguments, Trainer, AutoTokenizer\nimport torch\nfrom torch.utils.data import IterableDataset, DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-11T19:36:28.074249Z","iopub.execute_input":"2023-07-11T19:36:28.074506Z","iopub.status.idle":"2023-07-11T19:36:43.272338Z","shell.execute_reply.started":"2023-07-11T19:36:28.074481Z","shell.execute_reply":"2023-07-11T19:36:43.271371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = os.path.join(\"/kaggle\", \"input\", \"colie\", \"train\", \"train\")\ndata_dir_valid = os.path.join(\"/kaggle\", \"input\", \"colie\", \"valid\", \"valid\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:36:43.274327Z","iopub.execute_input":"2023-07-11T19:36:43.275195Z","iopub.status.idle":"2023-07-11T19:36:43.279831Z","shell.execute_reply.started":"2023-07-11T19:36:43.275158Z","shell.execute_reply":"2023-07-11T19:36:43.278884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"chunk_fname_parser = re.compile(r\"(\\d+)_(\\d+)\\.txt\")\n\n# train\nchunk_count = {}\nunparseable = []\nfor fname in os.listdir(data_dir):\n    match = chunk_fname_parser.fullmatch(fname)\n    if match is None:\n        unparseable.append(fname)\n        continue\n    book, chunk = match.group(1), match.group(2)\n    if book not in chunk_count:\n        chunk_count[book] = []\n    chunk_count[book].append(chunk)\nprint(f\"Total books train set: {len(chunk_count)}\")\nprint(f\"{len(unparseable)} chunk names could not be parsed.\")\nprint(\"Work in another notebook suggests that \\\"<BOOK_ID> (1)\\\" and \\\"<BOOK_ID>\\\" are duplicates\")\n    \n# valid\nchunk_count_valid = {}\nunparseable = []\nfor fname in os.listdir(data_dir_valid):\n    match = chunk_fname_parser.fullmatch(fname)\n    if match is None:\n        unparseable.append(fname)\n        continue\n    book, chunk = match.group(1), match.group(2)\n    if book not in chunk_count_valid:\n        chunk_count_valid[book] = []\n    chunk_count_valid[book].append(chunk)\nprint(f\"Total books valid set: {len(chunk_count_valid)}\")\nprint(f\"{len(unparseable)} chunks could not be parsed.\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:36:43.281346Z","iopub.execute_input":"2023-07-11T19:36:43.281980Z","iopub.status.idle":"2023-07-11T19:37:08.110584Z","shell.execute_reply.started":"2023-07-11T19:36:43.281948Z","shell.execute_reply":"2023-07-11T19:37:08.109577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CLASSNAME2NUMB = {\n    \"Viktorian\": 0,\n    \"Romantici\": 1,\n    \"Modernism\": 2,\n    \"PostModer\": 3,\n    \"OurDays\": 4\n}","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:08.113372Z","iopub.execute_input":"2023-07-11T19:37:08.113736Z","iopub.status.idle":"2023-07-11T19:37:08.118968Z","shell.execute_reply.started":"2023-07-11T19:37:08.113701Z","shell.execute_reply":"2023-07-11T19:37:08.117384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Do not forget to enable GPU for fast training","metadata":{}},{"cell_type":"markdown","source":"### Data loader","metadata":{}},{"cell_type":"code","source":"tokenizer = AutoTokenizer.from_pretrained(\"roberta-base\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:08.120373Z","iopub.execute_input":"2023-07-11T19:37:08.120836Z","iopub.status.idle":"2023-07-11T19:37:08.952805Z","shell.execute_reply.started":"2023-07-11T19:37:08.120803Z","shell.execute_reply":"2023-07-11T19:37:08.951805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDataset(IterableDataset):\n    def __init__(self):\n        super(TrainDataset).__init__()\n        y_train = pd.read_csv(os.path.join(\"/kaggle\", \"input\", \"colie\", \"train.csv\"))\n        self.removed_chunks_train = y_train.apply(lambda row: (row[0].split(\"_\")[0], row[1]), axis=1, raw=True)\n        self.removed_chunks_train = self.removed_chunks_train.groupby(by=[\"BOOK_id\", \"Epoch\"], as_index=False).size()\n        \n    def _book_chunks_generator(self):\n        for book, book_chunks in chunk_count.items():\n            content = \"\"\n            chunks = random.choices(book_chunks, k=5)\n            book_class = self.removed_chunks_train[self.removed_chunks_train[\"BOOK_id\"] == book].iat[0, 1]\n            label = torch.LongTensor([CLASSNAME2NUMB[book_class]]).cuda()\n#             label = torch.nn.functional.one_hot(torch.LongTensor([CLASSNAME2NUMB[book_class]]), num_classes=5).cuda()\n            for chunk in chunks:\n                try:\n                    with open(os.path.join(data_dir, f\"{book}_{chunk}.txt\"), \"r\", encoding=\"Windows-1252\") as f:\n                        content += f.read()\n                except UnicodeDecodeError:\n                    print(f\"UnicodeDecodeError with {book}_{chunk}.txt\")\n                    break\n            for content_chunk in np.array_split(list(content), max(1, len(content) // 2000)):  # 2000 chars ~= 500 English words\n                tokenized_inputs = tokenizer(\"\".join(content_chunk), return_tensors=\"pt\", padding=\"max_length\", truncation=True, max_length=512)\n                tokenized_inputs = tokenized_inputs['input_ids'].cuda()\n                yield (book, tokenized_inputs.squeeze(), label.squeeze())\n    \n    def __iter__(self):\n        return self._book_chunks_generator()\n\n\nclass ValidDataset(IterableDataset):\n    def __init__(self):\n        super(ValidDataset).__init__()\n        y_valid = pd.read_csv(os.path.join(\"/kaggle\", \"input\", \"colie\", \"valid.csv\"))\n        self.removed_chunks_valid = y_valid.apply(lambda row: (row[0].split(\"_\")[0], row[1]), axis=1, raw=True)\n        self.removed_chunks_valid = self.removed_chunks_valid.groupby(by=[\"BOOK_id\", \"Epoch\"], as_index=False).size()\n        \n    def _book_chunks_generator(self):\n        \"\"\"Difference w.r.t. TrainDataset._book_chunks_generator: sample in random order\n        \"\"\"\n        for book, book_chunks in random.sample(list(chunk_count_valid.items()), k=len(chunk_count_valid.items())):\n            content = \"\"\n            chunks = random.choices(book_chunks, k=5)\n            book_class = self.removed_chunks_valid[self.removed_chunks_valid[\"BOOK_id\"] == book].iat[0, 1]\n            label = torch.LongTensor([CLASSNAME2NUMB[book_class]]).cuda()\n#             label = torch.nn.functional.one_hot(torch.LongTensor([CLASSNAME2NUMB[book_class]]), num_classes=5).cuda()\n            for chunk in chunks:\n                try:\n                    with open(os.path.join(data_dir_valid, f\"{book}_{chunk}.txt\"), \"r\", encoding=\"Windows-1252\") as f:\n                        content += f.read()\n                except UnicodeDecodeError:\n                    print(f\"UnicodeDecodeError with {book}_{chunk}.txt\")\n                    break\n            for content_chunk in np.array_split(list(content), max(1, len(content) // 2000)):  # 2000 chars ~= 500 English words\n                tokenized_inputs = tokenizer(\"\".join(content_chunk), return_tensors=\"pt\", padding=\"max_length\", truncation=True, max_length=512)\n                tokenized_inputs = tokenized_inputs['input_ids'].cuda()\n                yield (book, tokenized_inputs.squeeze(), label.squeeze())\n    \n    def __iter__(self):\n        return self._book_chunks_generator()","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:08.954400Z","iopub.execute_input":"2023-07-11T19:37:08.954774Z","iopub.status.idle":"2023-07-11T19:37:08.974681Z","shell.execute_reply.started":"2023-07-11T19:37:08.954725Z","shell.execute_reply":"2023-07-11T19:37:08.973658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset()\ntrain_dataloader = DataLoader(train_dataset, batch_size=16)\n\nvalid_dataset = ValidDataset()\nvalid_dataloader = DataLoader(valid_dataset, batch_size=128)","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:08.976104Z","iopub.execute_input":"2023-07-11T19:37:08.976715Z","iopub.status.idle":"2023-07-11T19:37:12.800111Z","shell.execute_reply.started":"2023-07-11T19:37:08.976680Z","shell.execute_reply":"2023-07-11T19:37:12.799049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model definition and training","metadata":{}},{"cell_type":"code","source":"net = AutoModelForSequenceClassification.from_pretrained(\n    \"distilroberta-base\", num_labels=5, id2label={v: k for k, v in CLASSNAME2NUMB.items()}, label2id=CLASSNAME2NUMB\n).to(\"cuda\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:12.801548Z","iopub.execute_input":"2023-07-11T19:37:12.802044Z","iopub.status.idle":"2023-07-11T19:37:20.518248Z","shell.execute_reply.started":"2023-07-11T19:37:12.802003Z","shell.execute_reply":"2023-07-11T19:37:20.517218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:20.519600Z","iopub.execute_input":"2023-07-11T19:37:20.519970Z","iopub.status.idle":"2023-07-11T19:37:20.529529Z","shell.execute_reply.started":"2023-07-11T19:37:20.519933Z","shell.execute_reply":"2023-07-11T19:37:20.528234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Check trainable layers (only ``classifier`` head should be trainable for the moment)","metadata":{}},{"cell_type":"code","source":"all([param.requires_grad for param in net.parameters()])","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:20.533512Z","iopub.execute_input":"2023-07-11T19:37:20.534101Z","iopub.status.idle":"2023-07-11T19:37:20.543663Z","shell.execute_reply.started":"2023-07-11T19:37:20.534068Z","shell.execute_reply":"2023-07-11T19:37:20.542546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This is not what we expected, so we will freeze all params of the ``roberta`` sub-module (leaving only ``classifier`` unfreeezed)","metadata":{}},{"cell_type":"code","source":"for param in net.roberta.parameters():\n    param.requires_grad = False\nprint(all([param.requires_grad for param in net.parameters()]))\nprint(all([not param.requires_grad for param in net.roberta.parameters()]))\nprint(all([param.requires_grad for param in net.classifier.parameters()]))","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:20.545370Z","iopub.execute_input":"2023-07-11T19:37:20.546286Z","iopub.status.idle":"2023-07-11T19:37:20.557040Z","shell.execute_reply.started":"2023-07-11T19:37:20.546261Z","shell.execute_reply":"2023-07-11T19:37:20.556070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"ok! Let us proceed to training","metadata":{}},{"cell_type":"code","source":"criterion = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(net.parameters(), lr=0.01, momentum=0.9)\n\nLOGGING_PERIOD = 1000\nVALID_BATCHES_TO_USE = 100\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'max', factor=0.1, patience=3, min_lr=1E-5, verbose=True)\nbest_model_wts = None\nbest_acc = 0\nlearning_curve = []\nlearning_rates = []\n\nstart_time = time()\n\nfor epoch in range(1):\n    running_loss = 0.0\n    for i, data in enumerate(train_dataloader):\n        if i % 50 == 0:\n            print(f'[epoch {epoch + 1}, step {i + 1:5d}, {time()-start_time:.1f}s running]')\n        \n        books, inputs, labels = data\n        \n        # zero the parameter gradients\n        optimizer.zero_grad()\n        \n        # forward + backward + optimize\n        outputs = net(inputs)\n        \n        # NOTE: indexes where we change the book including left and right limits, e.g.  [0,2,4,7] if books=[1,1,2,2,3,3,3]\n        book_change_idxs = [0] + [i for i in range(1, len(books)) if books[i] != books[i-1]] + [len(books)]\n        outputs = torch.stack([torch.mean(outputs[\"logits\"][l: r], 0) for l,r in pairwise(book_change_idxs)])\n        labels = torch.index_select(labels, 0, torch.IntTensor(book_change_idxs[:-1]).cuda())\n#         outputs = outputs[\"logits\"]\n        \n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        # print statistics and update learning rate if needed\n        running_loss += loss.item()\n        if i % LOGGING_PERIOD == LOGGING_PERIOD - 1:    # print every LOGGING_PERIOD mini-batches\n            learning_rates.append(optimizer.param_groups[0]['lr'])\n            # Compute validation set loss and validation set accuracy:\n            val_loss = 0\n            total = 0\n            correct = 0\n            with torch.no_grad():\n                for counter, data in enumerate(valid_dataloader):\n                    if counter % VALID_BATCHES_TO_USE == VALID_BATCHES_TO_USE - 1:\n                        break  # limit validation set for speed purposes\n                    books, inputs, labels = data\n                    outputs = net(inputs)\n                    book_change_idxs = [0] + [i for i in range(1, len(books)) if books[i] != books[i-1]] + [len(books)]\n                    outputs = torch.stack([torch.mean(outputs[\"logits\"][l: r], 0) for l,r in pairwise(book_change_idxs)])\n                    labels = torch.index_select(labels, 0, torch.IntTensor(book_change_idxs[:-1]).cuda())\n#                     outputs = outputs[\"logits\"]\n                    val_loss += criterion(outputs, labels).item()\n                    total += labels.size(0)\n                    correct += (outputs.argmax(dim=1) == labels).sum().item()\n            running_loss /= LOGGING_PERIOD\n            val_loss /= counter\n            val_acc = correct / total\n            print(f'[epoch {epoch + 1}, step {i + 1:5d}]',\n                  f'loss: {running_loss:.3f}',\n                  f'val_loss: {val_loss:.3f}',\n                  f'val_acc: {val_acc:.3f}')\n            scheduler.step(val_acc)\n            learning_curve.append((i, running_loss, val_loss, val_acc))  # good for studying overfitting\n            running_loss = 0.0\n            # deep copy the model\n            if val_acc > best_acc:\n                print(\"Updated best model with current parameters\")\n                best_acc = val_acc\n                best_model_wts = copy.deepcopy(net.state_dict())\n            # SAVE BEST MODEL\n            modelsdirpath = os.path.join(\"/kaggle/working\", \"models\")\n            if not os.path.exists(modelsdirpath):\n                os.mkdir(modelsdirpath)\n            fname = os.path.join(modelsdirpath, \"net.pth\")\n            torch.save(best_model_wts, fname)\n            print(f\"best model saved to {fname}\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T19:37:20.558668Z","iopub.execute_input":"2023-07-11T19:37:20.559469Z","iopub.status.idle":"2023-07-11T20:46:23.898222Z","shell.execute_reply.started":"2023-07-11T19:37:20.559434Z","shell.execute_reply":"2023-07-11T20:46:23.897149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we will finetune the model by unfreezing all layers and using a small learning rate (1E-5)","metadata":{}},{"cell_type":"code","source":"for param in net.roberta.parameters():\n    param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2023-07-11T20:56:12.071963Z","iopub.execute_input":"2023-07-11T20:56:12.072420Z","iopub.status.idle":"2023-07-11T20:56:12.078795Z","shell.execute_reply.started":"2023-07-11T20:56:12.072381Z","shell.execute_reply":"2023-07-11T20:56:12.077297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.SGD(net.parameters(), lr=1E-5, momentum=0.9)\nLOGGING_PERIOD = 1000\n# keep other params\n\n\n# basically repeat training routine, but with no learning rate scheduler this time, and stopping at 2000 steps\n# (the model file has also another name)\nstart_time = time()\n\nfor epoch in range(1):\n    running_loss = 0.0\n    for i, data in enumerate(train_dataloader):\n        if i == 2000:\n            break\n        if i % 50 == 0:\n            print(f'[epoch {epoch + 1}, step {i + 1:5d}, {time()-start_time:.1f}s running]')\n        \n        books, inputs, labels = data\n        \n        # zero the parameter gradients\n        optimizer.zero_grad()\n        \n        # forward + backward + optimize\n        outputs = net(inputs)\n        \n        # NOTE: indexes where we change the book including left and right limits, e.g.  [0,2,4,7] if books=[1,1,2,2,3,3,3]\n        book_change_idxs = [0] + [i for i in range(1, len(books)) if books[i] != books[i-1]] + [len(books)]\n        outputs = torch.stack([torch.mean(outputs[\"logits\"][l: r], 0) for l,r in pairwise(book_change_idxs)])\n        labels = torch.index_select(labels, 0, torch.IntTensor(book_change_idxs[:-1]).cuda())\n#         outputs = outputs[\"logits\"]\n        \n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        # print statistics and update learning rate if needed\n        running_loss += loss.item()\n        if i % LOGGING_PERIOD == LOGGING_PERIOD - 1:    # print every LOGGING_PERIOD mini-batches\n            # Compute validation set loss and validation set accuracy:\n            val_loss = 0\n            total = 0\n            correct = 0\n            with torch.no_grad():\n                for counter, data in enumerate(valid_dataloader):\n                    if counter % VALID_BATCHES_TO_USE == VALID_BATCHES_TO_USE - 1:\n                        break  # limit validation set for speed purposes\n                    books, inputs, labels = data\n                    outputs = net(inputs)\n                    book_change_idxs = [0] + [i for i in range(1, len(books)) if books[i] != books[i-1]] + [len(books)]\n                    outputs = torch.stack([torch.mean(outputs[\"logits\"][l: r], 0) for l,r in pairwise(book_change_idxs)])\n                    labels = torch.index_select(labels, 0, torch.IntTensor(book_change_idxs[:-1]).cuda())\n#                     outputs = outputs[\"logits\"]\n                    val_loss += criterion(outputs, labels).item()\n                    total += labels.size(0)\n                    correct += (outputs.argmax(dim=1) == labels).sum().item()\n            running_loss /= LOGGING_PERIOD\n            val_loss /= counter\n            val_acc = correct / total\n            print(f'[epoch {epoch + 1}, step {i + 1:5d}]',\n                  f'loss: {running_loss:.3f}',\n                  f'val_loss: {val_loss:.3f}',\n                  f'val_acc: {val_acc:.3f}')\n            running_loss = 0.0\n            # deep copy the model\n            if val_acc > best_acc:\n                print(\"Updated best model with current parameters\")\n                best_acc = val_acc\n                best_model_wts = copy.deepcopy(net.state_dict())\n            # SAVE BEST MODEL\n            modelsdirpath = os.path.join(\"/kaggle/working\", \"models\")\n            if not os.path.exists(modelsdirpath):\n                os.mkdir(modelsdirpath)\n            fname = os.path.join(modelsdirpath, \"net_finetuned.pth\")\n            torch.save(best_model_wts, fname)\n            print(f\"best model saved to {fname}\")","metadata":{"execution":{"iopub.status.busy":"2023-07-11T20:56:20.764763Z","iopub.execute_input":"2023-07-11T20:56:20.765625Z","iopub.status.idle":"2023-07-11T21:31:35.451301Z","shell.execute_reply.started":"2023-07-11T20:56:20.765578Z","shell.execute_reply":"2023-07-11T21:31:35.450153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}