{"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":"cfg = {\n    \"num_proc\": 2,\n    # data\n    \"k_folds\": 5,\n    \"max_length\": 2048,\n    \"padding\": False,\n    \"stride\": 0,\n    \"data_dir\": \"D:/feedback\",\n    \"load_from_disk\": '../input/token-classification-approach-fpe-4ecac6/output', # if you already tokenized, you can load it through this\n    \"pad_multiple\": 8,\n    # model\n    #\"model_name_or_path\": 'allenai/longformer-large-4096',\n    'model_name_or_path':\"microsoft/deberta-v3-large\",\n    \"dropout\": 0.1,\n    # to put in TrainingArguments\n    \"trainingargs\": {\n        \"output_dir\": \"D:/feedback/tokend/debertav3largedata/attention\",\n        \"do_train\": True,\n        \"do_eval\": True,\n        \"per_device_train_batch_size\": 1,\n        \"per_device_eval_batch_size\": 4,\n        \"learning_rate\": 5e-7,\n        \"weight_decay\": 0.01,\n        \"num_train_epochs\": 2,\n        \"gradient_accumulation_steps\":4,\n        #\"gradient_checkpointing\":True,\n        #\"warmup_ratio\": 0.1,\n        \"optim\": 'adamw_torch',\n        \"logging_steps\": 500,\n        \"save_strategy\": \"epoch\",\n        \"evaluation_strategy\": \"epoch\",\n        #'load_best_model_at_end':True,\n        \"report_to\": \"none\",\n        \"group_by_length\": True,\n        \"save_total_limit\": 1,\n        \"metric_for_best_model\": \"loss\",\n        \"greater_is_better\": False,\n        \"seed\": 18,\n       # \"disable_tqdm\":True,\n        #'bf16 ':True\n        'fp16':True\n        # you should probably set \"fp16\" to True, but it doesn't really matter on Kaggle\n    }\n}","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:25.711421Z","iopub.execute_input":"2022-08-08T06:24:25.711834Z","iopub.status.idle":"2022-08-08T06:24:25.721862Z","shell.execute_reply.started":"2022-08-08T06:24:25.711799Z","shell.execute_reply":"2022-08-08T06:24:25.720771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\nimport pickle\nimport codecs\nimport warnings\nimport logging\nfrom functools import partial\nfrom pathlib import Path\nfrom itertools import chain\nfrom text_unidecode import unidecode\nfrom typing import Any, Optional, Tuple\nimport gc\nimport pandas as pd\nfrom sklearn.model_selection import KFold,GroupKFold\nfrom transformers import AutoTokenizer, set_seed\nfrom transformers import DebertaV2Tokenizer, DebertaV2ForTokenClassification\nfrom datasets import Dataset, load_from_disk\nfrom torch.utils.data import DataLoader, Dataset\nimport os\nimport gc\n#from log import _Logger\nimport random\nimport warnings\nfrom functools import reduce\nwarnings.filterwarnings(\"ignore\")\nimport numpy as np\nfrom numpy import ndarray\nimport scipy as sp\nimport torch\nfrom torch import inference_mode\nfrom torch import nn\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import _LRScheduler\nfrom IPython.display import display\nfrom transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup,AutoModelForTokenClassification,DataCollatorForTokenClassification\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:25.724015Z","iopub.execute_input":"2022-08-08T06:24:25.724604Z","iopub.status.idle":"2022-08-08T06:24:25.737760Z","shell.execute_reply.started":"2022-08-08T06:24:25.724567Z","shell.execute_reply":"2022-08-08T06:24:25.736725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = Path(cfg[\"data_dir\"])\n\nif cfg[\"load_from_disk\"]:\n    if not cfg[\"load_from_disk\"].endswith(\".dataset\"):\n        cfg[\"load_from_disk\"] += \".dataset\"\n    ds = load_from_disk(cfg['load_from_disk'])\n    \n    pkl_file = f\"{cfg['load_from_disk'][:-len('.dataset')]}_pkl\"\n    with open(pkl_file, \"rb\") as fp:\n        grouped = pickle.load(fp)\n    \n    print(\"Loading from saved files\")\ndisc_types = [\n    \"Claim\",\n    \"Concluding Statement\",\n    \"Counterclaim\",\n    \"Evidence\",\n    \"Lead\",\n    \"Position\",\n    \"Rebuttal\",\n]\ncls_tokens_map = {label: f\"[CLS_{label.upper()}]\" for label in disc_types}\nend_tokens_map = {label: f\"[END_{label.upper()}]\" for label in disc_types}\n\nlabel2id = {\n    \"Adequate\": 0,\n    \"Effective\": 1,\n    \"Ineffective\": 2,\n}\n\ntokenizer = AutoTokenizer.from_pretrained(cfg[\"model_name_or_path\"])\ntokenizer.add_special_tokens(\n    {\"additional_special_tokens\": list(cls_tokens_map.values())+list(end_tokens_map.values())}\n)\ncls_id_map = {\n    label: tokenizer.encode(tkn)[1]\n    for label, tkn in cls_tokens_map.items()\n}","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:25.741362Z","iopub.execute_input":"2022-08-08T06:24:25.743145Z","iopub.status.idle":"2022-08-08T06:24:27.789753Z","shell.execute_reply.started":"2022-08-08T06:24:25.743100Z","shell.execute_reply":"2022-08-08T06:24:27.788785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# basic kfold \ndef get_folds(df, k_folds=4):\n\n    kf = GroupKFold(n_splits=k_folds)\n    return [\n        val_idx\n        for _, val_idx in kf.split(df['input_ids'],ds['labels'],groups=ds['essay_id'])\n    ]\n\nfold_idxs = get_folds(ds, cfg[\"k_folds\"])","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:27.791174Z","iopub.execute_input":"2022-08-08T06:24:27.791640Z","iopub.status.idle":"2022-08-08T06:24:29.196207Z","shell.execute_reply.started":"2022-08-08T06:24:27.791594Z","shell.execute_reply":"2022-08-08T06:24:29.195192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keep_cols = {\"input_ids\", \"attention_mask\", \"labels\"}\ndata = ds.remove_columns([c for c in ds.column_names if c not in keep_cols]).to_pandas()","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.199331Z","iopub.execute_input":"2022-08-08T06:24:29.199755Z","iopub.status.idle":"2022-08-08T06:24:29.230955Z","shell.execute_reply.started":"2022-08-08T06:24:29.199701Z","shell.execute_reply":"2022-08-08T06:24:29.229873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.232213Z","iopub.execute_input":"2022-08-08T06:24:29.233225Z","iopub.status.idle":"2022-08-08T06:24:29.438241Z","shell.execute_reply.started":"2022-08-08T06:24:29.233186Z","shell.execute_reply":"2022-08-08T06:24:29.436849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader, Dataset\nimport numpy as np\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport torch\nfrom transformers import AutoModelWithLMHead, AutoTokenizer, AutoModel\nimport pickle\nclass TrainDataset(Dataset):\n    def __init__(self, df):\n        self.ids = df['input_ids'].values\n        self.att = df['attention_mask'].values\n        self.target = df['labels'].values\n\n    def __len__(self):\n        return len(self.target)\n\n    def __getitem__(self, item):\n        samples = {\n            'input_ids': self.ids[item],\n            'attention_mask': self.att[item],\n            'label': self.target[item],\n        }\n        return samples\n","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.440822Z","iopub.execute_input":"2022-08-08T06:24:29.441424Z","iopub.status.idle":"2022-08-08T06:24:29.451787Z","shell.execute_reply.started":"2022-08-08T06:24:29.441346Z","shell.execute_reply":"2022-08-08T06:24:29.450137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import Tensor\nfrom torch.nn import Module\nfrom transformers import AutoModel, AutoConfig\nfrom torch import Tensor\nfrom torch.nn import Module\nfrom torch.optim import Optimizer\nfrom torch.nn.modules.loss import _Loss\n\nclass AWP:\n    def __init__(\n        self,\n        model: Module,\n        criterion: _Loss,\n        optimizer: Optimizer,\n        apex: bool,\n        adv_param: str=\"weight\",\n        adv_lr: float=1.0,\n        adv_eps: float=0.01\n    ) -> None:\n        self.model = model\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.adv_param = adv_param\n        self.adv_lr = adv_lr\n        self.adv_eps = adv_eps\n        self.apex = apex\n        self.backup = {}\n        self.backup_eps = {}\n\n    def attack_backward(self, input_ids,attention_mask,labels) -> Tensor:\n        with torch.cuda.amp.autocast(enabled=self.apex):\n            self._save()\n            self._attack_step() # モデルを近傍の悪い方へ改変\n            y_preds = self.model(input_ids=sample['input_ids'],attention_mask=sample['attention_mask'],labels=labels)\n            self.optimizer.zero_grad()\n        return y_preds\n\n    def _attack_step(self) -> None:\n        e = 1e-6\n        for name, param in self.model.named_parameters():\n            if param.requires_grad and param.grad is not None and self.adv_param in name:\n                norm1 = torch.norm(param.grad)\n                norm2 = torch.norm(param.data.detach())\n                if norm1 != 0 and not torch.isnan(norm1):\n                    # 直前に損失関数に通してパラメータの勾配を取得できるようにしておく必要あり\n                    r_at = self.adv_lr * param.grad / (norm1 + e) * (norm2 + e)\n                    param.data.add_(r_at)\n                    param.data = torch.min(\n                        torch.max(\n                            param.data, self.backup_eps[name][0]), self.backup_eps[name][1]\n                    )\n\n    def _save(self) -> None:\n        for name, param in self.model.named_parameters():\n            if param.requires_grad and param.grad is not None and self.adv_param in name:\n                if name not in self.backup:\n                    self.backup[name] = param.data.clone()\n                    grad_eps = self.adv_eps * param.abs().detach()\n                    self.backup_eps[name] = (\n                        self.backup[name] - grad_eps,\n                        self.backup[name] + grad_eps,\n                    )\n\n    def _restore(self) -> None:\n        for name, param in self.model.named_parameters():\n            if name in self.backup:\n                param.data = self.backup[name]\n        self.backup = {}\n        self.backup_eps = {}\n\nclass FGM():\n    def __init__(self, model):\n        self.model = model\n        self.backup = {}\n\n    def attack(self, epsilon=1., emb_name='word_embeddings'):\n        for name, param in self.model.named_parameters():\n            if param.requires_grad and emb_name in name:\n                self.backup[name] = param.data.clone()\n                norm = torch.norm(param.grad)\n                if norm != 0:\n                    r_at = epsilon * param.grad / norm\n                    param.data.add_(r_at)\n\n    def restore(self, emb_name='word_embeddings'):\n        for name, param in self.model.named_parameters():\n            if param.requires_grad and emb_name in name:\n                assert name in self.backup\n                param.data = self.backup[name]\n            self.backup = {}","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.453620Z","iopub.execute_input":"2022-08-08T06:24:29.454449Z","iopub.status.idle":"2022-08-08T06:24:29.474287Z","shell.execute_reply.started":"2022-08-08T06:24:29.454411Z","shell.execute_reply":"2022-08-08T06:24:29.473272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"param = {\n    'apex': True,\n    'awp_eps': 1e-2,\n    'awp_lr': 1e-4,\n    'batch_size': 1, # 2\n    'betas': (0.9, 0.999),\n    'ckpt_name': 'deberta_v3_large',\n    'debug': True, # False\n    'decoder_lr': 1e-5,\n    'encoder_lr': 1e-5,\n    'eps': 1e-6,\n    'max_grad_norm': 1000,\n    'max_len': 2048, # 512\n    'min_lr': 5e-7,\n    'model_name': 'microsoft/deberta-v3-large',\n    'n_cycles': 1,\n    'n_epochs': 4, # 12\n    'n_eval_steps': 100,\n    'n_folds': 5, # 4\n    'n_gradient_accumulation_steps': 4,\n    'n_warmup_steps': 0,\n    'n_workers': 8,\n    'nth_awp_start_epoch': 1, # 4\n    'output_dir': './output/',\n    'print_freq': 100,\n    'scheduler_name': 'cosine',\n    'seed': 42,\n    'tar_token': '[TAR]',\n    'weight_decay': 0.01,\n    \"output_dir\": \"D:/feedback/tokend/debertav3largedata/tvm\",\n}\n\nclass Config:\n    def __init__(self, d: dict) -> None:\n        for k,v in d.items():\n            setattr(self, k, v)\n\ncfg = Config(d=param)\nimport gc\nfrom tqdm.notebook import tqdm\ndef valid_fn( dl: DataLoader, model: Module, criterion: _Loss) -> \"tuple[float, list[list[float]]]\":\n            model.eval()\n            preds = []\n            truths=[]\n            tot_loss = 0\n            for step, sample in enumerate(tqdm(dl)):\n                for k, v in sample.items():\n                    sample[k] = v.cuda()\n                target = sample['label']\n                with torch.cuda.amp.autocast(enabled=cfg.apex):\n                    loss = model(input_ids=sample['input_ids'],attention_mask=sample['attention_mask'],labels=target)\n                tot_loss += loss.item()\n                if cfg.n_gradient_accumulation_steps > 1:\n                    loss = loss / cfg.n_gradient_accumulation_steps\n                \n                if step % cfg.print_freq == 0 or step == (len(dl) - 1):\n                    print('EVAL: [{0}/{1}] '\n                        'Loss: {loss:.4f}({avg_loss:.4f}) '\n                        .format(step, len(dl),\n                                loss=loss.item()*cfg.n_gradient_accumulation_steps,\n                               avg_loss=tot_loss/(step+1)))\n                #preds.append(y_pred)\n            #preds = np.array(preds)\n            return tot_loss/(step+1)\ndef train(\n                            fold: int,\n                            train_loader: DataLoader,\n                            model: Module,\n                            optimizer: Optimizer,\n                            epoch: int,\n                            scheduler: _LRScheduler) -> \"tuple[float, float]\":\n\n            criterion =  nn.CrossEntropyLoss()\n            model.train()\n            #awp = AWP(\n            #    model, \n            #    criterion, \n            #    optimizer,\n            #    True,\n            #    adv_lr=1e-5, \n            #    adv_eps=1e-6\n            #)\n            #fgm = FGM(model)\n            scaler = torch.cuda.amp.GradScaler(enabled=cfg.apex)\n            global_step = 0\n            tot_loss = 0\n            for step, sample in enumerate(tqdm(train_loader)):\n                for k, v in sample.items():\n                    sample[k] = v.cuda()\n                target = sample['label']\n                with torch.cuda.amp.autocast(enabled=cfg.apex):\n                    loss = model(input_ids=sample['input_ids'],attention_mask=sample['attention_mask'],labels=target)\n                #loss = criterion(y_preds,target.view(-1))\n                tot_loss += loss.item()\n                if cfg.n_gradient_accumulation_steps > 1:\n                    loss = loss / cfg.n_gradient_accumulation_steps\n                scaler.scale(loss).backward()\n\n                # adversarial training\n                #fgm.attack() \n                #loss_adv = model(input_ids=sample['input_ids'],attention_mask=sample['attention_mask'],labels=target)\n                #scaler.scale(loss_adv).backward()\n                #fgm.restore()  \n                \n                grad_norm = torch.nn.utils.clip_grad_norm_(\n                    model.parameters(), \n                    cfg.max_grad_norm)\n                \n                #if 1 <= epoch:\n                #    loss = awp.attack_backward(input_ids=sample['input_ids'],attention_mask=sample['attention_mask'],labels=target)\n                #    scaler.scale(loss).backward()\n                #    awp._restore()\n               \n                \n                if (step + 1) % cfg.n_gradient_accumulation_steps == 0:\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad()\n                    #model.zero_grad()\n                    global_step += 1\n                    scheduler.step()\n                if step % cfg.print_freq == 0 or step == (len(train_loader) - 1):\n                    print('Epoch: [{0}][{1}/{2}] '\n                        'Loss: {loss:.4f}({avg_loss:.4f}) '\n                        'Grad: {grad_norm:.4f}  '\n                        'LR: {lr:.8f}  '\n                        .format(epoch + 1, step, len(train_loader),\n                                loss=loss.item()*cfg.n_gradient_accumulation_steps,\n                                avg_loss=tot_loss/(step+1),\n                                grad_norm=grad_norm,\n                                lr=scheduler.get_lr()[0]))\n                gc.collect()\n            return tot_loss/(step+1)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.477350Z","iopub.execute_input":"2022-08-08T06:24:29.477701Z","iopub.status.idle":"2022-08-08T06:24:29.499551Z","shell.execute_reply.started":"2022-08-08T06:24:29.477675Z","shell.execute_reply":"2022-08-08T06:24:29.498594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import Tensor\nfrom torch.nn import Module\nimport torch.nn as nn\nimport torch\nimport torch.nn.functional as F\nfrom transformers import AutoModel, AutoConfig\nfrom transformers.modeling_outputs import TokenClassifierOutput\nfrom typing import Optional, Tuple, Union\nclass CustomModel(Module):\n    def __init__(self) -> None:\n        super().__init__()\n        model_config = AutoConfig.from_pretrained(\n            cfg.model_name,\n        )\n        model_config.update(\n            {\n                \"num_labels\": 3,\n                \"cls_tokens\": list(cls_id_map.values()),\n                \"label2id\": label2id,\n                \"id2label\": {v:k for k, v in label2id.items()},\n            }\n        )\n        self.model_config=model_config\n        self.model = AutoModel.from_pretrained(cfg.model_name, config=model_config)\n        self.model.resize_token_embeddings(len(tokenizer)) \n        self.dropout = nn.Dropout(model_config.hidden_dropout_prob)\n        self.classifier = nn.Linear(model_config.hidden_size, model_config.num_labels)\n        \n        self._init_weights(self.classifier)\n        #self.model.gradient_checkpointing_enable()\n    \n    def _init_weights(self, module: Module) -> None:\n        if isinstance(module, nn.Linear):\n            module.weight.data.normal_(\n                mean=0.0, std=self.model_config.initializer_range)\n            if module.bias is not None:\n                module.bias.data.zero_()\n        elif isinstance(module, nn.Embedding):\n            module.weight.data.normal_(\n                mean=0.0, std=self.model_config.initializer_range)\n            if module.padding_idx is not None:\n                module.weight.data[module.padding_idx].zero_()\n        elif isinstance(module, nn.LayerNorm):\n            module.bias.data.zero_()\n            module.weight.data.fill_(1.0)\n\n\n\n    def forward(\n        self,\n        input_ids: Optional[torch.Tensor] = None,\n        attention_mask: Optional[torch.Tensor] = None,\n        token_type_ids: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.Tensor] = None,\n        inputs_embeds: Optional[torch.Tensor] = None,\n        labels: Optional[torch.Tensor] = None,\n        output_attentions: Optional[bool] = None,\n        output_hidden_states: Optional[bool] = None,\n        return_dict: Optional[bool] = None) -> Union[Tuple, TokenClassifierOutput]:\n        \n        r\"\"\"\n        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.\n        \"\"\"\n        #return_dict = return_dict if return_dict is not None else self.config.use_return_dict\n\n        outputs = self.model(\n            input_ids,\n            attention_mask=attention_mask,\n            token_type_ids=token_type_ids,\n            position_ids=position_ids,\n            inputs_embeds=inputs_embeds,\n            output_attentions=output_attentions,\n            output_hidden_states=output_hidden_states,\n            return_dict=return_dict,\n        )\n        \n        \n        sequence_output = outputs[0]\n        sequence_output = self.dropout(sequence_output)\n        logits = self.classifier(sequence_output)\n        \n        loss = None\n        if labels is not None:\n            loss_fct = nn.CrossEntropyLoss()\n            loss = loss_fct(logits.view(-1,3), labels.view(-1))\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.501060Z","iopub.execute_input":"2022-08-08T06:24:29.501624Z","iopub.status.idle":"2022-08-08T06:24:29.519714Z","shell.execute_reply.started":"2022-08-08T06:24:29.501588Z","shell.execute_reply":"2022-08-08T06:24:29.518674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"collator = DataCollatorForTokenClassification(\n    tokenizer=tokenizer, pad_to_multiple_of=8, padding=True\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.523319Z","iopub.execute_input":"2022-08-08T06:24:29.523803Z","iopub.status.idle":"2022-08-08T06:24:29.533828Z","shell.execute_reply.started":"2022-08-08T06:24:29.523605Z","shell.execute_reply":"2022-08-08T06:24:29.532988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(5):\n        train_idxs =  list(chain(*[i for f, i in enumerate(fold_idxs) if f != fold]))\n        dft=data.iloc[train_idxs]\n        dfv=data.iloc[fold_idxs[fold]]\n        \n        train_dataset = TrainDataset(dft)\n        valid_dataset = TrainDataset(dfv)\n\n\n        train_loader = DataLoader(train_dataset,\n                                batch_size=cfg.batch_size,\n                                shuffle=True,\n                                num_workers=cfg.n_workers, \n                                pin_memory=True, \n                                drop_last=True,\n                                collate_fn=collator)\n        valid_loader = DataLoader(valid_dataset,\n                                batch_size=cfg.batch_size,\n                                shuffle=False,\n                                num_workers=cfg.n_workers, \n                                pin_memory=True, \n                                drop_last=False,\n                                collate_fn=collator)\n        model=CustomModel()\n        #if fold == 0:\n        #    model.load_state_dict(torch.load(cfg.output_dir+f'/tvmdebertav3L_{fold}.bin')())\n        model.cuda()\n\n        def get_optimizer_params(model, encoder_lr, decoder_lr, weight_decay=0.0):\n            no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n            optimizer_parameters = [\n                {'params': [p for n, p in model.model.named_parameters() if not any(nd in n for nd in no_decay)],\n                'lr': encoder_lr, 'weight_decay': weight_decay},\n                {'params': [p for n, p in model.model.named_parameters() if any(nd in n for nd in no_decay)],\n                'lr': encoder_lr, 'weight_decay': 0.0},\n                {'params': [p for n, p in model.named_parameters() if \"model\" not in n],\n                'lr': decoder_lr, 'weight_decay': 0.0}\n            ]\n            return optimizer_parameters\n\n        optimizer_parameters = get_optimizer_params(model,\n                                                    encoder_lr=cfg.encoder_lr,\n                                                    decoder_lr=cfg.decoder_lr,\n                                                   weight_decay=cfg.weight_decay)\n        optimizer = AdamW(\n            optimizer_parameters, \n            lr=cfg.encoder_lr,\n            eps=cfg.eps, \n            betas=cfg.betas)\n\n\n        def get_scheduler(scheduler_name: str, optimizer: Optimizer, num_train_steps: int, n_cycles: int) -> _LRScheduler:\n            if scheduler_name == 'linear':\n                scheduler = get_linear_schedule_with_warmup(\n                    optimizer, num_warmup_steps=cfg.n_warmup_steps, num_training_steps=num_train_steps\n                )\n            elif scheduler_name == 'cosine':\n                scheduler = get_cosine_schedule_with_warmup(\n                    optimizer, num_warmup_steps=cfg.n_warmup_steps, num_training_steps=num_train_steps, num_cycles=n_cycles\n                )\n            return scheduler\n\n        num_train_steps = int(len(train_dataset) / cfg.batch_size *cfg.n_epochs)\n        scheduler = get_scheduler(\n            cfg.scheduler_name, optimizer, num_train_steps, cfg.n_cycles)\n        criterion=nn.CrossEntropyLoss()\n        best=100.0\n        es=0\n        stop=2\n        for epoch in range(cfg.n_epochs):\n            \n            avg_loss = train(\n                fold, \n                train_loader,  \n                model,  \n                optimizer, \n                epoch, \n                scheduler)\n            avg_val_loss =valid_fn(\n                        valid_loader,\n                        model,\n                        criterion)\n\n            # scoring\n            print(f'Epoch {epoch+1} - SCORE TRAIN: {avg_loss:.6f}  SCORE VALID: {avg_val_loss:.6f}')\n            if avg_val_loss < best:\n                best = avg_val_loss\n                #torch.save(model.state_dict(),cfg.output_dir+f'/tvmdebertav3L_{fold}.bin')\n            else:\n                es+=1\n                if es == stop:\n                    break\n            gc.collect()\n        break","metadata":{"execution":{"iopub.status.busy":"2022-08-08T06:24:29.535298Z","iopub.execute_input":"2022-08-08T06:24:29.535988Z","iopub.status.idle":"2022-08-08T06:24:37.248603Z","shell.execute_reply.started":"2022-08-08T06:24:29.535950Z","shell.execute_reply":"2022-08-08T06:24:37.246937Z"},"trusted":true},"execution_count":null,"outputs":[]}]}