{"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":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nINPUT_DIR = '../input/us-patent-phrase-to-phrase-matching/'\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"id":"fa3b873b","papermill":{"duration":0.039998,"end_time":"2022-03-21T11:33:53.291434","exception":false,"start_time":"2022-03-21T11:33:53.251436","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:58:49.310193Z","iopub.execute_input":"2022-06-29T19:58:49.310735Z","iopub.status.idle":"2022-06-29T19:58:49.320253Z","shell.execute_reply.started":"2022-06-29T19:58:49.310698Z","shell.execute_reply":"2022-06-29T19:58:49.31928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{"id":"1d0c4430","papermill":{"duration":0.02483,"end_time":"2022-03-21T11:33:53.341306","exception":false,"start_time":"2022-03-21T11:33:53.316476","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    wandb=True\n    competition='PPPM'\n    _wandb_kernel='nakama'\n    debug=False\n    apex=True\n    print_freq=100\n    num_workers=4\n    model=\"microsoft/deberta-v3-small\"\n    scheduler='cosine' # ['linear', 'cosine']\n    batch_scheduler=True\n    num_cycles=0.5\n    num_warmup_steps=0\n    epochs=4\n    encoder_lr=2e-5\n    decoder_lr=2e-5\n    min_lr=1e-6\n    eps=1e-6\n    betas=(0.9, 0.999)\n    batch_size=16\n    fc_dropout=0.2\n    target_size=1\n    max_len=512\n    weight_decay=0.01\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    seed=42\n    n_fold=4\n    trn_fold=[0, 1, 2, 3]\n    train=True\n    \nif CFG.debug:\n    CFG.epochs = 2\n    CFG.trn_fold = [0]","metadata":{"id":"48dd82bb","papermill":{"duration":0.03584,"end_time":"2022-03-21T11:33:53.402377","exception":false,"start_time":"2022-03-21T11:33:53.366537","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:41:39.350413Z","iopub.execute_input":"2022-06-29T19:41:39.350692Z","iopub.status.idle":"2022-06-29T19:41:39.357661Z","shell.execute_reply.started":"2022-06-29T19:41:39.350663Z","shell.execute_reply":"2022-06-29T19:41:39.356804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# wandb\n# ====================================================\nif CFG.wandb:\n    \n    import wandb\n\n    try:\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\n    except:\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\n    def class2dict(f):\n        return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\n    run = wandb.init(project='PPPM-Public', \n                     name=CFG.model,\n                     config=class2dict(CFG),\n                     group=CFG.model,\n                     job_type=\"train\",\n                     anonymous=anony)","metadata":{"id":"b88c983e","papermill":{"duration":0.03413,"end_time":"2022-03-21T11:33:53.523246","exception":false,"start_time":"2022-03-21T11:33:53.489116","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:41:39.362024Z","iopub.execute_input":"2022-06-29T19:41:39.36264Z","iopub.status.idle":"2022-06-29T19:41:48.851878Z","shell.execute_reply.started":"2022-06-29T19:41:39.362609Z","shell.execute_reply":"2022-06-29T19:41:48.851211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{"id":"f2ed8ef2","papermill":{"duration":0.024963,"end_time":"2022-03-21T11:33:53.573248","exception":false,"start_time":"2022-03-21T11:33:53.548285","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport os\nimport gc\nimport re\nimport ast\nimport sys\nimport copy\nimport json\nimport time\nimport math\nimport shutil\nimport string\nimport pickle\nimport random\nimport joblib\nimport itertools\nfrom pathlib import Path\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 f1_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nimport torch\nprint(f\"torch.__version__: {torch.__version__}\")\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 -y transformers')\nos.system('pip uninstall -y tokenizers')\nos.system('python -m pip install --no-index --find-links=../input/pppm-pip-wheels transformers')\nos.system('python -m pip install --no-index --find-links=../input/pppm-pip-wheels tokenizers')\nimport tokenizers\nimport transformers\nprint(f\"tokenizers.__version__: {tokenizers.__version__}\")\nprint(f\"transformers.__version__: {transformers.__version__}\")\nfrom transformers import AutoTokenizer, AutoModel, AutoConfig\nfrom transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\n%env TOKENIZERS_PARALLELISM=true\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"executionInfo":{"elapsed":20123,"status":"ok","timestamp":1644920080956,"user":{"displayName":"Yasufumi Nakama","photoUrl":"https://lh3.googleusercontent.com/a/default-user=s64","userId":"17486303986134302670"},"user_tz":-540},"id":"35916341","outputId":"06fa0ab8-a380-4f54-a98d-b7015b79d9e2","papermill":{"duration":27.238641,"end_time":"2022-03-21T11:34:20.895812","exception":false,"start_time":"2022-03-21T11:33:53.657171","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:41:48.856316Z","iopub.execute_input":"2022-06-29T19:41:48.858224Z","iopub.status.idle":"2022-06-29T19:42:15.161584Z","shell.execute_reply.started":"2022-06-29T19:41:48.858176Z","shell.execute_reply":"2022-06-29T19:42:15.16082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{"id":"fd586614","papermill":{"duration":0.029385,"end_time":"2022-03-21T11:34:21.160041","exception":false,"start_time":"2022-03-21T11:34:21.130656","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    score = sp.stats.pearsonr(y_true, y_pred)[0]\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=42)","metadata":{"id":"d5c0ccc6","papermill":{"duration":0.041568,"end_time":"2022-03-21T11:34:21.231169","exception":false,"start_time":"2022-03-21T11:34:21.189601","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:42:15.162891Z","iopub.execute_input":"2022-06-29T19:42:15.163205Z","iopub.status.idle":"2022-06-29T19:42:15.178057Z","shell.execute_reply.started":"2022-06-29T19:42:15.163167Z","shell.execute_reply":"2022-06-29T19:42:15.176634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"id":"cb3d8e1e","papermill":{"duration":0.028943,"end_time":"2022-03-21T11:34:21.290864","exception":false,"start_time":"2022-03-21T11:34:21.261921","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Data Loading\n# ====================================================\ntrain = pd.read_csv(INPUT_DIR+'train.csv')\ntest = pd.read_csv(INPUT_DIR+'test.csv')\nsubmission = pd.read_csv(INPUT_DIR+'sample_submission.csv')\nprint(f\"train.shape: {train.shape}\")\nprint(f\"test.shape: {test.shape}\")\nprint(f\"submission.shape: {submission.shape}\")\ndisplay(train.head())\ndisplay(test.head())\ndisplay(submission.head())","metadata":{"executionInfo":{"elapsed":2627,"status":"ok","timestamp":1644920084001,"user":{"displayName":"Yasufumi Nakama","photoUrl":"https://lh3.googleusercontent.com/a/default-user=s64","userId":"17486303986134302670"},"user_tz":-540},"id":"bef012d3","outputId":"d4d60dbc-510c-4f34-8d64-dd1d88c4808c","papermill":{"duration":0.860612,"end_time":"2022-03-21T11:34:22.180396","exception":false,"start_time":"2022-03-21T11:34:21.319784","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-29T19:42:15.180794Z","iopub.execute_input":"2022-06-29T19:42:15.18152Z","iopub.status.idle":"2022-06-29T19:42:15.489408Z","shell.execute_reply.started":"2022-06-29T19:42:15.181479Z","shell.execute_reply":"2022-06-29T19:42:15.488494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_cpc_texts():\n    contexts = []\n    pattern = '[A-Z]\\d+'\n    for file_name in os.listdir('../input/cpc-data/CPCSchemeXML202105'):\n        result = re.findall(pattern, file_name)\n        if result:\n            contexts.append(result)\n    contexts = sorted(set(sum(contexts, [])))\n    results = {}\n    for cpc in ['A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'Y']:\n        with open(f'../input/cpc-data/CPCTitleList202202/cpc-section-{cpc}_20220201.txt') as f:\n            s = f.read()\n        pattern = f'{cpc}\\t\\t.+'\n        result = re.findall(pattern, s)\n        cpc_result = result[0].lstrip(pattern)\n        for context in [c for c in contexts if c[0] == cpc]:\n            pattern = f'{context}\\t\\t.+'\n            result = re.findall(pattern, s)\n            results[context] = cpc_result + \". \" + result[0].lstrip(pattern)\n    return results\n\n\ncpc_texts = get_cpc_texts()\n\ntrain_df = pd.read_csv('../input/us-patent-phrase-to-phrase-matching/train.csv')\ntest_df = pd.read_csv('../input/us-patent-phrase-to-phrase-matching/test.csv')\n\ntrain_df['flag'] = 0\ntest_df['flag'] = 1\ntest_df['score'] = -1\n\nall_df = pd.concat([test_df, train_df], 0)\n\nall_df['context_text'] = all_df['context'].map(cpc_texts).apply(lambda x:x.lower())\nall_df = all_df.join(all_df.groupby('anchor').target.agg(list).rename('ref'), on='anchor')\nall_df['ref2'] = all_df.apply(lambda x:[i for i in x['ref'] if i != x['target']], axis=1)\nall_df['ref2'] = all_df.ref2.apply(lambda x: ', '.join(sorted(list(set(x)), key=x.index)))\nall_df['ref'] = all_df.ref.apply(lambda x:', '.join(sorted(list(set(x)), key=x.index)))\n\nall_df = all_df.join(all_df.groupby(['anchor', 'context']).target.agg(list).rename('ref3'), on=['anchor', 'context'])\nall_df['ref3'] = all_df.apply(lambda x: ', '.join([i for i in x['ref3'] if i != x['target']]), axis=1)\n\nall_df = all_df.join(all_df.groupby('context').anchor.agg('unique').rename('anchor_list'), on='context')\nall_df['anchor_list'] = all_df.apply(lambda x:', '.join([i for i in x['anchor_list'] if i != x['anchor']]), axis=1)\n\nall_df['text1'] = all_df['anchor'] + '[SEP]' + all_df['target'] + '[SEP]'  + all_df['context_text']\nall_df['text2'] = all_df['anchor'] + '[SEP]' + all_df['target'] + '[SEP]'  + all_df['context_text'] + '[SEP]'  + all_df['ref']\nall_df['text3'] = all_df['anchor'] + '[SEP]' + all_df['target'] + '[SEP]'  + all_df['context_text'] + '[SEP]'  + all_df['ref2']\nall_df['text4'] = all_df['anchor'] + '[SEP]' + all_df['target'] + '[SEP]'  + all_df['context_text'] + '[SEP]'  + all_df['ref2'] + ', ' + all_df['anchor_list']\nall_df['text5'] = all_df['anchor'] + '[SEP]' + all_df['target'] + '[SEP]'  + all_df['context_text'] + '[SEP]'  + all_df['ref3']\nall_df['text6'] = 'The similarity between anchor ' + all_df['anchor'] + ' and target ' + all_df['target'] + '. Context is ' + all_df['context_text'] + \\\n            '. Candidates are ' + all_df['ref3']\nall_df","metadata":{"execution":{"iopub.status.busy":"2022-06-29T19:42:15.490756Z","iopub.execute_input":"2022-06-29T19:42:15.491174Z","iopub.status.idle":"2022-06-29T19:42:16.376452Z","shell.execute_reply.started":"2022-06-29T19:42:15.491134Z","shell.execute_reply":"2022-06-29T19:42:16.375772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = all_df[all_df[\"flag\"]==0]\ntrain[\"text\"]  = train[\"text2\"]","metadata":{"execution":{"iopub.status.busy":"2022-06-29T19:42:24.38324Z","iopub.execute_input":"2022-06-29T19:42:24.383774Z","iopub.status.idle":"2022-06-29T19:42:24.416591Z","shell.execute_reply.started":"2022-06-29T19:42:24.383738Z","shell.execute_reply":"2022-06-29T19:42:24.415892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{"id":"9e05b6c4","papermill":{"duration":0.031837,"end_time":"2022-03-21T11:34:22.593468","exception":false,"start_time":"2022-03-21T11:34:22.561631","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CV split\n# ====================================================\ntrain['score_map'] = train['score'].map({0.00: 0, 0.25: 1, 0.50: 2, 0.75: 3, 1.00: 4})\nFold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(train, train['score_map'])):\n    train.loc[val_index, 'fold'] = int(n)\ntrain['fold'] = train['fold'].astype(int)\ndisplay(train.groupby('fold').size())","metadata":{"executionInfo":{"elapsed":12,"status":"ok","timestamp":1644920084528,"user":{"displayName":"Yasufumi Nakama","photoUrl":"https://lh3.googleusercontent.com/a/default-user=s64","userId":"17486303986134302670"},"user_tz":-540},"id":"3ba287c4","outputId":"307dc0e2-17d6-4bfe-9e95-fdd9b6303974","papermill":{"duration":0.052109,"end_time":"2022-03-21T11:34:22.677648","exception":false,"start_time":"2022-03-21T11:34:22.625539","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:19.467386Z","iopub.execute_input":"2022-05-30T01:47:19.467633Z","iopub.status.idle":"2022-05-30T01:47:19.504405Z","shell.execute_reply.started":"2022-05-30T01:47:19.4676Z","shell.execute_reply":"2022-05-30T01:47:19.503678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# tokenizer","metadata":{"id":"918a28aa","papermill":{"duration":0.032412,"end_time":"2022-03-21T11:34:22.813864","exception":false,"start_time":"2022-03-21T11:34:22.781452","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# tokenizer\n# ====================================================\ntokenizer = AutoTokenizer.from_pretrained(CFG.model)\ntokenizer.save_pretrained(OUTPUT_DIR+'tokenizer/')\nCFG.tokenizer = tokenizer","metadata":{"papermill":{"duration":5.588243,"end_time":"2022-03-21T11:34:28.435013","exception":false,"start_time":"2022-03-21T11:34:22.84677","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:19.52309Z","iopub.execute_input":"2022-05-30T01:47:19.523691Z","iopub.status.idle":"2022-05-30T01:47:22.63939Z","shell.execute_reply.started":"2022-05-30T01:47:19.523649Z","shell.execute_reply":"2022-05-30T01:47:22.638619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"14da40cf","papermill":{"duration":0.034425,"end_time":"2022-03-21T11:34:28.504726","exception":false,"start_time":"2022-03-21T11:34:28.470301","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Define max_len\n# ====================================================\nlengths_dict = {}\n\nlengths = []\ntk0 = tqdm(cpc_texts.values(), total=len(cpc_texts))\nfor text in tk0:\n    length = len(tokenizer(text, add_special_tokens=False)['input_ids'])\n    lengths.append(length)\nlengths_dict['context_text'] = lengths\n\nfor text_col in ['anchor', 'target','ref2']:\n    lengths = []\n    tk0 = tqdm(train[text_col].fillna(\"\").values, total=len(train))\n    for text in tk0:\n        length = len(tokenizer(text, add_special_tokens=False)['input_ids'])\n        lengths.append(length)\n    lengths_dict[text_col] = lengths\n    \nCFG.max_len = max(lengths_dict['anchor']) + max(lengths_dict['target'])\\\n                + max(lengths_dict['context_text']) + max(lengths_dict['ref2']) + 4 # CLS + SEP + SEP + SEP\nLOGGER.info(f\"max_len: {CFG.max_len}\")","metadata":{"executionInfo":{"elapsed":32827,"status":"ok","timestamp":1644920122500,"user":{"displayName":"Yasufumi Nakama","photoUrl":"https://lh3.googleusercontent.com/a/default-user=s64","userId":"17486303986134302670"},"user_tz":-540},"id":"c00327b0","outputId":"26e947da-b73a-494d-e776-906b037ac08a","papermill":{"duration":26.672702,"end_time":"2022-03-21T11:34:55.211912","exception":false,"start_time":"2022-03-21T11:34:28.53921","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:22.642024Z","iopub.execute_input":"2022-05-30T01:47:22.642386Z","iopub.status.idle":"2022-05-30T01:47:29.628918Z","shell.execute_reply.started":"2022-05-30T01:47:22.642331Z","shell.execute_reply":"2022-05-30T01:47:29.627983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\ndef prepare_input(cfg, text):\n    inputs = cfg.tokenizer(text,\n                           add_special_tokens=True,\n                           max_length=cfg.max_len,\n                           padding=\"max_length\",\n                           return_offsets_mapping=False)\n    for k, v in inputs.items():\n        inputs[k] = torch.tensor(v, dtype=torch.long)\n    return inputs\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, cfg, df):\n        self.cfg = cfg\n        self.texts = df['text'].values\n        self.labels = df['score'].values\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, item):\n        inputs = prepare_input(self.cfg, self.texts[item])\n        label = torch.tensor(self.labels[item], dtype=torch.float16)\n        return inputs, label","metadata":{"id":"9f791a19","papermill":{"duration":0.053757,"end_time":"2022-03-21T11:34:55.302054","exception":false,"start_time":"2022-03-21T11:34:55.248297","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.630418Z","iopub.execute_input":"2022-05-30T01:47:29.630824Z","iopub.status.idle":"2022-05-30T01:47:29.64162Z","shell.execute_reply.started":"2022-05-30T01:47:29.630785Z","shell.execute_reply":"2022-05-30T01:47:29.640891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntrain_dataset = TrainDataset(CFG, train)\ninputs, label = train_dataset[0]\nprint(inputs)\nprint(label)\n","metadata":{"executionInfo":{"elapsed":8,"status":"ok","timestamp":1644920122808,"user":{"displayName":"Yasufumi Nakama","photoUrl":"https://lh3.googleusercontent.com/a/default-user=s64","userId":"17486303986134302670"},"user_tz":-540},"id":"a200bd5b","outputId":"a30fde2b-86f9-4ebd-ae81-33f467b69836","papermill":{"duration":0.043429,"end_time":"2022-03-21T11:34:55.463328","exception":false,"start_time":"2022-03-21T11:34:55.419899","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.642971Z","iopub.execute_input":"2022-05-30T01:47:29.643271Z","iopub.status.idle":"2022-05-30T01:47:29.721192Z","shell.execute_reply.started":"2022-05-30T01:47:29.643234Z","shell.execute_reply":"2022-05-30T01:47:29.720434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"e04d6363","papermill":{"duration":0.036359,"end_time":"2022-03-21T11:34:55.535904","exception":false,"start_time":"2022-03-21T11:34:55.499545","status":"completed"},"tags":[]}},{"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 = AutoConfig.from_pretrained(cfg.model, output_hidden_states=True)\n        else:\n            self.config = torch.load(config_path)\n        if pretrained:\n            self.model = AutoModel.from_pretrained(cfg.model, config=self.config)\n        else:\n            self.model = AutoModel.from_config(self.config)\n        self.fc_dropout = nn.Dropout(cfg.fc_dropout)\n        self.fc = nn.Linear(self.config.hidden_size, self.cfg.target_size)\n        self._init_weights(self.fc)\n        self.attention = nn.Sequential(\n            nn.Linear(self.config.hidden_size, 512),\n            nn.Tanh(),\n            nn.Linear(512, 1),\n            nn.Softmax(dim=1)\n        )\n        self._init_weights(self.attention)\n        \n    def _init_weights(self, module):\n        if isinstance(module, nn.Linear):\n            module.weight.data.normal_(mean=0.0, std=self.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_(mean=0.0, std=self.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    def feature(self, inputs):\n        outputs = self.model(**inputs)\n        last_hidden_states = outputs[0]\n        # feature = torch.mean(last_hidden_states, 1)\n        weights = self.attention(last_hidden_states)\n        feature = torch.sum(weights * last_hidden_states, dim=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":{"id":"4c5bab44","papermill":{"duration":0.05059,"end_time":"2022-03-21T11:34:55.622912","exception":false,"start_time":"2022-03-21T11:34:55.572322","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.722673Z","iopub.execute_input":"2022-05-30T01:47:29.723114Z","iopub.status.idle":"2022-05-30T01:47:29.73766Z","shell.execute_reply.started":"2022-05-30T01:47:29.723076Z","shell.execute_reply":"2022-05-30T01:47:29.736854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helpler functions","metadata":{"id":"deee9675","papermill":{"duration":0.03649,"end_time":"2022-03-21T11:34:55.940757","exception":false,"start_time":"2022-03-21T11:34:55.904267","status":"completed"},"tags":[]}},{"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, (inputs, labels) in enumerate(train_loader):\n        for k, v in inputs.items():\n            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            y_preds = model(inputs)\n        loss = criterion(y_preds.view(-1, 1), labels.view(-1, 1))\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\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        if CFG.wandb:\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 valid_fn(valid_loader, model, criterion, device):\n    losses = AverageMeter()\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, (inputs, labels) in enumerate(valid_loader):\n        for k, v in inputs.items():\n            inputs[k] = v.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        with torch.no_grad():\n            y_preds = model(inputs)\n        loss = criterion(y_preds.view(-1, 1), labels.view(-1, 1))\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        preds.append(y_preds.sigmoid().to('cpu').numpy())\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    predictions = np.concatenate(predictions)\n    return losses.avg, predictions\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":{"id":"c8263b0c","papermill":{"duration":0.112662,"end_time":"2022-03-21T11:34:56.089768","exception":false,"start_time":"2022-03-21T11:34:55.977106","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.739417Z","iopub.execute_input":"2022-05-30T01:47:29.740007Z","iopub.status.idle":"2022-05-30T01:47:29.772692Z","shell.execute_reply.started":"2022-05-30T01:47:29.739968Z","shell.execute_reply":"2022-05-30T01:47:29.771893Z"},"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    train_folds = folds[folds['fold'] != fold].reset_index(drop=True)\n    valid_folds = folds[folds['fold'] == fold].reset_index(drop=True)\n    valid_labels = valid_folds['score'].values\n    \n    train_dataset = TrainDataset(CFG, train_folds)\n    valid_dataset = TrainDataset(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.BCEWithLogitsLoss(reduction=\"mean\")\n    criterion = nn.MSELoss(reduction=\"mean\")\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        avg_val_loss, predictions = valid_fn(valid_loader, model, criterion, device)\n        \n        # scoring\n        score = get_score(valid_labels, predictions)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                       f\"[fold{fold}] avg_train_loss\": avg_loss, \n                       f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                       f\"[fold{fold}] score\": score})\n        \n        if best_score < score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'model': model.state_dict(),\n                        'predictions': predictions},\n                        OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\")\n\n    predictions = torch.load(OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    valid_folds['pred'] = predictions\n\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return valid_folds","metadata":{"id":"bed940e1","papermill":{"duration":0.300071,"end_time":"2022-03-21T11:34:56.447127","exception":false,"start_time":"2022-03-21T11:34:56.147056","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.786302Z","iopub.execute_input":"2022-05-30T01:47:29.786698Z","iopub.status.idle":"2022-05-30T01:47:29.811139Z","shell.execute_reply.started":"2022-05-30T01:47:29.786669Z","shell.execute_reply":"2022-05-30T01:47:29.810453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    def get_result(oof_df):\n        labels = oof_df['score'].values\n        preds = oof_df['pred'].values\n        score = get_score(labels, preds)\n        LOGGER.info(f'Score: {score:<.4f}')\n    \n    if CFG.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        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        oof_df.to_pickle(OUTPUT_DIR+'oof_df.pkl')\n        \n    if CFG.wandb:\n        wandb.finish()","metadata":{"id":"6cc76b1e","papermill":{"duration":6464.068626,"end_time":"2022-03-21T13:22:40.573692","exception":false,"start_time":"2022-03-21T11:34:56.505066","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-30T01:47:29.814676Z","iopub.execute_input":"2022-05-30T01:47:29.815042Z","iopub.status.idle":"2022-05-30T01:48:42.701496Z","shell.execute_reply.started":"2022-05-30T01:47:29.81501Z","shell.execute_reply":"2022-05-30T01:48:42.700434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}