{"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":"# About this notebook\n- [Luke](https://arxiv.org/pdf/2010.01057v1.pdf)-base starter notebook\n- [Inference notebook](https://www.kaggle.com/yasufuminakama/jigsaw4-luke-base-starter-sub)\n- Approach References\n    - https://www.kaggle.com/c/jigsaw-toxic-severity-rating/discussion/286471\n    - https://www.kaggle.com/debarshichanda/pytorch-w-b-jigsaw-starter\n    - https://www.kaggle.com/debarshichanda/0-816-jigsaw-inference\n    - Thanks for sharing @debarshichanda","metadata":{}},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:51:25.505630Z","iopub.execute_input":"2022-02-05T11:51:25.506322Z","iopub.status.idle":"2022-02-05T11:51:25.538410Z","shell.execute_reply.started":"2022-02-05T11:51:25.506221Z","shell.execute_reply":"2022-02-05T11:51:25.537070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    competition='Jigsaw4'\n    _wandb_kernel='nakama'\n    debug=False\n    apex=True\n    print_freq=50\n    num_workers=4\n    model=\"studio-ousia/luke-base\"\n    scheduler='cosine' # ['linear', 'cosine']\n    batch_scheduler=True\n    num_cycles=0.5\n    num_warmup_steps=0\n    epochs=15#3\n    encoder_lr=1e-5\n    decoder_lr=1e-5\n    min_lr=1e-6\n    eps=1e-6\n    betas=(0.9, 0.999)\n    batch_size=64\n    fc_dropout=0.\n    text=\"text\"\n    target=\"target\"\n    target_size=1\n    head=32\n    tail=32\n    max_len=head+tail\n    weight_decay=0.01\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    margin=0.5\n    seed=42\n    n_fold=5\n    trn_fold=[0, 1, 2, 3, 4]\n    train=True","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:51:25.541055Z","iopub.execute_input":"2022-02-05T11:51:25.541665Z","iopub.status.idle":"2022-02-05T11:51:25.553772Z","shell.execute_reply.started":"2022-02-05T11:51:25.541622Z","shell.execute_reply":"2022-02-05T11:51:25.552771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n# ====================================================\n# wandb\n# ====================================================\nimport wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key=secret_value_0)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')\n\n    \ndef class2dict(f):\n    return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\nrun = wandb.init(project='Jigsaw4-Public', \n                 name=CFG.model,\n                 config=class2dict(CFG),\n                 group=CFG.model,\n                 job_type=\"train\",\n                 anonymous=anony)\n'''","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:51:25.557613Z","iopub.execute_input":"2022-02-05T11:51:25.557880Z","iopub.status.idle":"2022-02-05T11:51:25.570877Z","shell.execute_reply.started":"2022-02-05T11:51:25.557848Z","shell.execute_reply":"2022-02-05T11:51:25.569572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport os\nimport gc\nimport re\nimport sys\nimport json\nimport time\nimport math\nimport string\nimport pickle\nimport random\nimport joblib\nimport itertools\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\npd.set_option('display.max_rows', 500)\npd.set_option('display.max_columns', 500)\npd.set_option('display.width', 1000)\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import Parameter\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.utils.data import DataLoader, Dataset\n\nos.system('pip uninstall -q transformers -y')\nos.system('pip uninstall -q tokenizers -y')\nos.system('pip uninstall -q huggingface_hub -y')\n\nos.system('mkdir -p /tmp/pip/cache-tokenizers/')\nos.system('cp ../input/tokenizers-0103/tokenizers-0.10.3-cp37-cp37m-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_12_x86_64.manylinux2010_x86_64.whl /tmp/pip/cache-tokenizers/')\nos.system('pip install -q --no-index --find-links /tmp/pip/cache-tokenizers/ tokenizers')\n\nos.system('mkdir -p /tmp/pip/cache-huggingface-hub/')\nos.system('cp ../input/huggingface-hub-008/huggingface_hub-0.0.8-py3-none-any.whl /tmp/pip/cache-huggingface-hub/')\nos.system('pip install -q --no-index --find-links /tmp/pip/cache-huggingface-hub/ huggingface_hub')\n\nos.system('mkdir -p /tmp/pip/cache-transformers/')\nos.system('cp ../input/transformers-470/transformers-4.7.0-py3-none-any.whl /tmp/pip/cache-transformers/')\nos.system('pip install -q --no-index --find-links /tmp/pip/cache-transformers/ transformers')\n\nimport tokenizers\nimport transformers\nprint(f\"tokenizers.__version__: {tokenizers.__version__}\")\nprint(f\"transformers.__version__: {transformers.__version__}\")\nfrom transformers import LukeTokenizer, LukeModel, LukeConfig\nfrom transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:51:25.574165Z","iopub.execute_input":"2022-02-05T11:51:25.574558Z","iopub.status.idle":"2022-02-05T11:52:07.789547Z","shell.execute_reply.started":"2022-02-05T11:51:25.574513Z","shell.execute_reply":"2022-02-05T11:52:07.788388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(df):\n    score = len(df[df['less_toxic_pred'] < df['more_toxic_pred']]) / len(df)\n    return score\n\n\ndef get_logger(filename=OUTPUT_DIR+'train'):\n    from logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=f\"{filename}.log\")\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = get_logger()\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    \nseed_everything(seed=2345)","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:07.792311Z","iopub.execute_input":"2022-02-05T11:52:07.792887Z","iopub.status.idle":"2022-02-05T11:52:07.809338Z","shell.execute_reply.started":"2022-02-05T11:52:07.792841Z","shell.execute_reply":"2022-02-05T11:52:07.808224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Data Loading\n# ====================================================\ntrain = pd.read_csv('../input/jigsaw-toxic-severity-rating/validation_data.csv')\nif CFG.debug:\n    train = train.sample(n=100, random_state=CFG.seed).reset_index(drop=True)\ntest = pd.read_csv('../input/jigsaw-toxic-severity-rating/comments_to_score.csv')\nsubmission = pd.read_csv('../input/jigsaw-toxic-severity-rating/sample_submission.csv')\nprint(train.shape)\nprint(test.shape, submission.shape)\ndisplay(train.head())\ndisplay(test.head())\ndisplay(submission.head())","metadata":{"execution":{"iopub.status.busy":"2022-02-05T12:04:16.245041Z","iopub.execute_input":"2022-02-05T12:04:16.246294Z","iopub.status.idle":"2022-02-05T12:04:16.562601Z","shell.execute_reply.started":"2022-02-05T12:04:16.246248Z","shell.execute_reply":"2022-02-05T12:04:16.561287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import GroupKFold\n\nn_splits=5\nnrows = None\n\ndf = pd.read_csv(\"../input/jigsaw-toxic-severity-rating/validation_data.csv\", nrows=nrows)\ntexts = set(df.less_toxic.to_list() + df.more_toxic.to_list())\ntext2id = {t:id for id,t in enumerate(texts)}\ndf['less_id'] = df['less_toxic'].map(text2id)\ndf['more_id'] = df['more_toxic'].map(text2id)\n\n# Set array to store pair information\nlen_ids = len(text2id)\nidarr = np.zeros((len_ids,len_ids), dtype=bool)\n\nfor lid, mid in df[['less_id', 'more_id']].values:\n    min_id = min(lid, mid)\n    max_id = max(lid, mid)\n    idarr[max_id, min_id] = True\n\n# Recursively retrieve the text that is paired with the text whose id is i,\n# and store it's id in this_list.\n# then set idarr[i, j] to False\ndef add_ids(i, this_list):\n    for j in range(len_ids):\n        if idarr[i, j]:\n            idarr[i, j] = False\n            this_list.append(j)\n            this_list = add_ids(j,this_list)\n            #print(j,i)\n    for j in range(i+1,len_ids):\n        if idarr[j, i]:\n            idarr[j, i] = False\n            this_list.append(j)\n            this_list = add_ids(j,this_list)\n            #print(j,i)\n    return this_list\n\ngroup_list = []\nfor i in tqdm(range(len_ids)):\n    for j in range(i+1,len_ids):\n        if idarr[j, i]:\n            this_list = add_ids(i,[i])\n            #print(this_list)\n            group_list.append(this_list)\n\nid2groupid = {}\nfor gid,ids in enumerate(group_list):\n    for id in ids:\n        id2groupid[id] = gid\n\ndf['less_gid'] = df['less_id'].map(id2groupid)\ndf['more_gid'] = df['more_id'].map(id2groupid)\n\nprint('unique text counts:', len_ids)\nprint('grouped text counts:', len(group_list))\n\n# now we can use GroupKFold with group id\ngroup_kfold = GroupKFold(n_splits=n_splits)\n\n# Since df.less_gid and df.more_gid are the same, let's use df.less_gid here.\nfor fold, (trn, val) in enumerate(group_kfold.split(df, df, df.less_gid)): \n    df.loc[val , \"fold\"] = fold\n\ndf[\"fold\"] = df[\"fold\"].astype(int)\n\ndisplay(df.groupby('fold').size())\n\ntrain = df[list(train.columns)+['fold']].copy()\ntrain","metadata":{"execution":{"iopub.status.busy":"2022-02-05T12:04:23.146923Z","iopub.execute_input":"2022-02-05T12:04:23.147210Z","iopub.status.idle":"2022-02-05T12:06:10.396541Z","shell.execute_reply.started":"2022-02-05T12:04:23.147179Z","shell.execute_reply":"2022-02-05T12:06:10.395531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# tokenizer","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# tokenizer\n# ====================================================\ntokenizer = LukeTokenizer.from_pretrained(CFG.model, lowercase=True)\ntokenizer.save_pretrained(OUTPUT_DIR+'tokenizer/')\nCFG.tokenizer = tokenizer","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.852021Z","iopub.status.idle":"2022-02-05T11:52:08.852934Z","shell.execute_reply.started":"2022-02-05T11:52:08.852629Z","shell.execute_reply":"2022-02-05T11:52:08.852659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\ndef prepare_input(text, cfg):\n    if cfg.tail == 0:\n        inputs = cfg.tokenizer.encode_plus(text, \n                                           return_tensors=None, \n                                           add_special_tokens=True, \n                                           max_length=cfg.max_len,\n                                           pad_to_max_length=True,\n                                           truncation=True)\n        for k, v in inputs.items():\n            inputs[k] = torch.tensor(v, dtype=torch.long)\n    else:\n        inputs = cfg.tokenizer.encode_plus(text,\n                                           return_tensors=None, \n                                           add_special_tokens=True, \n                                           truncation=True)\n        for k, v in inputs.items():\n            v_length = len(v)\n            if v_length > cfg.max_len:\n                v = np.hstack([v[:cfg.head], v[-cfg.tail:]])\n            if k == 'input_ids':\n                new_v = np.ones(cfg.max_len) * cfg.tokenizer.pad_token_id\n            else:\n                new_v = np.zeros(cfg.max_len)\n            new_v[:v_length] = v \n            inputs[k] = torch.tensor(new_v, dtype=torch.long)\n    return inputs\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, cfg, df):\n        self.cfg = cfg\n        self.less_toxic = df['less_toxic'].fillna(\"none\").values\n        self.more_toxic = df['more_toxic'].fillna(\"none\").values\n\n    def __len__(self):\n        return len(self.less_toxic)\n\n    def __getitem__(self, item):\n        less_toxic_inputs = prepare_input(str(self.less_toxic[item]), self.cfg)\n        more_toxic_inputs = prepare_input(str(self.more_toxic[item]), self.cfg)\n        label = torch.tensor(1, dtype=torch.float)\n        return less_toxic_inputs, more_toxic_inputs, label\n\n\nclass TestDataset(Dataset):\n    def __init__(self, cfg, df):\n        self.cfg = cfg\n        self.text = df[cfg.text].fillna(\"none\").values\n\n    def __len__(self):\n        return len(self.text)\n\n    def __getitem__(self, item):\n        text = str(self.text[item])\n        inputs = prepare_input(text, self.cfg)\n        return inputs","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.854828Z","iopub.status.idle":"2022-02-05T11:52:08.855414Z","shell.execute_reply.started":"2022-02-05T11:52:08.855116Z","shell.execute_reply":"2022-02-05T11:52:08.855147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Model\n# ====================================================\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, config_path=None, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        if config_path is None:\n            self.config = LukeConfig.from_pretrained(cfg.model, output_hidden_states=True)\n        else:\n            self.config = torch.load(config_path)\n        if pretrained:\n            self.model = LukeModel.from_pretrained(cfg.model, config=self.config)\n        else:\n            self.model = LukeModel(self.config)\n        self.fc_dropout = nn.Dropout(cfg.fc_dropout)\n        self.fc = nn.Linear(self.config.hidden_size, cfg.target_size)\n        \n    def feature(self, inputs):\n        outputs = self.model(**inputs)\n        last_hidden_states = outputs[0]\n        feature = torch.mean(last_hidden_states, 1)\n        return feature\n\n    def forward(self, inputs):\n        feature = self.feature(inputs)\n        output = self.fc(self.fc_dropout(feature))\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.857152Z","iopub.status.idle":"2022-02-05T11:52:08.858072Z","shell.execute_reply.started":"2022-02-05T11:52:08.857762Z","shell.execute_reply":"2022-02-05T11:52:08.857793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helpler functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\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\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.apex)\n    losses = AverageMeter()\n    start = end = time.time()\n    global_step = 0\n    for step, (less_toxic_inputs, more_toxic_inputs, labels) in enumerate(train_loader):\n        for k, v in less_toxic_inputs.items():\n            less_toxic_inputs[k] = v.to(device)\n        for k, v in more_toxic_inputs.items():\n            more_toxic_inputs[k] = v.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        with torch.cuda.amp.autocast(enabled=CFG.apex):\n            less_toxic_y_preds = model(less_toxic_inputs)\n            more_toxic_y_preds = model(more_toxic_inputs)\n            loss = criterion(more_toxic_y_preds, less_toxic_y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        scaler.scale(loss).backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            global_step += 1\n            if CFG.batch_scheduler:\n                scheduler.step()\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.8f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        #wandb.log({f\"[fold{fold}] loss\": losses.val,\n        #           f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef inference_fn(test_loader, model, device):\n    preds = []\n    model.eval()\n    model.to(device)\n    tk0 = tqdm(test_loader, total=len(test_loader))\n    for inputs in tk0:\n        for k, v in inputs.items():\n            inputs[k] = v.to(device)\n        with torch.no_grad():\n            y_preds = model(inputs)\n        preds.append(y_preds.sigmoid().to('cpu').numpy())\n    predictions = np.concatenate(preds)\n    return predictions","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.860128Z","iopub.status.idle":"2022-02-05T11:52:08.860654Z","shell.execute_reply.started":"2022-02-05T11:52:08.860372Z","shell.execute_reply":"2022-02-05T11:52:08.860402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# train loop\n# ====================================================\ndef train_loop(folds, fold):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    \n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == fold].index\n    \n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    validation = folds.loc[val_idx].reset_index(drop=True)\n    \n    valid_folds = sorted(set(validation['less_toxic'].unique()) | set(validation['more_toxic'].unique()))\n    valid_folds = pd.DataFrame({'text': valid_folds}).reset_index()\n    \n    train_dataset = TrainDataset(CFG, train_folds)\n    valid_dataset = TestDataset(CFG, valid_folds)\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size,\n                              shuffle=True,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset,\n                              batch_size=CFG.batch_size,\n                              shuffle=False,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel(CFG, config_path=None, pretrained=True)\n    torch.save(model.config, OUTPUT_DIR+'config.pth')\n    model.to(device)\n    \n    def get_optimizer_params(model, encoder_lr, decoder_lr, weight_decay=0.0):\n        param_optimizer = list(model.named_parameters())\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(optimizer_parameters, lr=CFG.encoder_lr, eps=CFG.eps, betas=CFG.betas)\n    \n    # ====================================================\n    # scheduler\n    # ====================================================\n    def get_scheduler(cfg, optimizer, num_train_steps):\n        if cfg.scheduler=='linear':\n            scheduler = get_linear_schedule_with_warmup(\n                optimizer, num_warmup_steps=cfg.num_warmup_steps, num_training_steps=num_train_steps\n            )\n        elif cfg.scheduler=='cosine':\n            scheduler = get_cosine_schedule_with_warmup(\n                optimizer, num_warmup_steps=cfg.num_warmup_steps, num_training_steps=num_train_steps, num_cycles=cfg.num_cycles\n            )\n        return scheduler\n    \n    num_train_steps = int(len(train_folds) / CFG.batch_size * CFG.epochs)\n    scheduler = get_scheduler(CFG, optimizer, num_train_steps)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.MarginRankingLoss(margin=CFG.margin)\n    \n    best_score = 0.\n\n    for epoch in range(CFG.epochs):\n\n        start_time = time.time()\n\n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        preds = inference_fn(valid_loader, model, device)\n        \n        # scoring\n        valid_folds['pred'] = preds\n        if 'less_toxic_pred' in validation.columns:\n            validation = validation.drop(columns='less_toxic_pred')\n        if 'more_toxic_pred' in validation.columns:\n            validation = validation.drop(columns='more_toxic_pred')\n        rename_cols = {CFG.text: 'less_toxic', 'pred': 'less_toxic_pred'}\n        validation = validation.merge(valid_folds[[CFG.text, 'pred']].rename(columns=rename_cols), \n                                      on='less_toxic', how='left')\n        rename_cols = {CFG.text: 'more_toxic', 'pred': 'more_toxic_pred'}\n        validation = validation.merge(valid_folds[[CFG.text, 'pred']].rename(columns=rename_cols), \n                                      on='more_toxic', how='left')\n        score = get_score(validation)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n        #wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n        #           f\"[fold{fold}] avg_train_loss\": avg_loss, \n        #           f\"[fold{fold}] score\": score})\n        \n        if score > best_score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {score:.4f} Model')\n            torch.save({'model': model.state_dict(),\n                        'preds': preds},\n                        OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\")\n\n    preds = torch.load(OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\", \n                       map_location=torch.device('cpu'))['preds']\n    valid_folds['pred'] = preds\n    if 'less_toxic_pred' in validation.columns:\n        validation = validation.drop(columns='less_toxic_pred')\n    if 'more_toxic_pred' in validation.columns:\n        validation = validation.drop(columns='more_toxic_pred')\n    rename_cols = {CFG.text: 'less_toxic', 'pred': 'less_toxic_pred'}\n    validation = validation.merge(valid_folds[[CFG.text, 'pred']].rename(columns=rename_cols), \n                                  on='less_toxic', how='left')\n    rename_cols = {CFG.text: 'more_toxic', 'pred': 'more_toxic_pred'}\n    validation = validation.merge(valid_folds[[CFG.text, 'pred']].rename(columns=rename_cols), \n                                  on='more_toxic', how='left')\n\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return validation","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.864484Z","iopub.status.idle":"2022-02-05T11:52:08.865031Z","shell.execute_reply.started":"2022-02-05T11:52:08.864738Z","shell.execute_reply":"2022-02-05T11:52:08.864766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    def get_result(oof_df):\n        score = get_score(oof_df)\n        LOGGER.info(f'Score: {score:<.4f}')\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df = train_loop(train, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        oof_df = oof_df.reset_index(drop=True)\n        # CV result\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        # save result\n        oof_df.to_csv(OUTPUT_DIR+'oof_df.csv', index=False)\n    \n    #wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2022-02-05T11:52:08.866974Z","iopub.status.idle":"2022-02-05T11:52:08.867533Z","shell.execute_reply.started":"2022-02-05T11:52:08.867191Z","shell.execute_reply":"2022-02-05T11:52:08.867219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}