{"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":"markdown","source":"# Importing Libraries","metadata":{}},{"cell_type":"code","source":"import copy\nimport os\nimport joblib\nimport time\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport plotly.express as px\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom IPython.core.display import HTML, display\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import train_test_split\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom transformers import (AdamW, AutoConfig, AutoModel, AutoTokenizer,\n                          DataCollatorWithPadding)\n\n# Suppress warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nplt.style.use(\"ggplot\")\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-09T04:59:49.908813Z","iopub.execute_input":"2022-08-09T04:59:49.909133Z","iopub.status.idle":"2022-08-09T04:59:59.933515Z","shell.execute_reply.started":"2022-08-09T04:59:49.909055Z","shell.execute_reply":"2022-08-09T04:59:59.932472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('../input/feedback-prize-effectiveness/train.csv')\ntest_df = pd.read_csv('../input/feedback-prize-effectiveness/test.csv')","metadata":{"execution":{"iopub.status.busy":"2022-08-09T04:59:59.935829Z","iopub.execute_input":"2022-08-09T04:59:59.937071Z","iopub.status.idle":"2022-08-09T05:00:00.340163Z","shell.execute_reply.started":"2022-08-09T04:59:59.937031Z","shell.execute_reply":"2022-08-09T05:00:00.339004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"train: {train_df.shape}\")\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:00.345756Z","iopub.execute_input":"2022-08-09T05:00:00.348389Z","iopub.status.idle":"2022-08-09T05:00:00.377694Z","shell.execute_reply.started":"2022-08-09T05:00:00.348344Z","shell.execute_reply":"2022-08-09T05:00:00.376834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"train: {test_df.shape}\")\ntest_df","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:00.382555Z","iopub.execute_input":"2022-08-09T05:00:00.384655Z","iopub.status.idle":"2022-08-09T05:00:00.403861Z","shell.execute_reply.started":"2022-08-09T05:00:00.384619Z","shell.execute_reply":"2022-08-09T05:00:00.403030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"texts = []\nessay = \"\"\n\nprint(\"RANDOM EXAMPLES\\n\")\nfor essay_id in train_df.essay_id.unique()[:10]:\n    text = open(f'../input/feedback-prize-effectiveness/train/{essay_id}.txt').read()\n    essay += f'<td style=\"vertical-align:top; border-right: 1px solid #7accd8\">{text[:200]}</td>'\n    \ndisplay(HTML(f\"\"\"\n<table style=\"font-family: monospace;\">\n    <tr>\n         {essay}\n    </tr>\n</table>\n\"\"\"))\n\ndel essay","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:00:00.407630Z","iopub.execute_input":"2022-08-09T05:00:00.409729Z","iopub.status.idle":"2022-08-09T05:00:00.485545Z","shell.execute_reply.started":"2022-08-09T05:00:00.409694Z","shell.execute_reply":"2022-08-09T05:00:00.484582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Discourse Type\n\nEach essay element contains discourse type metadata. There are 7 discourse_type values with explainations taken from the data page.\n\n- `Lead` - an introduction that begins with a statistic, a quotation, a description, or some other device to grab the reader’s attention and point toward the thesis\n- `Position` - an opinion or conclusion on the main question\n- `Claim` - a claim that supports the position\n- `Counterclaim` - a claim that refutes another claim or gives an opposing reason to the position\n- `Rebuttal` - a claim that refutes a counterclaim\n- `Evidence` - ideas or examples that support claims, counterclaims, or rebuttals.\n- `Concluding` Statement - a concluding statement that restates the claims.","metadata":{}},{"cell_type":"code","source":"labels = ['Adequate', 'Effective', 'Ineffective']\n\nfig, axes = plt.subplots(1, 2, sharey=True, figsize=(22, 6))\n# plot Discourse type\nax = axes[0]\nsns.countplot(x=\"discourse_type\", data=train_df, linewidth=1.25, alpha=1, ax=ax, zorder=2)\nax.set_title(\"Discourse type distribution\")\n\n# plot Discourse effectiveness\nax = axes[1]\nsns.countplot(x=\"discourse_effectiveness\", data=train_df, ax=ax)\nax.set_title(\"Discourse Effectiveness distribution\")\n\nfig.show()\n\n\n# plot Discourse Effectiveness distribution per Discourse Type\ndiscourse_types = train_df.discourse_type.unique()\n\nfig, axes = plt.subplots(2, 4, sharex='col', sharey='row', figsize=(25, 10))\nfor i, discourse_type in enumerate(discourse_types):\n    ax = axes.flatten()[i]\n    filtered_df = train_df[train_df.discourse_type == discourse_type]\n    sns.countplot(x=\"discourse_effectiveness\", data=filtered_df, ax=ax, order=labels)\n    ax.set_title(discourse_type)\n    ax.set(xlabel=\"Discourse Effectiveness\", ylabel=None)\n    \nfig.delaxes(axes[1,3])\nfig.suptitle('Discourse Effectiveness distribution per Discourse Type', fontsize=15)\nplt.show()\n\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:00:00.489872Z","iopub.execute_input":"2022-08-09T05:00:00.492228Z","iopub.status.idle":"2022-08-09T05:00:02.063050Z","shell.execute_reply.started":"2022-08-09T05:00:00.492186Z","shell.execute_reply":"2022-08-09T05:00:02.061906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_examples_for_discourse_type(discourse_type):\n    filt = train_df[train_df.discourse_type==\"Lead\"].sample(frac=1, random_state=420)\n    display(HTML(\n        f\"\"\"\n        <h4><code>{discourse_type}</code> examples</h4>\n        <table>\n            <tr>\n              <th width=33%>Ineffective</th>\n              <th width=33%>Adequate</th>\n              <th width=33%>Effective</th>\n            </tr>\n            <tr>\n              <td>{filt[filt.discourse_effectiveness=='Ineffective'].iloc[0].discourse_text}</td>\n              <td>{filt[filt.discourse_effectiveness=='Adequate'].iloc[0].discourse_text}</td>\n              <td>{filt[filt.discourse_effectiveness=='Effective'].iloc[0].discourse_text}</td>\n            </tr>\n        </table>\n        \"\"\"\n    ))\n    \n\n[show_examples_for_discourse_type(dt) for dt in discourse_types];\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:00:02.064781Z","iopub.execute_input":"2022-08-09T05:00:02.065501Z","iopub.status.idle":"2022-08-09T05:00:02.147980Z","shell.execute_reply.started":"2022-08-09T05:00:02.065462Z","shell.execute_reply":"2022-08-09T05:00:02.146938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Configuration","metadata":{}},{"cell_type":"code","source":"CONFIG = {\n    \"seed\": 69,\n    \"epochs\": 5,\n    \"model_name\": \"microsoft/deberta-v3-base\",\n    \"drop\": 0.2,\n    \"fast_tokenizer\": True,\n    \"freeze\": False,\n    \"n_accumulate\": 3,\n    \"train_batch_size\": 8,\n    \"valid_batch_size\": 16,\n    \"max_length\": 512,\n    \"learning_rate\": 3e-5,\n    \"min_lr\": 1e-6,\n    \"T_max\": 500,\n    \"weight_decay\": 0.01,\n    \"num_classes\": 3,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n  }\n\nCONFIG[\"tokenizer\"] = AutoTokenizer.from_pretrained(\n    CONFIG['model_name'], use_fast=CONFIG[\"fast_tokenizer\"])\n\nTRAIN_DIR = \"../input/feedback-prize-effectiveness/train\"\nTEST_DIR = \"../input/feedback-prize-effectiveness/test\"","metadata":{"execution":{"iopub.status.busy":"2022-08-10T03:27:24.473279Z","iopub.execute_input":"2022-08-10T03:27:24.473908Z","iopub.status.idle":"2022-08-10T03:27:24.495979Z","shell.execute_reply.started":"2022-08-10T03:27:24.473874Z","shell.execute_reply":"2022-08-10T03:27:24.494476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n#   for REPRODUCIBILITY.\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:40.656632Z","iopub.execute_input":"2022-08-09T05:00:40.657098Z","iopub.status.idle":"2022-08-09T05:00:40.665430Z","shell.execute_reply.started":"2022-08-09T05:00:40.657054Z","shell.execute_reply":"2022-08-09T05:00:40.664447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"encoder = LabelEncoder()\ntrain_df['discourse_effectiveness_label'] = encoder.fit_transform(train_df['discourse_effectiveness'])\n\nwith open(\"le.pkl\", \"wb\") as fp:\n    joblib.dump(encoder, fp)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:40.666963Z","iopub.execute_input":"2022-08-09T05:00:40.668889Z","iopub.status.idle":"2022-08-09T05:00:40.691533Z","shell.execute_reply.started":"2022-08-09T05:00:40.668847Z","shell.execute_reply":"2022-08-09T05:00:40.690327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeedBackDataset(Dataset):\n    def __init__(self, df, tokenizer, max_length, data_path):\n        self.df = df\n        self.max_len = max_length\n        self.tokenizer = tokenizer\n        self.data_path = data_path\n        self.discourse_text = df['discourse_text'].values\n        self.essay_id = df['essay_id'].values\n        self.discourse_type = df['discourse_type'].values\n        self.targets = df['discourse_effectiveness_label'].values\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        discourse_text = self.discourse_text[index]\n        discourse_type = self.discourse_type[index]\n        essay_path = os.path.join(\n            self.data_path, f\"{self.essay_id[index]}.txt\")\n        essay = open(essay_path, 'r').read()\n\n        text = discourse_type + \" \" + self.tokenizer.sep_token + \\\n            discourse_text + \" \" + self.tokenizer.sep_token + \" \" + essay\n        inputs = self.tokenizer.encode_plus(\n            text,\n            truncation=True,\n            add_special_tokens=True,\n            max_length=self.max_len\n        )\n\n        return {\n            'input_ids': inputs['input_ids'],\n            'attention_mask': inputs['attention_mask'],\n            'target': self.targets[index]\n        }\n","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:40.694492Z","iopub.execute_input":"2022-08-09T05:00:40.695103Z","iopub.status.idle":"2022-08-09T05:00:40.704811Z","shell.execute_reply.started":"2022-08-09T05:00:40.695044Z","shell.execute_reply":"2022-08-09T05:00:40.703865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Defining Model","metadata":{}},{"cell_type":"code","source":"class MeanPooling(nn.Module):\n    def __init__(self):\n        super(MeanPooling, self).__init__()\n        \n    def forward(self, last_hidden_state, attention_mask):\n        input_mask_expanded = attention_mask.unsqueeze(-1).expand(last_hidden_state.size()).float()\n        sum_embeddings = torch.sum(last_hidden_state * input_mask_expanded, 1)\n        sum_mask = input_mask_expanded.sum(1)\n        sum_mask = torch.clamp(sum_mask, min=1e-9)\n        mean_embeddings = sum_embeddings / sum_mask\n        return mean_embeddings\n    \n\nclass FeedBackModel(nn.Module):\n    def __init__(self, model_name, dropout, freeze=False):\n        super(FeedBackModel, self).__init__()\n        self.drop = nn.Dropout(p=dropout)\n\n        self.model = AutoModel.from_pretrained(model_name)\n        if freeze:\n            for parameter in self.model.parameters():\n                parameter.requires_grad = False\n        \n        self.config = AutoConfig.from_pretrained(model_name)\n        self.mpool = MeanPooling()\n        self.fc = nn.Sequential(\n            nn.Linear(self.config.hidden_size, CONFIG['num_classes']),\n        )\n        \n    def forward(self, ids, mask):        \n        out = self.model(input_ids=ids,attention_mask=mask,\n                         output_hidden_states=False)\n        out = self.mpool(out.last_hidden_state, mask)\n        out = self.drop(out)\n        out = self.fc(out)\n#         out = self.drop(out)\n#         out = self.fc2\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:40.706224Z","iopub.execute_input":"2022-08-09T05:00:40.706888Z","iopub.status.idle":"2022-08-09T05:00:40.720437Z","shell.execute_reply.started":"2022-08-09T05:00:40.706850Z","shell.execute_reply":"2022-08-09T05:00:40.719532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device):\n    model.train()\n\n    total = 0\n    running_loss = 0.0\n    correct = 0\n    lr = []\n    bar = tqdm(dataloader, total=len(dataloader))\n    steps = len(dataloader)\n    for step, data in enumerate(bar):\n        ids = data[\"input_ids\"].to(device, dtype=torch.long)\n        mask = data[\"attention_mask\"].to(device, dtype=torch.long)\n        targets = data[\"target\"].to(device, dtype=torch.long)\n\n        batch_size = ids.size(0)\n\n        outputs = model(ids, mask)\n\n        loss = criterion(outputs, targets)\n        loss = loss / CONFIG['n_accumulate']\n        loss.backward()\n        \n        if (step + 1) % CONFIG['n_accumulate'] == 0 or step == steps:\n            optimizer.step()\n            optimizer.zero_grad()\n            if scheduler:\n                scheduler.step()\n\n        running_loss += loss.item() * batch_size\n        total += batch_size\n\n        _, predictions = outputs.max(1)\n        correct += (predictions == targets).float().sum().item()\n\n        epoch_loss = running_loss / total\n        acc = correct / total\n\n        bar.set_postfix(\n            Loss=epoch_loss, Accuracy=acc*100, LR=optimizer.param_groups[0]['lr'])\n        \n        lr.append(optimizer.param_groups[0]['lr'])\n\n    return epoch_loss, acc, lr\n\n\n@torch.no_grad()\ndef evaluate(model, dataloader, device):\n    model.eval()\n\n    total = 0\n    running_loss = 0.0\n    correct = 0\n\n    for data in dataloader:\n        ids = data[\"input_ids\"].to(device, dtype=torch.long)\n        mask = data[\"attention_mask\"].to(device, dtype=torch.long)\n        targets = data[\"target\"].to(device, dtype=torch.long)\n\n        batch_size = ids.size(0)\n\n        outputs = model(ids, mask)\n\n        loss = criterion(outputs, targets)\n\n        running_loss += loss.item() * batch_size\n        total += batch_size\n        \n        _, predictions = outputs.max(1)\n        correct += (predictions == targets).float().sum().item()\n\n    epoch_loss = running_loss / total\n    acc = correct / total\n\n    print(\"Validation Loss: {:.4f} Accuracy: {:.2f}%\".format(epoch_loss, acc * 100))\n    return epoch_loss, acc","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:00:40.722077Z","iopub.execute_input":"2022-08-09T05:00:40.722784Z","iopub.status.idle":"2022-08-09T05:00:40.738354Z","shell.execute_reply.started":"2022-08-09T05:00:40.722743Z","shell.execute_reply":"2022-08-09T05:00:40.737353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def start_training(model, optimizer, scheduler, device, num_epochs):\n    start = time.time()\n    best_epoch_loss = np.inf\n    history = {\"Train Loss\": [], \"Valid Loss\": [], \"Train Acc\": [], \"Valid Acc\": [], \"LR\": []}\n\n    for epoch in range(1, num_epochs + 1):\n        print(\"Epoch: \", epoch)\n        train_epoch_loss, train_epoch_acc, epoch_lr = train_one_epoch(\n            model, optimizer, scheduler, dataloader=train_loader, device=CONFIG[\"device\"]\n        )\n\n        val_epoch_loss, valid_epoch_acc = evaluate(\n            model, valid_loader, device=CONFIG[\"device\"]\n        )\n\n        history[\"Train Loss\"].append(train_epoch_loss)\n        history[\"Valid Loss\"].append(val_epoch_loss)\n        history[\"Train Acc\"].append(train_epoch_acc)\n        history[\"Valid Acc\"].append(valid_epoch_acc)\n        history[\"LR\"].extend(epoch_lr)\n\n        # deep copy the model\n        if val_epoch_loss <= best_epoch_loss:\n            print(\n                f\"Validation Loss Improved ({best_epoch_loss} ---> {val_epoch_loss})\"\n            )\n            best_epoch_loss = val_epoch_loss\n            best_epoch_acc = valid_epoch_acc\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"best_feedback.bin\"\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            print(f\"Model Saved\")\n\n        print()\n\n    end = time.time()\n    time_elapsed = end - start\n    print(\n        \"Training complete in {:.0f}h {:.0f}m {:.0f}s\".format(\n            time_elapsed // 3600,\n            (time_elapsed % 3600) // 60,\n            (time_elapsed % 3600) % 60,\n        )\n    )\n    print(\n        \"Best Loss: {:.4f} Best Accuracy: {:.2f}\".format(\n            best_epoch_loss, best_epoch_acc * 100\n        )\n    )\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n\n    return model, history","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:00:40.740379Z","iopub.execute_input":"2022-08-09T05:00:40.740899Z","iopub.status.idle":"2022-08-09T05:00:40.753368Z","shell.execute_reply.started":"2022-08-09T05:00:40.740863Z","shell.execute_reply":"2022-08-09T05:00:40.752456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train & Valid SET","metadata":{}},{"cell_type":"code","source":"# df_train = train_df.sample(frac=0.8, random_state=42)\n# df_valid = train_df.drop(df_train.index) \n\ndf_train, df_valid = train_test_split(\n    train_df, test_size=0.2, random_state=42, stratify = train_df.discourse_effectiveness_label)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:52.086943Z","iopub.execute_input":"2022-08-09T05:00:52.087355Z","iopub.status.idle":"2022-08-09T05:00:52.121076Z","shell.execute_reply.started":"2022-08-09T05:00:52.087318Z","shell.execute_reply":"2022-08-09T05:00:52.120165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape, df_valid.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:52.793531Z","iopub.execute_input":"2022-08-09T05:00:52.794498Z","iopub.status.idle":"2022-08-09T05:00:52.801965Z","shell.execute_reply.started":"2022-08-09T05:00:52.794453Z","shell.execute_reply":"2022-08-09T05:00:52.800786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntrain_dataset = FeedBackDataset(\n    df_train, tokenizer=CONFIG[\"tokenizer\"], max_length=CONFIG[\"max_length\"], data_path=TRAIN_DIR\n)\nvalid_dataset = FeedBackDataset(\n    df_valid, tokenizer=CONFIG[\"tokenizer\"], max_length=CONFIG[\"max_length\"], data_path=TRAIN_DIR\n)\n\ncollate_fn = DataCollatorWithPadding(tokenizer=CONFIG['tokenizer'])\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CONFIG[\"train_batch_size\"],\n    collate_fn=collate_fn,\n    num_workers=2,\n    shuffle=True,\n    pin_memory=True,\n)\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=CONFIG[\"valid_batch_size\"],\n    collate_fn=collate_fn,\n    num_workers=2,\n    shuffle=False,\n    pin_memory=True,\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-09T05:00:53.788870Z","iopub.execute_input":"2022-08-09T05:00:53.789233Z","iopub.status.idle":"2022-08-09T05:00:53.797775Z","shell.execute_reply.started":"2022-08-09T05:00:53.789203Z","shell.execute_reply":"2022-08-09T05:00:53.796441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# BeGin Training","metadata":{}},{"cell_type":"code","source":"def get_freezed_parameters(module):\n    \"\"\"\n    Returns names of freezed parameters of the given module.\n    \"\"\"\n    \n    freezed_parameters = []\n    for name, parameter in module.named_parameters():\n        if not parameter.requires_grad:\n            freezed_parameters.append(name)\n            \n    return freezed_parameters","metadata":{"execution":{"iopub.status.busy":"2022-08-05T04:51:18.426609Z","iopub.execute_input":"2022-08-05T04:51:18.427099Z","iopub.status.idle":"2022-08-05T04:51:18.435380Z","shell.execute_reply.started":"2022-08-05T04:51:18.427051Z","shell.execute_reply":"2022-08-05T04:51:18.433186Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = FeedBackModel(CONFIG['model_name'], CONFIG[\"drop\"], CONFIG[\"freeze\"])\nmodel.to(CONFIG['device'])\n\nmodel_parameters = filter(lambda parameter: parameter.requires_grad, model.parameters())\noptimizer = AdamW(model_parameters, lr=CONFIG['learning_rate'])\nscheduler = lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=CONFIG['T_max'], eta_min=CONFIG['min_lr'])\ncriterion = nn.CrossEntropyLoss()\n\nmodel, history = start_training(\n    model, optimizer, scheduler, device=CONFIG['device'], num_epochs=CONFIG['epochs'])\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-09T05:01:28.016126Z","iopub.execute_input":"2022-08-09T05:01:28.016534Z","iopub.status.idle":"2022-08-09T05:01:56.478805Z","shell.execute_reply.started":"2022-08-09T05:01:28.016501Z","shell.execute_reply":"2022-08-09T05:01:56.476182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_history(history):\n\n    plt.figure(figsize=(20,6))\n#     fig, axes = plt.subplots(1, 2, sharey=True, figsize=(22, 6))\n    \n    plt.subplot(1,2,1)\n    for k in [\"Train Loss\", \"Valid Loss\"]:\n        plt.plot(history[k])\n\n    plt.title('Loss')\n    plt.xlabel('epochs')\n    plt.ylabel('loss')\n    plt.legend(['train', 'valid'], loc='upper left')\n\n    plt.subplot(1,2,2)\n    for k in [\"Train Acc\", \"Valid Acc\"]:\n        plt.plot(history[k])\n\n    plt.title('Accuracy')\n    plt.xlabel('epochs')\n    plt.ylabel('accuracy')\n    plt.legend(['train', 'valid'], loc='upper left')\n    \n    plt.show()\n\n\n\nplot_history(history)","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**LB Submission:**\nhttps://www.kaggle.com/code/anantgupt/feedback-infer-submission","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}