{"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, os, time, math\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\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 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\")\n\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-07-14T11:23:07.892499Z","iopub.execute_input":"2022-07-14T11:23:07.893091Z","iopub.status.idle":"2022-07-14T11:23:15.342853Z","shell.execute_reply.started":"2022-07-14T11:23:07.892954Z","shell.execute_reply":"2022-07-14T11:23:15.341677Z"},"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-07-14T11:23:15.345005Z","iopub.execute_input":"2022-07-14T11:23:15.345896Z","iopub.status.idle":"2022-07-14T11:23:15.635291Z","shell.execute_reply.started":"2022-07-14T11:23:15.345856Z","shell.execute_reply":"2022-07-14T11:23:15.634184Z"},"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-07-14T11:23:15.637310Z","iopub.execute_input":"2022-07-14T11:23:15.637797Z","iopub.status.idle":"2022-07-14T11:23:15.663511Z","shell.execute_reply.started":"2022-07-14T11:23:15.637752Z","shell.execute_reply":"2022-07-14T11:23:15.661198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"train: {test_df.shape}\")\ntest_df","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:23:15.666470Z","iopub.execute_input":"2022-07-14T11:23:15.667448Z","iopub.status.idle":"2022-07-14T11:23:15.681814Z","shell.execute_reply.started":"2022-07-14T11:23:15.667383Z","shell.execute_reply":"2022-07-14T11:23:15.680783Z"},"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-07-14T11:23:15.683151Z","iopub.execute_input":"2022-07-14T11:23:15.684079Z","iopub.status.idle":"2022-07-14T11:23:15.746113Z","shell.execute_reply.started":"2022-07-14T11:23:15.684045Z","shell.execute_reply":"2022-07-14T11:23:15.744843Z"},"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-07-14T11:23:15.748214Z","iopub.execute_input":"2022-07-14T11:23:15.749020Z","iopub.status.idle":"2022-07-14T11:23:16.983463Z","shell.execute_reply.started":"2022-07-14T11:23:15.748968Z","shell.execute_reply":"2022-07-14T11:23:16.982406Z"},"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-07-14T11:23:16.985118Z","iopub.execute_input":"2022-07-14T11:23:16.985739Z","iopub.status.idle":"2022-07-14T11:23:17.078264Z","shell.execute_reply.started":"2022-07-14T11:23:16.985697Z","shell.execute_reply":"2022-07-14T11:23:17.077047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Configuration","metadata":{}},{"cell_type":"code","source":"CONFIG = {\n    \"seed\": 42,\n    \"epochs\": 5,\n    \"model_name\": \"microsoft/deberta-v3-base\",\n    \"train_batch_size\": 12,\n    \"valid_batch_size\": 16,\n    \"max_length\": 512,\n    \"learning_rate\": 1e-5,\n#     \"min_lr\": 1e-6,\n#     \"T_max\": 500, \n#     \"weight_decay\": 1e-6,\n    \"num_classes\": 3,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n  }\n\nCONFIG[\"tokenizer\"] = AutoTokenizer.from_pretrained(CONFIG['model_name'])\n\nTRAIN_DIR = \"../input/feedback-prize-effectiveness/train\"\nTEST_DIR = \"../input/feedback-prize-effectiveness/test\"","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:23:17.080016Z","iopub.execute_input":"2022-07-14T11:23:17.080629Z","iopub.status.idle":"2022-07-14T11:23:29.226859Z","shell.execute_reply.started":"2022-07-14T11:23:17.080593Z","shell.execute_reply":"2022-07-14T11:23:29.225498Z"},"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-07-14T11:23:29.228822Z","iopub.execute_input":"2022-07-14T11:23:29.229230Z","iopub.status.idle":"2022-07-14T11:23:29.240611Z","shell.execute_reply.started":"2022-07-14T11:23:29.229189Z","shell.execute_reply":"2022-07-14T11:23:29.238671Z"},"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'])","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:23:29.246703Z","iopub.execute_input":"2022-07-14T11:23:29.247073Z","iopub.status.idle":"2022-07-14T11:23:29.262245Z","shell.execute_reply.started":"2022-07-14T11:23:29.247041Z","shell.execute_reply":"2022-07-14T11:23:29.261034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_DIR = \"../input/feedback-prize-effectiveness/\"\n\ndef get_essay(essay_id, is_train=True):\n    parent_path = INPUT_DIR + 'train' if is_train else INPUT_DIR + 'test'\n    essay_path = os.path.join(parent_path, f\"{essay_id}.txt\")\n    essay_text = open(essay_path, 'r').read()\n    return essay_text\n\ntrain_df['essay_text']  = train_df['essay_id'].apply(lambda x: get_essay(x, is_train=True))\n\nSEP = CONFIG['tokenizer'].sep_token\n\ntrain_df['text'] = train_df['discourse_type'] + ' ' + train_df['discourse_text'] + SEP + train_df['essay_text']","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:23:29.264922Z","iopub.execute_input":"2022-07-14T11:23:29.265360Z","iopub.status.idle":"2022-07-14T11:24:06.664101Z","shell.execute_reply.started":"2022-07-14T11:23:29.265322Z","shell.execute_reply":"2022-07-14T11:24:06.662909Z"},"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.text = df['text'].values\n        self.discourse_text = df['discourse_text'].values\n        self.targets = df['discourse_effectiveness_label'].values\n        self.essay_id = df['essay_id'].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        #essay_path = os.path.join(self.data_path, f\"{self.essay_id[index]}.txt\")\n        #essay = open(essay_path, 'r').read()\n        #text = discourse_text + \" \" + self.tokenizer.sep_token + \" \" + essay\n        text = self.text[index]\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        }","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.665743Z","iopub.execute_input":"2022-07-14T11:24:06.666147Z","iopub.status.idle":"2022-07-14T11:24:06.677103Z","shell.execute_reply.started":"2022-07-14T11:24:06.666105Z","shell.execute_reply":"2022-07-14T11:24:06.675462Z"},"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):\n        super(FeedBackModel, self).__init__()\n        self.drop = nn.Dropout(p=0.05)\n\n        self.model = AutoModel.from_pretrained(model_name)\n        self.config = AutoConfig.from_pretrained(model_name)\n        self.mpool = MeanPooling()\n        \n        self.fc = nn.Sequential(\n            nn.Linear(self.config.hidden_size, CONFIG['num_classes']),\n        )\n#         self.fc2 = nn.Sequential(\n#             nn.Linear(1024, 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-07-14T11:24:06.679496Z","iopub.execute_input":"2022-07-14T11:24:06.680473Z","iopub.status.idle":"2022-07-14T11:24:06.694579Z","shell.execute_reply.started":"2022-07-14T11:24:06.680423Z","shell.execute_reply":"2022-07-14T11:24:06.693429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\ndef train_one_epoch(model, optimizer, dataloader, device):\n    model.train()\n\n    total = 0\n    running_loss = 0.0\n    correct = 0\n    start = end = time.time()\n\n    bar = tqdm(dataloader, total=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.backward()\n\n        optimizer.step()\n        optimizer.zero_grad()\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(Loss=epoch_loss, Accuracy=acc*100)\n        if step % 100 == 0:\n            print('step: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss} '\n                  .format(step, len(dataloader), \n                          remain=timeSince(start, float(step+1)/len(dataloader)),\n                          loss=epoch_loss))\n\n    return epoch_loss, acc\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\n","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.698122Z","iopub.execute_input":"2022-07-14T11:24:06.698454Z","iopub.status.idle":"2022-07-14T11:24:06.717559Z","shell.execute_reply.started":"2022-07-14T11:24:06.698421Z","shell.execute_reply":"2022-07-14T11:24:06.716616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def start_training(model, optimizer, device, num_epochs):\n    start = time.time()\n    best_epoch_loss = np.inf\n    history = {\"Train Loss\": [], \"Valid Loss\": [], \"Train Acc\": [], \"Valid Acc\": []}\n\n    for epoch in range(1, num_epochs + 1):\n        print(\"Epoch: \", epoch)\n        train_epoch_loss, train_epoch_acc = train_one_epoch(\n            model, optimizer, 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\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":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.719361Z","iopub.execute_input":"2022-07-14T11:24:06.719844Z","iopub.status.idle":"2022-07-14T11:24:06.734433Z","shell.execute_reply.started":"2022-07-14T11:24:06.719806Z","shell.execute_reply":"2022-07-14T11:24:06.733371Z"},"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)\ndf_valid = train_df.drop(df_train.index) ","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.736128Z","iopub.execute_input":"2022-07-14T11:24:06.736570Z","iopub.status.idle":"2022-07-14T11:24:06.778992Z","shell.execute_reply.started":"2022-07-14T11:24:06.736531Z","shell.execute_reply":"2022-07-14T11:24:06.777946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape, df_valid.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.781599Z","iopub.execute_input":"2022-07-14T11:24:06.782440Z","iopub.status.idle":"2022-07-14T11:24:06.791453Z","shell.execute_reply.started":"2022-07-14T11:24:06.782360Z","shell.execute_reply":"2022-07-14T11:24:06.789942Z"},"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-07-14T11:24:06.794164Z","iopub.execute_input":"2022-07-14T11:24:06.794939Z","iopub.status.idle":"2022-07-14T11:24:06.806742Z","shell.execute_reply.started":"2022-07-14T11:24:06.794883Z","shell.execute_reply":"2022-07-14T11:24:06.805151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# BeGin Training","metadata":{}},{"cell_type":"code","source":"model = FeedBackModel(CONFIG['model_name'])\nmodel.to(CONFIG['device'])\n\noptimizer = AdamW(model.parameters(), lr=CONFIG['learning_rate'])\ncriterion = nn.CrossEntropyLoss()\nmodel, history = start_training(\n    model, optimizer, device=CONFIG['device'], num_epochs=CONFIG['epochs'])\n","metadata":{"execution":{"iopub.status.busy":"2022-07-14T11:24:06.808636Z","iopub.execute_input":"2022-07-14T11:24:06.810161Z"},"trusted":true},"execution_count":null,"outputs":[]}]}