{"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":"# try:\n# except Exception as e:\n#     from IPython import embed; embed()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:19.235643Z","iopub.execute_input":"2022-08-10T23:33:19.236753Z","iopub.status.idle":"2022-08-10T23:33:19.259647Z","shell.execute_reply.started":"2022-08-10T23:33:19.236652Z","shell.execute_reply":"2022-08-10T23:33:19.258737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, sys, gc\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"\n\nimport re\nimport copy\nimport json\nimport time\nimport math\nimport pickle\nimport random\nimport itertools\nimport warnings\nfrom collections import OrderedDict\nfrom pathlib import Path\nfrom bisect import bisect\nimport multiprocessing\n\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 import tqdm\nfrom sklearn.metrics import f1_score, precision_recall_fscore_support\nfrom sklearn.model_selection import GroupShuffleSplit, GroupKFold\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torchvision\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\n# Transformers\nimport tokenizers\nimport transformers\nfrom transformers import AutoTokenizer, AutoModel, AutoConfig\nfrom transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\n\nfrom transformers import BertConfig\nfrom transformers.models.bert.modeling_bert import BertEncoder\n\nprint(f\"tokenizers.__version__: {tokenizers.__version__}\")\nprint(f\"transformers.__version__: {transformers.__version__}\")\nprint(f\"torch.__version__: {torch.__version__}\")\n\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:19.716265Z","iopub.execute_input":"2022-08-10T23:33:19.717056Z","iopub.status.idle":"2022-08-10T23:33:28.541309Z","shell.execute_reply.started":"2022-08-10T23:33:19.717017Z","shell.execute_reply":"2022-08-10T23:33:28.539806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{"papermill":{"duration":0.021348,"end_time":"2022-02-08T03:58:01.901379","exception":false,"start_time":"2022-02-08T03:58:01.880031","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\n\nclass Args():\n    # Senttng default parameters\n    \n    def __init__(self, args_d=None):\n        if args_d is not None:\n            for k, v in args_d.items():\n                self.__setattr__(k, v)\n            \n        return None\n    \n    def get_argslist(self):\n        arg_list = [p for p in  dir(self) if (not p.startswith('_')) and ( (not p.startswith('get')) and (not p.startswith('_') ))]\n        return arg_list\n    \n    def get_argsdict(self):\n        arg_list = self.get_argslist()\n        arg_dict = {k : getattr(self, k) for k in arg_list}\n        return arg_dict\n    \n    def __repr__(self):\n        s = 10*' =' + ' ArgList' + 10*' =' + '\\n'\n        for k, v in self.get_argsdict().items():\n            s += f'{k:30s}: {v}\\n'\n            \n        return s\n\n    def __str__(self):\n        return self.__repr__()\n    \n    \n    def _save_to_file(self, filepath='args.json'):\n        args_d = self.get_argsdict()\n        \n        for k in args_d.keys():\n            if type(args_d[k]) is np.ndarray:\n                args_d[k] = args_d[k].tolist()\n                \n        with open(filepath, 'w') as f:\n            json.dump(args_d, f)\n        \n        return None\n    \n    def _read_from_file(filepath='args.json'):\n        with open(filepath, 'r') as f:\n            args_d = json.load(f)\n        \n        return Args(args_d)\n    \n    \n    \nclass ModelArgs(Args):\n    do_training = False\n    do_inference = False\n    do_kaggle_inference = True\n    \n    competition       = 'AI4CODE'\n    \n    # Splits\n    n_folds    = [5, 25][1]\n    test_split = [0.15, 0.04][1]\n    split_seed = 1227\n    train_seed = 1457\n    max_dataset_size = [100, 500, 1_000, 2500, 5_000, 10_000, 20_000, 50_000, -1][-1]\n    max_add_dataset_size = [100, -1][-1]\n    \n    # Training\n    epochs       = [4, 15, 20, 30, 50][3]\n    trn_fold     = [0] #[i for i in range(0, n_folds)]\n    batch_size   = [2, 38, 50, 70, 260, 128, 80, 576, 120, 128, 150, 180, 200, 300][-3]\n    val_bs_mult  = 3\n    save_top_k   = 8\n    num_workers  = 0 if multiprocessing.cpu_count() <= 4 else multiprocessing.cpu_count()//2\n    device       = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n    apex         = device != 'cpu'\n    pin_memory   = device != 'cpu'\n    \n    initial_ckpt_path = None #'run_42_MATRIXMdCo_microsoft_codebert_base_PL/microsoft-codebert-base_F0_E15_TLnan_VL1.0837_VS0.8692.ckpt'\n    restore_scheduler = False\n    restore_optimizer = False\n    \n    encoder_lr = [1e-5, 5e-5, 1e-5, 5e-6][-1]\n    decoder_lr = [1e-4, 5e-4, 1e-3, 5e-6][-1]\n    eps   = 1e-7\n    betas = (0.9, 0.999)\n    \n    fc_dropout     = [0.0, 0.15][0]\n    p_rnd_mask_tkn = [0.0, 0.15][1]\n    \n    weight_decay = 0.02\n    gradient_accumulation_steps = 2\n    max_grad_norm = [1000, 500][1]\n    \n    \n    # Scheduler\n    scheduler_type = ['linear', 'cosine'][1]\n    use_scheduler = [True, False][0]\n    num_cycles=0.50\n    num_warmup_steps=0\n    \n    \n    # Criteria\n    criterion = ['MSELoss', 'BCEWithLogitsLoss', 'FocalLoss', 'CrossEntropyLoss', 'PositionLoss', 'ListMLELoss'][-2]\n    fl_gamma = 3.0\n    fl_positve_weighting = -1\n    \n    ploss_scale = 10.0\n    ploss_use_position_weights = True\n    ploss_position_weights_factor = 1.0\n    ploss_position_weights_clamp_value = [100.0, float(batch_size)][1]\n    ploss_use_max_pooling = False\n    ploss_use_abs_pos_enc = True\n    ploss_interpolate_pe  = True\n    ploss_use_pe_projection = False\n    ploss_pe_lenght = 200\n    ploss_pe_dim = None\n    ploss_use_code_context = False\n    ploss_n_code_context = 3\n    \n    \n    # Model\n    model = [\n        \"tals/roberta_python\",\n        \"distilbert-base-multilingual-cased\", # MD1@T64@B260@[B0.7923|M0.7961|E0.869]\n        \"distilbert-base-uncased\",            # MD1@T64@B260@[B0.7920|M0.7987|E0.857]\n        \"microsoft/codebert-base\",            # MD1@T64@B128@[B0.8031|M0.810-|E0.870]\n        \"microsoft/deberta-base\",\n        \"microsoft/deberta-v3-base\",\n        \"roberta-base\",\n        \"roberta-large\",\n        \"microsoft/deberta-v2-xlarge\",\n        \"bert-large-cased\",\n        \"bigscience/bloom-1b3\",\n        \"bigscience/bloom-350m\"\n    ][3]\n    \n    use_batch_head = False\n    batch_head_nhlayers = 4\n    \n    code_use_abs_pos_enc = False\n    code_interpolate_pe = True\n    code_use_pe_projection = False\n    code_pe_lenght = 200\n    code_pe_dim = None    \n        \n    sum_hidden_states = False\n    sum_hidden_states_n = 4\n    text_cleaner_type = [None, 1, 2][1]\n    \n\n    \n    model_type = ['order', 'abs', 'matrix', 'sorting'][-2]\n    d_max = 1\n    order_train_only_md_cells = False\n    \n    sort_n_hidden_layers = 4\n    sort_use_abs_pos_enc = True\n    sort_interpolate_pe  = True\n    sort_use_pe_projection = False\n    sort_pe_lenght = 200\n    sort_pe_dim = None\n    \n    mtx_use_metanotebooks   = True\n    mtx_use_rnd_source_cuts = True\n    \n    abs_train_md_and_code_cells = True\n    abs_md_sample_default_rank  = 0.5\n    \n    max_md_tkn_len   = [None, 121, 150, 64, 64, 32, 32, 77, 64, 128][-1]\n    max_code_tkn_len = [None, 285, 150, 64, 20, 32, 32, 77, 64, 128][-1]\n    max_abs_tkn_len  = [512-3][0]\n    \n    n_md_feat_norm   = 180\n    n_code_feat_norm = 120\n    use_sample_weights = True\n    \n    output_dim = {'matrix': 256, 'sorting':512}.get(model_type, 1)\n    output_l2_embedding = (model_type == 'matrix')\n    \n    # Directories\n    dataset_dir = '../input/AI4Code'\n    add_trn_ds_dir = [None, '../input/GitHubDS_85kNB'][0]\n    working_dir = './'\n#     checkpoints_dir = f'./run_47_{\"DM\" + str(d_max) if model_type==\"order\" else model_type.upper() + (\"MdCo\" if abs_train_md_and_code_cells else \"Md\")}_{model.replace(\"/\", \"_\").replace(\"-\", \"_\")}_{\"\".join( [c for c in criterion if str.isupper(c)] )}'\n    checkpoints_dir = '../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast'\ntry:\n    display\n    \nexcept Exception:\n    display = print\n    \n    \nCFG = ModelArgs()\nCFG.working_dir     = Path( CFG.working_dir )\nCFG.checkpoints_dir = Path( CFG.checkpoints_dir )\nCFG.dataset_dir     = Path( CFG.dataset_dir )\n\nif CFG.add_trn_ds_dir is not None:\n    CFG.add_trn_ds_dir = Path( CFG.add_trn_ds_dir )\n    assert CFG.add_trn_ds_dir.exists()\n    \n\nCFG.working_dir.mkdir(exist_ok=True)\nCFG.checkpoints_dir.mkdir(exist_ok=True)\n\nassert CFG.dataset_dir.exists()\n\nassert CFG.text_cleaner_type in [None, 1, 2]\n\nif CFG.model_type == 'abs':\n    assert CFG.criterion in ['MSELoss']\n    \nelif CFG.model_type == 'order':\n    assert CFG.criterion in ['BCEWithLogitsLoss', 'FocalLoss']\n\nelif CFG.model_type == 'matrix':\n    assert CFG.criterion in ['PositionLoss']\n\nelif CFG.model_type == 'sorting':\n    assert CFG.criterion in ['ListMLELoss']\n    \nelse:\n    raise NotImplementedError(f'{CFG.model_type}??')\n\nif CFG.initial_ckpt_path is not None:\n    assert os.path.exists(CFG.initial_ckpt_path)\n    \nCFG","metadata":{"papermill":{"duration":0.031489,"end_time":"2022-02-08T03:58:01.954404","exception":false,"start_time":"2022-02-08T03:58:01.922915","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:33:29.624065Z","iopub.execute_input":"2022-08-10T23:33:29.625070Z","iopub.status.idle":"2022-08-10T23:33:29.689996Z","shell.execute_reply.started":"2022-08-10T23:33:29.625026Z","shell.execute_reply":"2022-08-10T23:33:29.689045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    ckpt_cfg_d = torch.load('../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/model_cfg.pth')\n\n    for k_ckpt, v_ckpt in ckpt_cfg_d.items():\n        v_now = CFG.get_argsdict().get(k_ckpt, None)\n        if v_ckpt != v_now:\n            print(f'k_ckpt: \"{k_ckpt}\"')\n            print(f'\\tv_ckpt: {v_ckpt}')\n            print(f'\\tv_now:  {v_now}')\n            print()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:29.696050Z","iopub.execute_input":"2022-08-10T23:33:29.696666Z","iopub.status.idle":"2022-08-10T23:33:29.706846Z","shell.execute_reply.started":"2022-08-10T23:33:29.696626Z","shell.execute_reply":"2022-08-10T23:33:29.705744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions for scoring","metadata":{"papermill":{"duration":0.024244,"end_time":"2022-02-08T03:58:28.676381","exception":false,"start_time":"2022-02-08T03:58:28.652137","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Utils","metadata":{"papermill":{"duration":0.024089,"end_time":"2022-02-08T03:58:28.866389","exception":false,"start_time":"2022-02-08T03:58:28.8423","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def save_obj(obj, filename):\n    with open(filename, 'wb') as f:\n        pickle.dump(obj, f)\n    print(f'Saved: {filename}')\n    return None\n\ndef load_obj(filename):\n    with open(filename, 'rb') as f:\n        obj = pickle.load(f)\n    print(f'Loaded: {filename}')\n    return obj","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:29.708629Z","iopub.execute_input":"2022-08-10T23:33:29.709447Z","iopub.status.idle":"2022-08-10T23:33:29.717662Z","shell.execute_reply.started":"2022-08-10T23:33:29.709398Z","shell.execute_reply":"2022-08-10T23:33:29.716661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_logger(filename=os.path.join(CFG.checkpoints_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\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 ","metadata":{"papermill":{"duration":0.039958,"end_time":"2022-02-08T03:58:28.930582","exception":false,"start_time":"2022-02-08T03:58:28.890624","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:33:29.719015Z","iopub.execute_input":"2022-08-10T23:33:29.719943Z","iopub.status.idle":"2022-08-10T23:33:29.738820Z","shell.execute_reply.started":"2022-08-10T23:33:29.719905Z","shell.execute_reply":"2022-08-10T23:33:29.737522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def count_inversions(a):\n    inversions = 0\n    sorted_so_far = []\n    for i, u in enumerate(a):\n        j = bisect(sorted_so_far, u)\n        inversions += i - j\n        sorted_so_far.insert(j, u)\n        \n    return inversions\n\ndef kendall_tau(ground_truth, predictions):\n    assert len(ground_truth) == len(predictions), f\"ground_truth.shape={ground_truth.shape}  predictions.shape={predictions.shape}\"\n    \n    total_inversions = 0\n    total_2max = 0  # twice the maximum possible inversions across all instances\n    for gt, pred in zip(ground_truth, predictions):\n        ranks = [gt.index(x) for x in pred]  # rank predicted order in terms of ground truth\n        total_inversions += count_inversions(ranks)\n        n = len(gt)\n        total_2max += n * (n - 1)\n        \n    return 1 - 4 * total_inversions / total_2max","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:29.740621Z","iopub.execute_input":"2022-08-10T23:33:29.741038Z","iopub.status.idle":"2022-08-10T23:33:29.757900Z","shell.execute_reply.started":"2022-08-10T23:33:29.741001Z","shell.execute_reply":"2022-08-10T23:33:29.756902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_orders_from_prediction(preds_d, ds, code_rank_column='code_rank_pct'):\n    \n    assert code_rank_column in ('code_rank_pct', 'rank_pct', None)\n    \n    pred_df = pd.DataFrame({k:v for k, v in preds_d.items() if k in ['sample_ids_v', 'preds_v']})\n\n    pred_df = pd.concat(\n        [\n            pd.DataFrame( pred_df.sample_ids_v.str.split('_').to_list(), columns=[\"id\", \"cell_id\"]),\n            pred_df.preds_v\n        ],\n        axis=1\n    )\n\n    pred_df.columns = pred_df.columns.to_list()[:2] + ['rank_pct']\n\n    pred_df = pred_df.sort_values(['id', 'rank_pct'])\n    \n    if code_rank_column is not None:\n        code_rank_df = ds.df_code[code_rank_column].reset_index()\n        code_rank_df.columns = list(code_rank_df.columns)[:2] + ['rank_pct']\n\n        # Adding code cells\n        pred_df = pd.concat(\n            [\n                code_rank_df,\n                pred_df\n            ]\n        ).sort_values(['id', 'rank_pct'])\n        \n        \n    pred_df = pred_df.reset_index(drop=True)\n\n    pred_orders = pred_df.groupby('id')['cell_id'].apply(list)\n    \n    return pred_orders, pred_df\n\n\ndef calc_competition_score(preds_d, ds, code_rank_column='code_rank_pct'):\n    \"\"\"\n    code_rank_column in ('code_rank_pct', 'rank_pct', None)\n    \"\"\"\n    pred_orders, pred_df = get_orders_from_prediction(\n        preds_d,\n        ds,\n        code_rank_column=code_rank_column,\n    )\n    \n    gt_orders = ds.get_gt_orders()\n    \n    return kendall_tau(gt_orders, pred_orders)\n\n\ndef calc_acc_score(preds_d, ds):\n    acc = ( (preds_d['preds_v'] > 0.0) == (preds_d['targets_v'] > 0.5) ).mean()\n    return acc\n\n\ndef calc_score(preds_d, ds, cfg):\n    if cfg.model_type == 'abs':\n        return calc_competition_score(\n            preds_d,\n            ds,\n            code_rank_column=(None if cfg.abs_train_md_and_code_cells else 'code_rank_pct')\n        )\n    \n    elif cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n        return calc_competition_score(\n            preds_d,\n            ds,\n            code_rank_column=None,\n        )\n        \n    elif cfg.model_type == 'order':\n        return calc_acc_score(preds_d, ds)\n        \n    else:\n        raise NotImplementedError(f'model_type = {cfg.model_type}')\n        \n    return None\n","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:29.760776Z","iopub.execute_input":"2022-08-10T23:33:29.761187Z","iopub.status.idle":"2022-08-10T23:33:29.779463Z","shell.execute_reply.started":"2022-08-10T23:33:29.761150Z","shell.execute_reply":"2022-08-10T23:33:29.778059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cleaning Texts","metadata":{}},{"cell_type":"code","source":"from nltk.stem import WordNetLemmatizer\n\"\"\"first author: yuanzhe zhou\"\"\"\n\nstemmer = WordNetLemmatizer()\n\ndef preprocess_text_v2(txt):\n    txt = str(txt)\n    \n    # Remove all the special characters\n    txt = re.sub(r'\\W', ' ', txt)\n\n    # remove all single characters\n    txt = re.sub(r'\\s+[a-zA-Z]\\s+', ' ', txt)\n\n    # Remove single characters from the start\n    txt = re.sub(r'\\^[a-zA-Z]\\s+', ' ', txt)\n\n    # Substituting multiple spaces with single space\n    txt = re.sub(r'\\s+', ' ', txt, flags=re.I)\n\n    # Removing prefixed 'b'\n    txt = re.sub(r'^b\\s+', '', txt)\n\n    # Converting to Lowercase\n    txt = txt.lower()\n    \n    # Lemmatization\n    tokens = txt.split()\n    tokens = [stemmer.lemmatize(word) for word in tokens if len(word) > 3]\n    txt = ' '.join(tokens)\n    return txt","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:29.919206Z","iopub.execute_input":"2022-08-10T23:33:29.919567Z","iopub.status.idle":"2022-08-10T23:33:29.932759Z","shell.execute_reply.started":"2022-08-10T23:33:29.919527Z","shell.execute_reply":"2022-08-10T23:33:29.931823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_both(txt, version=1):\n    txt = re.sub('\\n+', '\\n', txt)\n    txt = re.sub('\\t', '', txt)\n    txt = re.sub('\\r', '', txt)    \n    \n    if version == 2:\n        txt = preprocess_text_v2(txt)\n        \n    return txt\n\ndef process_code(txt, version=1):\n    txt = process_both(txt, version=version)\n    \n    txt = re.sub(r'b\"[a-zA-Z0-9+/=\\\\\\n]*\"', 'BINARY', txt)\n    txt = re.sub(r\"b'[a-zA-Z0-9+/=\\\\\\n]*'\", 'BINARY', txt)\n    txt = re.sub(r'.array\\([0-9., \\[\\]]*\\)', '.array(values)', txt)\n    \n    \n    return txt\n\n\ndef process_markdown(txt, version=1):\n    txt = process_both(txt, version=version)\n    \n    txt = re.sub('^#+ ?', '', txt)\n    txt = re.sub(r'\\*+', '', txt)\n    txt = re.sub(r'base64,[a-zA-Z0-9+/=\\\\\\n]*\"?/?', 'IMG', txt)\n    txt = re.sub(r'IMG\\+\"[a-zA-Z0-9+/=\\\\\\n]*\"', 'IMG', txt)\n    txt = re.sub(r'</?\\w*>', '', txt)\n    \n    return txt\n\ndef process_by_type(type_txt, version=1):\n    if type_txt[0] == 'code':\n        return process_code(type_txt[1], version=version)\n    elif type_txt[0] == 'markdown':\n        return process_markdown(type_txt[1], version=version)\n    else:\n        raise NotImplementedError(f'cell_type = {type_txt[0]}')\n\ndef process_all(txt, version=1):\n    txt = process_markdown(txt, version=version)\n    txt = process_code(txt, version=version)\n    \n    return txt","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:30.008147Z","iopub.execute_input":"2022-08-10T23:33:30.008543Z","iopub.status.idle":"2022-08-10T23:33:30.021441Z","shell.execute_reply.started":"2022-08-10T23:33:30.008506Z","shell.execute_reply":"2022-08-10T23:33:30.020451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"papermill":{"duration":0.02389,"end_time":"2022-02-08T03:58:28.982039","exception":false,"start_time":"2022-02-08T03:58:28.958149","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def read_notebook_smp(path):\n    with open(path, 'r') as f:\n        nb_d = json.load(f)\n\n    index_v = np.array( list(nb_d['cell_type'].keys()), dtype=np.object_ )\n    \n    nb_df = pd.DataFrame(\n        {\n            'cell_type': pd.Categorical([nb_d['cell_type'][k] for k in index_v]),\n            'source':    np.array([nb_d['source'][k] for k in index_v]),\n        },\n        index=index_v,\n        ).assign(id=path.stem).rename_axis('cell_id')\n    \n    return nb_df\n\n\ndef read_notebook(path, add_code_rank=True):\n#     nb = (\n#         pd.read_json(\n#             path,\n#             dtype={'cell_type': 'category', 'source': 'str'})\n#         .assign(id=path.stem)\n#         .rename_axis('cell_id')\n#     )\n\n    nb = read_notebook_smp(path)\n    \n    if add_code_rank:\n        f_code = (nb.cell_type == 'code')\n        n_code_cells = f_code.sum()\n        if n_code_cells > 0:\n            code_rank = np.arange(n_code_cells)\n            code_rank_pct = (code_rank / (n_code_cells-1.0)).astype(np.float32)\n            code_rank[n_code_cells:] = -1\n            code_rank_pct[n_code_cells:] = np.nan\n            \n            code_ranks_df = pd.DataFrame(\n                data={\n                    'code_rank': code_rank,\n                    'code_rank_pct': code_rank_pct,\n                },\n                index=nb[f_code].index,\n            ).rename_axis('cell_id')\n            \n            nb = nb.merge(code_ranks_df, how='outer', on='cell_id')\n            \n            nb.code_rank = nb.code_rank.fillna(-1).astype(np.int64)\n    \n    return nb\n\ndef get_ranks(base, derived):\n    return np.array( [base.index(d) for d in derived] )\n\n\ndef read_dataset(CFG, is_training, use_cache=False, text_cleaner_type=None, additional_dataset=False):\n    if additional_dataset:\n        dataset_dir = CFG.add_trn_ds_dir\n        n_max = CFG.max_add_dataset_size\n        split = 'train'\n        \n    else:\n        dataset_dir = CFG.dataset_dir\n        \n        if is_training:\n            n_max = CFG.max_dataset_size\n            split = 'train'\n\n        else:\n            n_max = -1\n            split = 'test'\n    \n    cache_file = CFG.working_dir / f'{str(dataset_dir.stem) + \"_\" if additional_dataset else \"\"}{split}_{n_max if n_max > 0 else \"TOT\"}{\"_CLN\" + (\"v\" + str(text_cleaner_type) if text_cleaner_type > 1 else \"\") if text_cleaner_type is not None else \"\"}.pickle'\n    \n    print('Cache filename:', cache_file)\n    \n    if use_cache and cache_file.exists():\n        df_nb, df_orders, df_ancestors = load_obj(cache_file)\n        \n    else:    \n        paths_v = sorted(\n            list( (dataset_dir / split).glob('*.json') )\n        )\n\n        if n_max != -1:\n            paths_v = paths_v[:n_max]\n\n        notebooks_train = [\n            read_notebook(path) for path in tqdm(paths_v, desc='Reading Notebooks...')\n        ]\n\n        df_nb = (\n            pd.concat(notebooks_train)\n            .set_index('id', append=True)\n            .swaplevel()\n            .sort_index(level='id', sort_remaining=False)\n        )\n\n        if is_training:\n            df_orders = pd.read_csv(\n                dataset_dir / 'train_orders.csv',\n                index_col='id',\n                squeeze=True,\n            ).str.split().to_frame() # Split the string representation of cell_ids into a list\n            \n            df_ancestors = pd.read_csv(dataset_dir / 'train_ancestors.csv', index_col='id')\n        \n            df_orders_ = df_orders.join(\n                df_nb.reset_index('cell_id').groupby('id')['cell_id'].apply(list),\n                how='right',\n            )\n            \n            try:\n                ranks = {}\n                for id_, cell_order, cell_id in df_orders_.itertuples():\n                    rank = get_ranks(cell_order, cell_id)\n                    rank_pct = rank / (len(rank) - 1)\n                    ranks[id_] = {'cell_id': cell_id, 'rank': rank, 'rank_pct':rank_pct}\n            except Exception as e:\n                print(f'ERROR en {id_}', file=sys.stderr)\n                raise e\n\n            df_ranks = (\n                pd.DataFrame\n                .from_dict(ranks, orient='index')\n                .rename_axis('id')\n                .apply(pd.Series.explode)\n                .set_index('cell_id', append=True)\n            )\n            df_ranks['rank'] = df_ranks['rank'].astype(np.int64)\n            df_ranks['rank_pct'] = df_ranks['rank_pct'].astype(np.float32)\n\n            df_nb = df_nb.reset_index().merge(df_ranks, on=[\"id\", \"cell_id\"]).merge(df_ancestors, on=[\"id\"]).set_index([\"id\", \"cell_id\"])\n        \n        else:\n            df_orders = None\n            df_ancestors = None\n            \n        if text_cleaner_type is not None:\n            print('Cleaning text ...')\n            text_cleaner_fn = lambda type_txt: process_by_type(type_txt, version=text_cleaner_type)\n            \n            df_nb['source'] = df_nb[ ['cell_type', 'source'] ].apply(text_cleaner_fn, axis=1)\n            \n        if use_cache:\n            print('Saving cache ...')\n            save_obj( [df_nb, df_orders, df_ancestors], cache_file)\n    \n    return df_nb, df_orders, df_ancestors\n","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:30.177688Z","iopub.execute_input":"2022-08-10T23:33:30.178123Z","iopub.status.idle":"2022-08-10T23:33:30.209890Z","shell.execute_reply.started":"2022-08-10T23:33:30.178084Z","shell.execute_reply":"2022-08-10T23:33:30.208885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_training:\n    LOGGER = get_logger()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:30.256533Z","iopub.execute_input":"2022-08-10T23:33:30.256888Z","iopub.status.idle":"2022-08-10T23:33:30.264222Z","shell.execute_reply.started":"2022-08-10T23:33:30.256854Z","shell.execute_reply":"2022-08-10T23:33:30.263308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_kaggle_inference:\n    df_nb_tst, _, _ = read_dataset(\n        CFG,\n        is_training=False,\n        use_cache=False,\n        text_cleaner_type=CFG.text_cleaner_type,\n    )","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-08-10T23:33:30.339686Z","iopub.execute_input":"2022-08-10T23:33:30.340071Z","iopub.status.idle":"2022-08-10T23:33:30.437616Z","shell.execute_reply.started":"2022-08-10T23:33:30.340036Z","shell.execute_reply":"2022-08-10T23:33:30.436492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kaggle_cache_path = \"../input/ai4code-ds-cache-train-cln\"\nif (CFG.do_training or CFG.do_inference) and os.path.exists(kaggle_cache_path):\n    old_working_dir = CFG.working_dir\n    CFG.working_dir = Path(kaggle_cache_path)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:33:30.439581Z","iopub.execute_input":"2022-08-10T23:33:30.440238Z","iopub.status.idle":"2022-08-10T23:33:30.449064Z","shell.execute_reply.started":"2022-08-10T23:33:30.440201Z","shell.execute_reply":"2022-08-10T23:33:30.448146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if (CFG.do_training or CFG.do_inference):\n    df_nb, df_orders, df_ancestors = read_dataset(\n        CFG,\n        is_training=True,\n        use_cache=True,\n        text_cleaner_type=CFG.text_cleaner_type,\n    )\n    \n    sample_sub = pd.read_csv( CFG.dataset_dir / 'sample_submission.csv', index_col=0)\n        ","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-08-10T23:33:30.511900Z","iopub.execute_input":"2022-08-10T23:33:30.512317Z","iopub.status.idle":"2022-08-10T23:34:06.842568Z","shell.execute_reply.started":"2022-08-10T23:33:30.512281Z","shell.execute_reply":"2022-08-10T23:34:06.841459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference and os.path.exists(kaggle_cache_path):\n    CFG.working_dir = old_working_dir","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:06.845854Z","iopub.execute_input":"2022-08-10T23:34:06.846151Z","iopub.status.idle":"2022-08-10T23:34:06.853097Z","shell.execute_reply.started":"2022-08-10T23:34:06.846123Z","shell.execute_reply":"2022-08-10T23:34:06.849797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Example","metadata":{}},{"cell_type":"code","source":"if CFG.do_training:\n    nb_id = df_nb.index.unique('id')[10]\n\n    # Get an example notebook\n    print('Notebook:', nb_id)\n\n    nb = df_nb.loc[nb_id]\n    if df_orders is None:\n        print(\"Disordered notebook:\")\n\n        display(nb)\n\n    else:\n        print(\"Ordered notebook:\")\n        nb_ordered = nb.loc[ df_orders.loc[nb_id].cell_order ]\n        display( nb_ordered )\n        \n        \n    df_nb['code_rank'].hist(bins=100, range=(1,100))\n    plt.show()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-08-10T23:34:06.854623Z","iopub.execute_input":"2022-08-10T23:34:06.854983Z","iopub.status.idle":"2022-08-10T23:34:06.872530Z","shell.execute_reply.started":"2022-08-10T23:34:06.854947Z","shell.execute_reply":"2022-08-10T23:34:06.871472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{"papermill":{"duration":0.027305,"end_time":"2022-02-08T03:58:30.634713","exception":false,"start_time":"2022-02-08T03:58:30.607408","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if CFG.do_training or CFG.do_inference:\n    # Test split\n    folds = np.zeros(len(df_nb), dtype=np.int64)\n\n    splitter = GroupShuffleSplit(\n        n_splits=1,\n        test_size=CFG.test_split,\n        random_state=CFG.split_seed\n    )\n    trn_val_idx, tst_idx = next(splitter.split(df_nb, groups=df_nb[\"ancestor_id\"]))\n\n    folds[tst_idx] = -1\n\n    # CV split\n    Fold = GroupKFold(\n        n_splits=CFG.n_folds,\n    )\n\n    groups     = df_nb[\"ancestor_id\"][trn_val_idx]\n    global_idx = np.arange( len(df_nb) )[trn_val_idx]\n\n    trn_idx_v, val_idx_v = [], []\n    for i_f, (trn_idx, val_idx) in enumerate(Fold.split(global_idx, groups=groups)):\n        trn_idx_v.append( global_idx[trn_idx] )\n        val_idx_v.append( global_idx[val_idx] )\n\n        folds[val_idx_v[-1]] = i_f\n\n    df_nb['folds'] = folds","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:06.875201Z","iopub.execute_input":"2022-08-10T23:34:06.875532Z","iopub.status.idle":"2022-08-10T23:34:20.370321Z","shell.execute_reply.started":"2022-08-10T23:34:06.875506Z","shell.execute_reply":"2022-08-10T23:34:20.369246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Additional training dataset","metadata":{}},{"cell_type":"code","source":"if (CFG.do_training or CFG.do_inference) and (CFG.add_trn_ds_dir is not None):\n    df_nb_add, df_orders_add, df_ancestors_add = read_dataset(\n        CFG,\n        is_training=True,\n        use_cache=True,\n        additional_dataset=True,\n        text_cleaner_type=CFG.text_cleaner_type,\n    )\n    \n    df_nb_add['folds'] = CFG.n_folds\n    df_nb = pd.concat( [df_nb, df_nb_add] )\n\n    del(df_nb_add)\n    del(df_orders_add)\n    del(df_ancestors_add)\n    \n    del(df_orders)\n    del(df_ancestors)\n    \n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.371729Z","iopub.execute_input":"2022-08-10T23:34:20.372668Z","iopub.status.idle":"2022-08-10T23:34:20.380593Z","shell.execute_reply.started":"2022-08-10T23:34:20.372628Z","shell.execute_reply":"2022-08-10T23:34:20.378906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tokenizer","metadata":{"papermill":{"duration":0.036113,"end_time":"2022-02-08T03:58:30.846837","exception":false,"start_time":"2022-02-08T03:58:30.810724","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# tokenizer\n# ====================================================\n\ntokenizer_path = CFG.checkpoints_dir / f\"tokenizer-{CFG.model.replace('/', '-')}\"\n\nif not tokenizer_path.exists():\n    tokenizer = AutoTokenizer.from_pretrained(CFG.model)\n    tokenizer.save_pretrained(tokenizer_path)\n    \nelse:\n    print('Loading:', tokenizer_path)\n    tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)\n","metadata":{"papermill":{"duration":8.030176,"end_time":"2022-02-08T03:58:38.904883","exception":false,"start_time":"2022-02-08T03:58:30.874707","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:34:20.381842Z","iopub.execute_input":"2022-08-10T23:34:20.382718Z","iopub.status.idle":"2022-08-10T23:34:20.585443Z","shell.execute_reply.started":"2022-08-10T23:34:20.382679Z","shell.execute_reply":"2022-08-10T23:34:20.584250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#tokenizer(\n#    'hello',\n#    'love',\n#    return_token_type_ids=True,\n#    max_length=10,\n#    padding=\"max_length\",\n#    truncation=True,\n#)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.587287Z","iopub.execute_input":"2022-08-10T23:34:20.587738Z","iopub.status.idle":"2022-08-10T23:34:20.596840Z","shell.execute_reply.started":"2022-08-10T23:34:20.587701Z","shell.execute_reply":"2022-08-10T23:34:20.595907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TokenLen","metadata":{"papermill":{"duration":0.034998,"end_time":"2022-02-08T03:58:38.971851","exception":false,"start_time":"2022-02-08T03:58:38.936853","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_token_len(df_nb, tokenizer_code, tokenizer_md):\n    \n    cell_type_v = df_nb['cell_type'].values\n    text_v      = df_nb['source'].fillna(\"\").values\n    len_v       = np.zeros(len(df_nb), dtype=np.int64)\n    \n    itt = tqdm(\n        enumerate( zip(cell_type_v, text_v) ),\n        total=len(text_v),\n    )\n    \n    for i_t, (cell_type, text) in itt:\n        if cell_type == 'code':\n            n = len(tokenizer_code.encode(text, add_special_tokens=False))\n            \n        elif cell_type == 'markdown':\n            n = len(tokenizer_md.encode(  text, add_special_tokens=False))\n            \n        else:\n            raise NotImplementedError(f'cell_type = {cell_type}')\n            \n        len_v[i_t] = n\n\n    return len_v\n\ndef calc_max_tkn_len(tkn_len_v, inc_pct_max=10.0, verbose=True):\n    q_vs_tknlen_v = []\n    for q in np.linspace(0.8, 1.0, 100):\n        q_vs_tknlen_v.append( [q, np.quantile(tkn_len_v, q)] )\n\n    q_vs_tknlen_v = np.array(q_vs_tknlen_v)\n\n    slope_v = np.gradient(q_vs_tknlen_v[:,1], q_vs_tknlen_v[:,0], edge_order=2)[1:]\n    x_v = 0.5 * (q_vs_tknlen_v[1:,0] + q_vs_tknlen_v[:-1,0])\n    inc_pct = slope_v/q_vs_tknlen_v[1:,1]\n\n    q_min = round( x_v[ min([i for i in range(len(inc_pct)) if (inc_pct > inc_pct_max)[i:].all() ]) ], 2)\n    l_max = int( np.quantile(tkn_len_v, q_min) )\n    \n    if verbose:\n        print(f'l_max = {l_max}, q_min={q_min:0.2f}')\n        plt.plot(q_vs_tknlen_v[:,0], q_vs_tknlen_v[:,1], label='q vs tknlen')\n        plt.plot(x_v, inc_pct, label='inc pct')\n        plt.plot([q_min, q_min], [0,500], label=f'q_min={q_min:0.2f}')\n        plt.plot([0.8, 1.0], [l_max,l_max], label=f'l_max={l_max:d}')\n        plt.ylim((0, 500))\n        plt.legend()\n        plt.grid()\n        plt.show()\n        \n    return l_max","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.598615Z","iopub.execute_input":"2022-08-10T23:34:20.599086Z","iopub.status.idle":"2022-08-10T23:34:20.618344Z","shell.execute_reply.started":"2022-08-10T23:34:20.599014Z","shell.execute_reply":"2022-08-10T23:34:20.617368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if (CFG.max_md_tkn_len is None) or (CFG.max_code_tkn_len is None):\n    len_v = get_token_len(\n        df_nb,\n        tokenizer_code,\n        tokenizer_md\n    )\n    \n    df_nb['token_len'] = len_v\n    \n    md_tkn_len_v = df_nb.token_len[df_nb.cell_type=='markdown'].values\n    code_tkn_len_v = df_nb.token_len[df_nb.cell_type=='code'].values\n    \n    print('Markdown length')\n    CFG.max_md_tkn_len = calc_max_tkn_len(tkn_len_v=md_tkn_len_v, inc_pct_max=7.0, verbose=True)\n    \n    print('Code length')\n    CFG.max_code_tkn_len = calc_max_tkn_len(tkn_len_v=code_tkn_len_v, inc_pct_max=7.0, verbose=True)\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-08-10T23:34:20.620082Z","iopub.execute_input":"2022-08-10T23:34:20.620761Z","iopub.status.idle":"2022-08-10T23:34:20.631788Z","shell.execute_reply.started":"2022-08-10T23:34:20.620722Z","shell.execute_reply":"2022-08-10T23:34:20.630896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# New Features","metadata":{}},{"cell_type":"code","source":"def calc_features(cfg, df, tokenizer, use_cache=True):\n    all_nb_id_v = np.unique( [nb_id for nb_id, cell_id in df.index.values] )\n    n_nb = len(all_nb_id_v)\n    \n    max_code_tkn_len = cfg.max_code_tkn_len\n    n_md_feat_norm = cfg.n_md_feat_norm\n    n_code_feat_norm = cfg.n_code_feat_norm\n    \n    cache_path = cfg.working_dir / f'feat_cache_{tokenizer.__class__.__name__}_{n_nb}_{max_code_tkn_len}_{n_md_feat_norm}_{n_code_feat_norm}.pickle'\n    \n    if (not use_cache) or (not cache_path.exists()):\n        nb_feat_d = {}\n        for nb_id in tqdm(all_nb_id_v):\n            nb = df.loc[nb_id]\n\n            n_cells = len(nb)\n\n            f_code = (nb.cell_type == 'code')\n            n_code = f_code.sum()\n            n_md = n_cells\n\n            code_text_v = nb[f_code].source.values.tolist()\n\n            batch_code_tkn_v = tokenizer.batch_encode_plus(\n                batch_text_or_text_pairs=code_text_v,\n                add_special_tokens=False,\n                padding=False,\n                truncation=True,\n                max_length=max_code_tkn_len,\n                return_token_type_ids=False,\n                return_attention_mask=False,\n            )['input_ids']\n\n            code_tkn_v = []\n            for i_c, code_tkn in enumerate(batch_code_tkn_v):\n                if i_c > 0:\n                    code_tkn_v.append( tokenizer.sep_token_id )\n                    \n                code_tkn_v.extend( code_tkn )\n\n            feat_d = {\n                'n_md':n_md,\n                'n_code':n_code,\n                'n_cells':n_cells,\n                'code_tkn_v':code_tkn_v,\n                'feat_v': np.array(\n                    [\n                        min(1.0, n_md/n_md_feat_norm),\n                        min(1.0, n_code/n_code_feat_norm),\n                        min(1.0, n_cells/(n_md_feat_norm + n_code_feat_norm)),\n                    ],\n                    dtype=np.float32\n                )\n            }\n\n            nb_feat_d[nb_id] = feat_d\n            \n        if use_cache:\n            save_obj(nb_feat_d, cache_path)\n\n    else:\n        nb_feat_d = load_obj(cache_path)\n\n    return nb_feat_d","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.636366Z","iopub.execute_input":"2022-08-10T23:34:20.636789Z","iopub.status.idle":"2022-08-10T23:34:20.648077Z","shell.execute_reply.started":"2022-08-10T23:34:20.636763Z","shell.execute_reply":"2022-08-10T23:34:20.646746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.model_type == 'abs' and (CFG.do_training or CFG.do_inference):\n    nb_feat_d = calc_features(\n        CFG, \n        df_nb,\n        tokenizer,\n        use_cache=True\n    )\n    \n    # Features stats\n    n_code_tkn_v = []\n    for nb_id in nb_feat_d.keys():\n        n_code_tkn_v.append( len( nb_feat_d[nb_id]['code_tkn_v'] ) )\n\n    n_code_tkn_v = np.array(n_code_tkn_v)\n\n    _ = plt.hist(n_code_tkn_v, bins=100, range=(0, CFG.max_abs_tkn_len-CFG.max_md_tkn_len))\n    plt.show()\n\n    prop_full_nb = (np.array(n_code_tkn_v) < CFG.max_abs_tkn_len-CFG.max_md_tkn_len).mean()\n    print( f\"prop_full_nb = {prop_full_nb * 100:0.01f}%\" )\n    \nelse:\n    nb_feat_d = None\n    \n    \nif CFG.model_type == 'abs' and CFG.do_kaggle_inference:\n    nb_feat_d = calc_features(\n        CFG,\n        df_nb_tst,\n        tokenizer,\n        use_cache=False\n    )\n    \n","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.649434Z","iopub.execute_input":"2022-08-10T23:34:20.650479Z","iopub.status.idle":"2022-08-10T23:34:20.662710Z","shell.execute_reply.started":"2022-08-10T23:34:20.650384Z","shell.execute_reply":"2022-08-10T23:34:20.661722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def join_tkn_seq(\n    sample_a,\n    sample_b,\n    tokenizer,\n    n_pad=None,\n    ):\n    \n    if tokenizer.__class__.__name__ == 'DistilBertTokenizerFast':\n        input_ids      = [tokenizer.cls_token_id] + sample_a + [tokenizer.sep_token_id] + sample_b + [tokenizer.sep_token_id]\n        attention_mask = [1 for _ in input_ids]\n        token_type_ids = [0] * (len(sample_a) + 2) + [1] * (len(sample_b) + 1)\n        \n        if n_pad is not None:\n            n_pad += 3\n        \n        return_token_type_ids = False\n        \n        \n    elif tokenizer.__class__.__name__ == 'RobertaTokenizerFast':\n        input_ids      = [tokenizer.bos_token_id] + sample_a + [tokenizer.eos_token_id, tokenizer.sep_token_id] + sample_b + [tokenizer.eos_token_id]\n        attention_mask = [1 for _ in input_ids]\n        token_type_ids = [0] * (len(sample_a) + 2) + [1] * (len(sample_b) + 2)\n        \n        if n_pad is not None:\n            n_pad += 4\n        \n        return_token_type_ids = False\n        \n    else:\n        raise NotImplementedError(f'tokenizer = {tokenizer.__class__.__name__}')\n    \n    \n    if n_pad is not None:\n        if len(input_ids) < n_pad:\n            to_pad = n_pad - len(input_ids)\n            input_ids      = input_ids      + [tokenizer.pad_token_id] * to_pad\n            attention_mask = attention_mask + [0]                      * to_pad\n            token_type_ids = token_type_ids + [0]                      * to_pad\n\n        elif len(input_ids) > n_pad:          \n            input_ids      = input_ids[:n_pad]\n            attention_mask = attention_mask[:n_pad]\n            token_type_ids = token_type_ids[:n_pad]\n\n    inputs = transformers.tokenization_utils_base.BatchEncoding(\n        {\n            'input_ids':      torch.tensor(input_ids,      dtype=torch.long),\n            'attention_mask': torch.tensor(attention_mask, dtype=torch.long),\n            'token_type_ids': torch.tensor(token_type_ids, dtype=torch.long),\n        }\n    )\n    \n    if not return_token_type_ids:\n        del( inputs['token_type_ids'] )\n    \n    return inputs","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.664108Z","iopub.execute_input":"2022-08-10T23:34:20.664734Z","iopub.status.idle":"2022-08-10T23:34:20.677446Z","shell.execute_reply.started":"2022-08-10T23:34:20.664699Z","shell.execute_reply":"2022-08-10T23:34:20.676463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AbsTrainDataset():\n    def __init__(\n        self,\n        df,\n        nb_feat_d,\n        cfg,\n        tokenizer,\n        i_fold=0,\n        training=False,\n        only_md_samples=False,\n        md_sample_default_rank=0.5,\n    ):\n        \n        self.tokenizer = tokenizer\n        self.cfg = cfg\n        self.nb_feat_d = nb_feat_d\n        \n        # #\n        n_cells_squared = sum( [self.nb_feat_d[nb_id]['n_cells'] ** 2 for nb_id in self.nb_feat_d.keys()] )\n        n_cells         = sum( [self.nb_feat_d[nb_id]['n_cells'] for nb_id in self.nb_feat_d.keys()] )\n        n_nb = len(self.nb_feat_d.keys())\n        for nb_id in self.nb_feat_d.keys():\n            self.nb_feat_d[nb_id]['w'] = min(200, self.nb_feat_d[nb_id]['n_cells'] ** 2 / n_cells_squared * n_cells)\n        # #\n        \n        \n        if i_fold is None:\n            self.df = df\n            \n        else:\n            if training:\n                self.df = df[(df.folds != i_fold) & (df.folds >= 0)]\n            else:\n                self.df = df[(df.folds == i_fold)]\n        \n        self.df_md   = self.df[(self.df.cell_type == 'markdown')]\n        self.df_code = self.df[(self.df.cell_type == 'code')]\n            \n        if only_md_samples:\n            self.sample_nb_id_v = self.df_md.index.map(lambda x: x[0]).values\n            self.sample_id_v    = self.df_md.index.map(lambda x: '_'.join(x)).values\n            self.text_v         = self.df_md['source']\n            self.code_rank_v    = self.df_md['code_rank_pct']\n            \n            if 'rank_pct' in self.df.columns:\n                self.target_v = self.df_md['rank_pct'].astype(np.float32)\n\n            else:\n                self.target_v = None\n            \n        else:\n            self.sample_id_v    = self.df.index.map(lambda x: '_'.join(x)).values\n            self.sample_nb_id_v = self.df.index.map(lambda x: x[0]).values\n            self.text_v         = self.df['source']\n            self.code_rank_v    = self.df['code_rank_pct']\n            \n            if 'rank_pct' in self.df.columns:\n                self.target_v = self.df['rank_pct'].astype(np.float32)\n\n            else:\n                self.target_v = None\n        \n        \n        self.code_rank_v = np.nan_to_num(\n            self.code_rank_v.values,\n            nan=md_sample_default_rank,\n        ).astype(np.float32)\n        \n        return None\n    \n    \n    def __getitem__(self, idx):\n        \n        sample_id    = self.sample_id_v[idx]\n        sample_nb_id = self.sample_nb_id_v[idx]\n        \n        sample_a = self.tokenizer.encode(\n            text=self.text_v[idx],\n            add_special_tokens=False,\n            max_length=self.cfg.max_md_tkn_len,\n            padding=False,\n            truncation=True,\n            return_offsets_mapping=False,\n        )\n        sample_b = self.nb_feat_d[sample_nb_id]['code_tkn_v']\n        n_cells  = self.nb_feat_d[sample_nb_id]['n_cells']\n        \n        \n        inputs = join_tkn_seq(\n            sample_a,\n            sample_b,\n            tokenizer,\n            n_pad=self.cfg.max_abs_tkn_len,\n        )\n   \n        inputs['features'] = torch.tensor(\n            np.concatenate(\n                [\n                    self.code_rank_v[idx, None],\n                    self.nb_feat_d[sample_nb_id]['feat_v'],\n                ],\n                axis=0\n            ),\n            dtype=torch.float32,\n        )\n        \n        data = {\n            'sample_id': sample_id,\n            'inputs': inputs,\n            'w': self.nb_feat_d[sample_nb_id]['w'],\n        }\n        \n        \n        if self.target_v is not None:\n            target = self.target_v[idx]\n            data['target'] = torch.tensor(target, dtype=torch.float32)\n            \n        return data\n\n    \n    def collate_fn(self, data_v):\n        ret_data_d = {k: list() for k in data_v[0].keys()}\n        \n        for k in ret_data_d.keys():\n            for i_s in range(len(data_v)):\n                ret_data_d[k].append( data_v[i_s][k] )\n                \n        for k in ret_data_d.keys():\n#             print(k, type(ret_data_d[k][0]) )\n            if type(ret_data_d[k][0]) in [np.ndarray, np.int64, np.int32, np.float64, np.int32,  int, float]:\n                ret_data_d[k] = torch.tensor( ret_data_d[k] )\n            \n            elif type(ret_data_d[k][0]) is torch.Tensor:\n                ret_data_d[k] = torch.stack( ret_data_d[k] )\n                \n            elif type(ret_data_d[k][0]) is transformers.tokenization_utils_base.BatchEncoding:\n                ret_data_d[k] = transformers.tokenization_utils_base.BatchEncoding(\n                    self.collate_fn( ret_data_d[k] )\n                )\n                \n            else:\n                ret_data_d[k] = np.array( ret_data_d[k] )\n            \n        return ret_data_d\n    \n    \n    def __len__(self):\n        return len(self.sample_id_v)\n    \n    def get_gt_orders(self):\n        gt_orders = self.df.reset_index().sort_values(['id', 'rank_pct']).groupby('id')['cell_id'].apply(list)\n        return gt_orders\n    \n# ds_trn = AbsTrainDataset(\n#     df=df_nb,\n#     nb_feat_d=nb_feat_d,\n#     cfg=CFG,\n#     tokenizer=tokenizer,\n#     i_fold=0,\n#     training=True,\n#     only_md_samples=False,\n#     md_sample_default_rank=0.5,\n# )\n# dl_trn = DataLoader( ds_trn, batch_size=32, collate_fn=ds_trn.collate_fn, shuffle=True)\n\n# for data in dl_trn:\n#     break","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-08-10T23:34:20.678903Z","iopub.execute_input":"2022-08-10T23:34:20.679536Z","iopub.status.idle":"2022-08-10T23:34:20.705163Z","shell.execute_reply.started":"2022-08-10T23:34:20.679501Z","shell.execute_reply":"2022-08-10T23:34:20.704131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class OrderTrainDataset():\n    def __init__(\n        self,\n        df,\n        cfg,\n        tokenizer,\n        i_fold=0,\n        training=False,\n        p_rnd_mask_tkn=0.0,\n        only_md_cells=False,\n    ):\n        \n        self.tokenizer = tokenizer\n        self.cfg = cfg\n        \n        self.d_max = self.cfg.d_max\n        self.i_fold = i_fold\n        self.training = training        \n        self.p_rnd_mask_tkn = p_rnd_mask_tkn\n        self.only_md_cells = only_md_cells\n        \n        if i_fold is None:\n            self.df = df\n            \n        else:\n            if training:\n                self.df = df[(df.folds != i_fold) & (df.folds >= 0)]\n            else:\n                self.df = df[(df.folds == i_fold)]\n        \n        \n        if self.only_md_cells:\n            self.df = self.df[self.df['cell_type'] != 'code']\n        \n        self.sorted_df = self.df.sort_values( ['id', 'rank'] ).reset_index()\n        \n        self.add_idx_matrix()\n        \n        return None\n    \n    \n    def add_idx_matrix(self, use_cache=True):\n        cache_path = self.cfg.working_dir / f'idx_matrix_L{len(self.sorted_df)}_D{self.d_max}_{\"TRN\" if self.training else \"noTRN\"}_F{self.i_fold}_{\"MD\" if self.only_md_cells else \"\"}.pickle'\n        \n        if (not use_cache) or (not cache_path.exists()):\n            idx_matrix = []\n            last_nb_id = ''\n            i_start = 0\n            n_cells = 0\n            \n            n_samples = len(self.sorted_df)\n            for i_global, (nb_id, cell_id, rank) in enumerate(\n                tqdm(\n                    self.sorted_df[['id', 'cell_id', 'rank']].values,\n                    desc='id_matrix ...',\n                )\n            ):\n                \n                if (nb_id != last_nb_id) or (i_global + 1 == n_samples):\n                    \n                    i_end = n_samples if (i_global + 1 == n_samples) else i_global\n                    \n                    if i_end - i_start > self.d_max + 1:\n                        for i_ in range(1, n_cells+1):\n                            idx_matrix[-i_][-1] = i_end\n\n                        for _ in range(self.d_max):\n                            if (len(idx_matrix) > 0) and (idx_matrix[-1][1] == last_nb_id):\n                                idx_matrix.pop()\n                                \n                    else:\n                        # We dont have enoght room to build anchor, positives, neg sample \n                        for i_ in range(1, n_cells+1):\n                            idx_matrix.pop()\n                        \n                    last_nb_id = nb_id\n                    i_start = i_global\n                    n_cells = 0\n                    \n                    \n                if (i_global + self.d_max < n_samples):\n                    idx_matrix.append( [i_global, nb_id, cell_id, rank, i_start, -1] )\n                    n_cells += 1\n\n            idx_matrix = np.array( idx_matrix, dtype=np.object )\n            \n            if use_cache:\n                save_obj(idx_matrix, cache_path)\n                \n        else:\n            idx_matrix = load_obj( cache_path )\n\n        self.idx_matrix = idx_matrix\n        return idx_matrix\n    \n\n    def tokenize_sample(self, text, text_pair):\n        inputs = self.tokenizer.encode(\n            text=text,\n            text_pair=text_pair, \n            add_special_tokens=True,\n            max_length=max(self.cfg.max_md_tkn_len, self.cfg.max_code_tkn_len),\n            padding=False,\n            truncation=True,\n        )\n\n        return inputs\n        \n        \n    def tokenize_samples(self, sample_a, sample_b):\n        txt_a = sample_a.source\n        txt_b = sample_b.source\n        \n        if sample_a.cell_type == 'code':\n            pair_a = f'code cell {sample_a.code_rank}'\n        else:\n            pair_a = 'text cell'\n        \n        if sample_b.cell_type == 'code':\n            pair_b = f'code cell {sample_b.code_rank}'\n        else:\n            pair_b = 'text cell'\n        \n        tkn_a = self.tokenize_sample(txt_a, pair_a)[1:-1]\n        tkn_b = self.tokenize_sample(txt_b, pair_b)[1:-1]\n        \n        inputs = join_tkn_seq(\n            tkn_a,\n            tkn_b,\n            tokenizer=self.tokenizer,\n            n_pad=2*max(self.cfg.max_code_tkn_len, self.cfg.max_md_tkn_len),\n        )\n        \n        return inputs\n    \n    \n    def _mask_tokens(self, inputs):\n        for i in range(len(inputs['input_ids'])):\n            if inputs['input_ids'][i] not in self.tokenizer.all_special_ids:\n                if np.random.random() < self.p_rnd_mask_tkn:\n                    inputs['input_ids'][i] = self.tokenizer.mask_token_id\n\n            elif inputs['input_ids'][i] == self.tokenizer.pad_token_id:\n                break\n                        \n        return None\n            \n        \n    def __getitem__(self, idx):\n        idx, d = divmod(idx, self.d_max)\n        d += 1\n        \n        i_pos, nb_id, cell_id, rank, i_start, i_end = self.idx_matrix[idx]\n        \n        i_neg = i_pos\n        while i_pos <= i_neg <= i_pos + self.d_max:\n            i_neg = random.randint(i_start, i_end - 1)\n        \n        anc_sample = self.sorted_df.iloc[i_pos]\n        pos_sample = self.sorted_df.iloc[i_pos + d]\n        neg_sample = self.sorted_df.iloc[i_neg]\n        \n        pos_tkn = self.tokenize_samples(anc_sample, pos_sample)\n        neg_tkn = self.tokenize_samples(anc_sample, neg_sample)\n        \n        # Masking some tokens\n        if self.training and self.p_rnd_mask_tkn > 0.0: \n            self._mask_tokens(pos_tkn)\n            self._mask_tokens(neg_tkn)\n            \n        data = {\n            'anc_sample': anc_sample,\n            'pos_sample': pos_sample,\n            'neg_sample': neg_sample,\n            \n            'pos_tkn': pos_tkn,\n            'neg_tkn': neg_tkn,\n            \n            'anc_rank': anc_sample['rank'],\n            'pos_rank': pos_sample['rank'],\n            'neg_rank': neg_sample['rank'],\n        }\n            \n        return data\n    \n\n    def collate_fn(self, data_v):\n#         print('len: ', len(data_v), 'type', type(data_v), data_v[0].keys())\n        ret_data_d = {k: list() for k in data_v[0].keys()}\n        \n        for k in ret_data_d.keys():\n            for i_s in range(len(data_v)):\n                ret_data_d[k].append( data_v[i_s][k] )\n                \n        for k in ret_data_d.keys():\n            #print(k, type(ret_data_d[k][0]) )\n            if type(ret_data_d[k][0]) in [np.ndarray, np.int64, np.int32, np.float64, np.int32,  int, float]:\n                ret_data_d[k] = torch.tensor( ret_data_d[k] )\n            \n            elif type(ret_data_d[k][0]) is torch.Tensor:\n                ret_data_d[k] = torch.stack( ret_data_d[k] )\n                \n            elif type(ret_data_d[k][0]) is transformers.tokenization_utils_base.BatchEncoding:\n                ret_data_d[k] = transformers.tokenization_utils_base.BatchEncoding(\n                    self.collate_fn( ret_data_d[k] )\n                )\n                \n            else:\n                ret_data_d[k] = np.array( ret_data_d[k] )\n        \n        ## \n        \n        if 'pos_tkn' in ret_data_d.keys():\n            BS = ret_data_d['pos_tkn']['input_ids'].shape[0]\n\n            ret_data_d['inputs'] = transformers.tokenization_utils_base.BatchEncoding(\n                {\n                    input_key: torch.cat(\n                        [\n                            ret_data_d['pos_tkn'][input_key],\n                            ret_data_d['neg_tkn'][input_key],\n                        ]\n                    )\n                    for input_key in ret_data_d['pos_tkn'].keys()\n                }\n            )\n\n            ret_data_d['target'] = torch.cat(\n                [\n                    torch.ones(BS),\n                    torch.zeros(BS),\n                ]\n            )\n            \n            sample_id = []\n            for a,b in zip(ret_data_d['anc_sample'][:, :2], ret_data_d['pos_sample'][:, :2]):\n                sample_id.append('_'.join(a) + '|' + '_'.join(b))\n\n            for a,b in zip(ret_data_d['anc_sample'][:, :2], ret_data_d['neg_sample'][:, :2]):\n                sample_id.append('_'.join(a) + '|' + '_'.join(b))\n                \n            ret_data_d['sample_id'] = np.array(sample_id)\n        \n        return ret_data_d\n    \n    def __len__(self):\n        return len(self.idx_matrix) * self.d_max\n    \n    def get_gt_orders(self):\n        gt_orders = self.df.reset_index().sort_values(['id', 'rank_pct']).groupby('id')['cell_id'].apply(list)\n        return gt_orders\n    \n    \n# ds_trn = OrderTrainDataset(\n#     df=df_nb,\n#     cfg=CFG,\n#     tokenizer=tokenizer,\n#     i_fold=0,\n#     training=True,\n#     p_rnd_mask_tkn=0.0,\n#     only_md_cells=True,\n# )\n# dl_trn = DataLoader(ds_trn, batch_size=2, collate_fn=ds_trn.collate_fn)\n\n# for data in dl_trn:\n#     break\n\n#ds_trn, ds_val, ds_tst= get_datasets(df_nb, i_fold=0, cfg=CFG, tokenizer=tokenizer)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.706874Z","iopub.execute_input":"2022-08-10T23:34:20.707308Z","iopub.status.idle":"2022-08-10T23:34:20.743708Z","shell.execute_reply.started":"2022-08-10T23:34:20.707263Z","shell.execute_reply":"2022-08-10T23:34:20.742761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def shrink_nb(nb_df_sorted, max_n_cells=128):\n    n_cells = len(nb_df_sorted)\n    \n    if n_cells > max_n_cells:\n        n_cut = n_cells - max_n_cells \n        f = np.ones(n_cells, dtype=np.bool)\n        f[np.random.permutation(n_cells)[:n_cut]] = False\n\n        nb_df_sorted = nb_df_sorted.iloc[f]\n        \n    return nb_df_sorted\n\n\ndef re_rank_nb(nb_df_sorted):\n    n_cells = len(nb_df_sorted)\n\n    rank = np.arange(n_cells, dtype=np.int64)\n    rank_pct = rank / (n_cells-1)\n\n    f_code = (nb_df_sorted.cell_type == 'code')\n    n_code = f_code.sum()\n\n    code_rank     = - np.ones(n_cells, dtype=np.int64)\n    code_rank[f_code]     = np.arange(n_code, dtype=np.int64)\n    \n    code_rank_pct = np.nan * np.ones(n_cells, dtype=np.float32)\n    code_rank_pct[f_code] = np.arange(n_code, dtype=np.float32) / (n_code-1)\n\n    nb_df_sorted['rank'] = rank\n    nb_df_sorted['rank_pct'] = rank_pct\n\n    nb_df_sorted['code_rank'] = code_rank\n    nb_df_sorted['code_rank_pct'] = code_rank_pct\n\n    nb_df = nb_df_sorted.sort_values('code_rank_pct')\n    \n    return nb_df\n\ndef merge_nbs(nb_df_sorted, nb_df_pair_sorted):\n    n_cells = len( nb_df_sorted )\n\n    assert len( nb_df_pair_sorted ) > 4\n\n    if n_cells > 1:\n        i_join = random.randint(1, n_cells-1)\n        \n    else:\n        i_join = 1\n\n    nb_df_meged_sorted = pd.concat(\n        [\n            nb_df_sorted.iloc[:i_join],\n            nb_df_pair_sorted.iloc[2:-2],\n            nb_df_sorted.iloc[i_join:],\n        ],\n        axis=0,\n    )\n    \n    return nb_df_meged_sorted","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.745181Z","iopub.execute_input":"2022-08-10T23:34:20.746531Z","iopub.status.idle":"2022-08-10T23:34:20.758856Z","shell.execute_reply.started":"2022-08-10T23:34:20.746481Z","shell.execute_reply":"2022-08-10T23:34:20.757847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MatrixTrainDataset():\n    def __init__(\n        self,\n        df,\n        cfg,\n        tokenizer,\n        i_fold=0,\n        training=False,\n        md_sample_default_rank=0.5,\n        p_rnd_mask_tkn=0.0,\n        global_rank_target=False,\n        use_metanotebooks=False,\n        use_rnd_source_cuts=False,\n        max_cell_len=None,\n    ):\n        \n        self.tokenizer = tokenizer\n        self.cfg = cfg\n        self.md_sample_default_rank = md_sample_default_rank\n        self.p_rnd_mask_tkn = p_rnd_mask_tkn\n        self.training = training\n        self.global_rank_target = global_rank_target\n        \n        self.max_cell_len = max_cell_len\n        \n        self.use_metanotebooks = use_metanotebooks\n        self.use_rnd_source_cuts = use_rnd_source_cuts\n        \n        assert (self.use_metanotebooks == False) or (max_cell_len is not None)\n        \n        if i_fold is None:\n            self.df = df\n            \n        else:\n            if self.training:\n                self.df = df[(df.folds != i_fold) & (df.folds >= 0)]\n            else:\n                self.df = df[(df.folds == i_fold)]\n        \n        \n        self.all_nb_id_v = np.unique( self.df.index.map(lambda i: i[0]) )\n        \n        return None\n    \n    \n    def __getitem__(self, idx):\n        \n        nb_id = self.all_nb_id_v[idx]\n        nb_df = self.df.loc[nb_id]\n        nb_df['sample_id'] = nb_df.index.map(lambda x: nb_id + '_' + x).values\n        \n        if self.training and (self.use_metanotebooks or (self.max_cell_len is not None)):\n            nb_df = nb_df.sort_values('rank')\n            \n            if self.use_metanotebooks:\n                used_idx_v = [idx]\n                while len(nb_df) < self.max_cell_len:\n                    idx_merge = idx\n                    while (idx_merge in used_idx_v) or (len(nb_df_merge) < 5):\n                        idx_merge = random.randint(0, self.all_nb_id_v.shape[0] - 1)\n                        \n                        nb_id_merge = self.all_nb_id_v[idx_merge]\n                        nb_df_merge = self.df.loc[nb_id_merge]\n                    \n                    used_idx_v.append(idx_merge)\n                    nb_df_merge['sample_id'] = nb_df_merge.index.map(lambda x: nb_id_merge + '_' + x).values    \n                    nb_df_merge = nb_df_merge.sort_values('rank')\n                    \n#                     print('len(nb_df) =', len(nb_df))\n#                     print('len(nb_df_merge) =', len(nb_df_merge))\n                    \n                    nb_df_merge = shrink_nb(\n                        nb_df_merge,\n                        max_n_cells=self.max_cell_len - len(nb_df) + 4\n                    )\n                    \n                    nb_df = merge_nbs(nb_df, nb_df_merge)\n                    \n                    \n                    \n            if self.max_cell_len is not None:\n                nb_df = shrink_nb(\n                    nb_df,\n                    max_n_cells=self.max_cell_len\n                )\n\n            nb_df = re_rank_nb(nb_df)\n            \n        \n        \n        tkn_len = max(self.cfg.max_md_tkn_len, self.cfg.max_code_tkn_len)\n        text_v = []\n        for cell_type, source, code_rank in nb_df[['cell_type', 'source', 'code_rank']].values:\n            \n            if self.use_rnd_source_cuts:\n                text_a_split_v = source.split(' ')\n                n_words = len(text_a_split_v)\n                n_words_max = tkn_len // 4 # It is an approximation \n                \n                if n_words > n_words_max:\n                    i_s = random.randint(0, n_words - n_words_max - 1)\n                    \n                else:\n                    i_s = 0\n                \n                text_a = ' '.join( text_a_split_v[i_s:i_s + tkn_len] ) \n                \n            else:\n                # Optimization:\n                text_a = ' '.join( source.split(' ')[:tkn_len] ) \n                #text_a = source\n\n            if cell_type == 'code':\n                text_b = f'code cell {code_rank}'\n            else:\n                text_b = f'text cell'\n                \n            text_v.append( [text_a, text_b] )\n    \n    \n        inputs = self.tokenizer.batch_encode_plus(\n            text_v,\n            max_length=tkn_len,\n            padding=\"max_length\",\n            truncation=True,\n            return_offsets_mapping=False,\n        )\n        \n        # Masking some tokens\n        if self.training and self.p_rnd_mask_tkn > 0.0 and (self.tokenizer.mask_token_id is not None): \n            for i in range(len(inputs['input_ids'])):\n                for j in range(len(inputs['input_ids'][i])):\n                    if inputs['input_ids'][i][j] not in self.tokenizer.all_special_ids:\n                        if np.random.random() < self.p_rnd_mask_tkn:\n                            inputs['input_ids'][i][j] = self.tokenizer.mask_token_id\n\n                    elif inputs['input_ids'][i][j] == self.tokenizer.pad_token_id:\n                        break\n\n        for k in inputs.keys():\n            inputs[k] = torch.tensor( inputs[k], dtype=torch.long)\n            \n\n        f_code = (nb_df.cell_type == 'code').values\n        f_md   = ~f_code\n\n        n_code  = f_code.sum()\n        n_cells = f_code.shape[0]\n        n_md    = n_cells - n_code\n\n        cte_features = [\n            min(1.0, n_md/self.cfg.n_md_feat_norm),\n            min(1.0, n_code/self.cfg.n_code_feat_norm),\n            min(1.0, n_cells/(self.cfg.n_md_feat_norm + self.cfg.n_code_feat_norm)),\n        ]\n        \n        # Lineal\n        #sample_w = n_cells /  (self.cfg.n_md_feat_norm + self.cfg.n_code_feat_norm)\n        \n        # Squared\n        sample_w = 10.0 * (n_cells /  (self.cfg.n_md_feat_norm + self.cfg.n_code_feat_norm)) ** 2\n\n        features = np.array(\n            [\n                cte_features + [cr]\n                for cr in np.nan_to_num( nb_df.code_rank_pct.values, nan=self.md_sample_default_rank)\n            ],\n            dtype=np.float32,\n        )\n        \n        inputs['features'] = torch.tensor(features, dtype=torch.float32)\n        \n        \n        \n        data = {\n            'sample_id': nb_df['sample_id'].values,\n            'nb_id':    nb_id,\n            'cells_id': nb_df.index.values,\n            \n            'w': torch.tensor(sample_w),\n\n            'inputs':  inputs,\n            'f_code':  torch.tensor(f_code),\n            'f_md':    torch.tensor(f_md),\n            \n            'n_code':  n_code,\n            'n_md':    n_md,\n            'n_cells': n_cells,\n            'nb_df':   nb_df,\n        }\n        \n        \n\n        if 'rank' in nb_df.columns :\n            rank_v = nb_df['rank'].values\n            \n            if self.global_rank_target:\n                data['target'] = torch.tensor(\n                    rank_v,\n                    dtype=torch.long\n                )\n                \n            else:\n            \n                code_rank_v = rank_v[f_code]\n                md_rank_v   = rank_v[f_md]\n\n                data['target'] = torch.tensor(\n                    [bisect(code_rank_v, md_rank) for md_rank in md_rank_v],\n                    dtype=torch.long\n                )\n            \n        return data\n\n    def __len__(self):\n        return len(self.all_nb_id_v)\n    \n    def get_gt_orders(self):\n        gt_orders = self.df.reset_index().sort_values(['id', 'rank_pct']).groupby('id')['cell_id'].apply(list)\n        return gt_orders\n    \n    def collate_fn(self, data_v):\n        assert len(data_v) == 1 \n        return data_v[0]\n    \n    \n# ds_trn = MatrixTrainDataset(\n#     df=df_nb,\n#     cfg=CFG,\n#     tokenizer=tokenizer,\n#     i_fold=0,\n#     training=True,\n#     md_sample_default_rank=0.5,\n#     p_rnd_mask_tkn=0.1,\n#     global_rank_target=False,\n#     use_metanotebooks=True,\n#     use_rnd_source_cuts=True,\n#     max_cell_len=128,\n# )\n\n# dl_trn = DataLoader(ds_trn, collate_fn=ds_trn.collate_fn, batch_size=1)\n\n# for data in dl_trn:\n#     break","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.760461Z","iopub.execute_input":"2022-08-10T23:34:20.761092Z","iopub.status.idle":"2022-08-10T23:34:20.791209Z","shell.execute_reply.started":"2022-08-10T23:34:20.761057Z","shell.execute_reply":"2022-08-10T23:34:20.790349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"papermill":{"duration":0.031084,"end_time":"2022-02-08T03:59:08.497761","exception":false,"start_time":"2022-02-08T03:59:08.466677","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=3.0, positve_weighting=-1, reduction=\"none\"):\n        super().__init__()\n        \n        assert (0 <= positve_weighting <= 1) or positve_weighting == -1, \"Weighting factor in range (0,1) to balance positive vs negative examples or -1 for ignore.\"\n        \n        self.alpha = positve_weighting\n        self.gamma = gamma\n        \n        self.reduction = reduction\n        \n        return None\n\n    def forward(self, inputs, targets):\n        focal_loss = torchvision.ops.sigmoid_focal_loss(\n            inputs,\n            targets,\n            alpha=self.alpha,\n            gamma=self.gamma,\n            reduction=self.reduction,\n        )\n        \n        return focal_loss\n        \n        \ndef l2_norm(x, axis=1, eps=1e-6):\n    norm = torch.norm(x, 2, axis, True) + eps\n    output = torch.div(x, norm)\n    return output\n\n        \nclass Arcface(nn.Module):\n    # implementation of additive margin softmax loss in https://arxiv.org/abs/1801.05599    \n    def __init__(\n        self, \n        arc_classnum=2,\n        arc_bottleneck=None,\n        embedding_size=32,\n        arc_s=0.5,\n        arc_m=64,\n        reduction='mean',\n        ignore_index=-1,\n    ):\n        super().__init__()\n        \n        self.classnum       = arc_classnum\n        self.arc_bottleneck = arc_bottleneck\n        self.embedding_size = embedding_size\n        self.s              = arc_s  # the margin value, default is 0.5\n        self.m              = arc_m  # scalar value default is 64, see normface https://arxiv.org/abs/1704.06369\n        self.reduction      = reduction\n        self.ignore_index   = ignore_index\n        \n        \n        self.use_bottleneck = (self.arc_bottleneck is not None)\n\n        if self.use_bottleneck:\n            self.kernel_A = nn.Parameter(\n                torch.Tensor(self.embedding_size, self.arc_bottleneck)\n            )\n            self.kernel_A.data.uniform_(-1, 1).renorm_(2,1,1e-5).mul_(1e5)\n            \n            self.kernel_B = nn.Parameter(\n                torch.Tensor(self.arc_bottleneck, self.classnum)\n            )\n            self.kernel_B.data.uniform_(-1, 1).renorm_(2,1,1e-5).mul_(1e5)\n            \n        else:\n            self.kernel = nn.Parameter(\n                torch.Tensor(self.embedding_size, self.classnum)\n            )\n            self.kernel.data.uniform_(-1, 1).renorm_(2,1,1e-5).mul_(1e5)\n\n\n        self.cos_m = np.cos(self.m)\n        self.sin_m = np.sin(self.m)\n        \n        self.mm = (self.sin_m * self.m) # issue 1\n        self.threshold = np.cos(np.pi - self.m)\n        \n        \n        self.criterion = nn.CrossEntropyLoss(\n            weight=None,\n            reduction=self.reduction,\n            ignore_index=self.ignore_index,\n        )\n        \n        return None\n    \n    \n    def project(self, embbedings):\n        # weights norm\n        if self.use_bottleneck:\n            kernel_norm = l2_norm(self.kernel_A @ self.kernel_B, axis=0)\n        else:\n            kernel_norm = l2_norm(self.kernel, axis=0)\n            \n        # cos(theta+m)\n        cos_theta = torch.mm(embbedings, kernel_norm)\n        cos_theta = cos_theta.clamp(-1,1) # for numerical stability\n        output = self.s * cos_theta\n        return output\n        \n        \n    def forward(self, embbedings, label):\n        \n        nB = len(embbedings)\n        \n        if self.use_bottleneck:\n            kernel_norm = l2_norm(self.kernel_A @ self.kernel_B, axis=0)\n            \n        else:\n            kernel_norm = l2_norm(self.kernel, axis=0)\n            \n        # cos(theta+m)\n\n        cos_theta = torch.mm(embbedings, kernel_norm)\n        cos_theta = cos_theta.clamp(-1,1) # for numerical stability\n        \n        cos_theta_2 = torch.pow(cos_theta, 2).type(cos_theta.dtype)\n        sin_theta_2 = 1.0 - cos_theta_2\n        sin_theta   = torch.sqrt(sin_theta_2)\n        cos_theta_m = (cos_theta * self.cos_m - sin_theta * self.sin_m)\n        \n        # this condition controls the theta+m should in range [0, pi]\n        #  0<=theta+m<=pi\n        # -m<=theta<=pi-m\n        \n        cond_v    = cos_theta - self.threshold\n        cond_mask = cond_v <= 0\n        \n        keep_val = (cos_theta - self.mm) # when theta not in [0,pi], use cosface instead\n        \n        cos_theta_m[cond_mask] = keep_val[cond_mask]\n        output = cos_theta * 1.0 # a little bit hacky way to prevent in_place operation on cos_theta\n        \n        \n        f = (label  != self.ignore_index)\n        idx_ = torch.arange(0, nB, dtype=torch.long)[f]\n        label_ = label[f]\n        \n        output[idx_, label_] = cos_theta_m[idx_, label_]\n        output *= self.s # scale up in order to make softmax work, first introduced in normface\n        \n        loss = self.criterion(\n            output,\n            label,\n        )\n        \n        return loss #, output, cos_theta\n    \n    @torch.jit.ignore\n    def get_trainable_weights(self, verbose=True):\n        trainable_params_v = [p for p in self.parameters() if p.requires_grad ]\n        \n        if verbose:\n            n_w = 0\n            for p in trainable_params_v:\n                n_w += np.prod(p.shape)\n\n            print(f' - Arcface: total trainable weights: {n_w/1e6:0.03} M')\n\n            \n        return trainable_params_v","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.794643Z","iopub.execute_input":"2022-08-10T23:34:20.794922Z","iopub.status.idle":"2022-08-10T23:34:20.817652Z","shell.execute_reply.started":"2022-08-10T23:34:20.794898Z","shell.execute_reply":"2022-08-10T23:34:20.816676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Reshape(nn.Module):\n    def __init__(self, *args):\n        super().__init__()\n        self.shape = args\n\n    def forward(self, x):\n        return x.view(self.shape)\n    \n    \nclass L2NormLayer(nn.Module):\n    def __init__(self, axis=1, eps=1e-6):\n        super().__init__()\n        self.axis = axis\n        self.eps = eps\n        return None\n    \n    def forward(self, x):\n        norm = torch.norm(x, 2, self.axis, True) + self.eps\n        output = torch.div(x, norm)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.819106Z","iopub.execute_input":"2022-08-10T23:34:20.819744Z","iopub.status.idle":"2022-08-10T23:34:20.831734Z","shell.execute_reply.started":"2022-08-10T23:34:20.819708Z","shell.execute_reply":"2022-08-10T23:34:20.830827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PositionWeightsMSELoss(nn.Module):\n    def __init__(\n        self,\n        reduction='mean',\n        position_weights_factor=1.0,\n        position_weights_clamp_value=10.0,\n    ):\n        super().__init__()\n        \n        self.reduction = reduction\n        self.position_weights_factor = position_weights_factor\n        self.position_weights_clamp_value = position_weights_clamp_value\n        return None\n\n    def forward(self, pos_mat, target, f_code, f_md):\n        n_code = f_code.sum()\n        n_md = f_md.sum()\n        \n        w = (\n            self.position_weights_factor * torch.abs(\n                torch.arange(\n                    0, n_code + 1,\n                    device=target.device\n                )[:,None] - target\n            ) + 1.0\n        ).T\n        \n        if True:\n            for i in range(n_md):\n                for j in range(n_md):\n                    if i != j:\n                        if target[i] < target[j]:\n                            w[i] += (torch.arange(0, n_code + 1, device=target.device) > target[j]).type(w.dtype)\n\n                        elif target[i] > target[j]:\n                            w[i] += (torch.arange(0, n_code + 1, device=target.device) < target[j]).type(w.dtype)\n        \n        w = torch.clamp(w, 0.0, self.position_weights_clamp_value).type(pos_mat.dtype)\n            \n        target_oh = torch.zeros_like(pos_mat)\n        target_oh[torch.arange(target.shape[0]), target] = 1.0\n        \n        loss = (w * torch.square(pos_mat.softmax(axis=-1) - target_oh)).sum(axis=-1)\n        \n        if self.reduction == 'sum':\n            loss = loss.sum(axis=0)\n            \n        elif self.reduction == 'mean':\n            loss = loss.mean(axis=0)\n            \n        elif self.reduction == 'none':\n            pass\n        \n        else:\n            raise NotImplementedError(f'reduction = {self.reduction} ???')\n        \n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.833455Z","iopub.execute_input":"2022-08-10T23:34:20.834155Z","iopub.status.idle":"2022-08-10T23:34:20.847169Z","shell.execute_reply.started":"2022-08-10T23:34:20.834119Z","shell.execute_reply":"2022-08-10T23:34:20.846038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def l2_norm(x, axis=1, eps=1e-6):\n    norm = torch.norm(x, 2, axis, True) + eps\n    output = torch.div(x, norm)\n    return output\n\n\ndef sinusoid_positional_encoding_ref(length, dimensions):\n    \"\"\" returns pe with pe.shape == (length, dimensions)\"\"\"\n    def get_position_angle_vec(position):\n        return [np.pi * position / np.power(length, 2 * (i // 2) / (dimensions-2))\n                for i in range(dimensions)]\n\n    PE = np.array([get_position_angle_vec(i) for i in range(length)])\n    PE[:, 0::2] = np.sin(PE[:, 0::2])  # dim 2i\n    PE[:, 1::2] = np.cos(PE[:, 1::2])  # dim 2i+1\n    return PE\n\npe = sinusoid_positional_encoding_ref(1000, 128)\n\nclass PositionLoss(nn.Module):\n    def __init__(\n        self,\n        embedding_dim=256,\n        scale=10.0,\n        ignore_index=-1,\n        reduction='mean',\n        eps=1e-6,\n        use_position_weights=False,\n        position_weights_factor=1.0,\n        position_weights_clamp_value=10.0,\n        use_abs_pos_enc=False,\n        interpolate_pe=False,\n        use_pe_projection=False,\n        pe_lenght=10_000,\n        pe_dim=256,\n        use_code_context=False,\n        n_code_context=2,\n        use_max_pooling=False,\n    ):\n        super().__init__()\n        \n        self.embedding_dim = embedding_dim\n        self.s = scale\n        self.ignore_index = ignore_index\n        self.reduction = reduction\n        self.eps = eps\n        \n        self.use_position_weights = use_position_weights\n        self.position_weights_factor = position_weights_factor\n        self.position_weights_clamp_value = position_weights_clamp_value\n        \n        self.use_abs_pos_enc = use_abs_pos_enc\n        self.interpolate_pe = interpolate_pe\n        self.use_pe_projection = use_pe_projection\n        self.pe_lenght = pe_lenght\n        self.pe_dim = pe_dim\n            \n        self.use_code_context = use_code_context\n        self.n_code_context = n_code_context\n        self.use_max_pooling = use_max_pooling\n        \n        self._build_layers()\n        \n        return None\n    \n    \n    def _build_layers(self):\n        # First code cell\n        self.boc = nn.Parameter(\n            torch.Tensor(1, self.embedding_dim)\n        )\n        self.boc.data.uniform_(-1, 1).renorm_(2,0,1e-5).mul_(1e5)\n        \n        # Last code cell\n        self.eoc = nn.Parameter(\n            torch.Tensor(1, self.embedding_dim)\n        )\n        self.eoc.data.uniform_(-1, 1).renorm_(2,0,1e-5).mul_(1e5)\n        \n        if self.use_abs_pos_enc:\n            if self.use_pe_projection:\n                self.abs_pe_cte = torch.tensor(\n                    sinusoid_positional_encoding_ref(\n                        self.pe_lenght,\n                        self.pe_dim,\n                    ),\n                    dtype=torch.float32\n                )\n\n                self.abs_pe_proj = nn.Linear(\n                    self.pe_dim,\n                    self.embedding_dim,\n                    bias=False,\n                )\n            \n            else:  \n                self.abs_pe_cte = nn.Parameter(\n                    torch.Tensor(self.pe_lenght, 1, self.embedding_dim)\n                )\n                self.abs_pe_cte.data.normal_(mean=0.0, std=0.02)\n            \n        \n        if self.use_code_context:\n            self.rel_pe = nn.Parameter(\n                torch.Tensor(2*self.n_code_context, self.embedding_dim)\n            )\n            self.rel_pe.data.normal_(mean=0.0, std=0.02)\n            \n            self.context_layer = nn.Sequential(\n                nn.Linear(\n                    2*self.embedding_dim,\n                    1,\n                    bias=False,\n                ),\n                nn.Sigmoid(),\n            )\n        \n        if self.use_position_weights:\n            self.criterion = PositionWeightsMSELoss(\n                reduction=self.reduction,\n                position_weights_factor=self.position_weights_factor,\n                position_weights_clamp_value=self.position_weights_clamp_value,\n            )\n            \n        else:\n            self.criterion = nn.CrossEntropyLoss(\n                weight=None,\n                reduction=self.reduction,\n                ignore_index=self.ignore_index,\n            )\n        \n        return None\n    \n    \n    def _norm_and_cat(self, y_pred, f_code, f_md):\n        \n        y_pred_ud = y_pred[:,:2]  # N, 2, E\n        \n        boc = l2_norm(self.boc, axis=1, eps=self.eps)\n        eoc = l2_norm(self.eoc, axis=1, eps=self.eps)\n        \n        if self.use_abs_pos_enc:\n            N_code = f_code.sum()\n            pe_cte = self.abs_pe_cte.to(y_pred_ud.device).type(y_pred_ud.dtype)\n            \n            if self.interpolate_pe:\n                pe_idx_v = torch.round(\n                    (self.pe_lenght - 1) * torch.linspace(0.0, 1.0, N_code),\n                    decimals=0,\n                ).type(torch.int64)\n                \n                pe_cte = pe_cte[pe_idx_v]\n                \n            else:\n                if N_code > self.pe_lenght:\n                    pe_cte = torch.cat(\n                        [pe_cte] + (N_code - self.pe_lenght) * [pe_cte[-1:]]\n                    )\n                    \n                else:\n                    pe_cte = pe_cte[:N_code]\n            \n            if self.use_pe_projection:\n                pe = self.abs_pe_proj(pe_cte).reshape( (-1, 1, self.embedding_dim) )\n                \n            else:\n                pe = pe_cte.reshape( (-1, 1, self.embedding_dim) )\n            \n            y_code_ud = y_pred_ud[f_code] + pe\n            \n        else:\n            y_code_ud = y_pred_ud[f_code]\n            \n            \n        y_code_ud_norm = l2_norm(y_code_ud, axis=2, eps=self.eps)\n        \n        y_code_ud_norm = torch.concat(\n            [\n                boc,\n                y_code_ud_norm.reshape( (-1, self.embedding_dim) ),\n                eoc,\n            ],\n            axis=0\n        )\n\n        y_md_ud_norm = l2_norm( y_pred_ud[f_md], axis=2, eps=self.eps )\n        \n        return y_code_ud_norm.reshape( (-1, 2, self.embedding_dim) ), y_md_ud_norm.reshape( (-1, 2, self.embedding_dim) )\n    \n    \n    def build_context_attention_matrix(self, y_pred, f_code, f_md):\n        y_code = y_pred[f_code, 2]\n        y_md   = y_pred[f_md,   2]\n        \n        N_code = y_code.shape[0]\n        N_md   = y_md.shape[0]\n        \n        \n        y_code_pe_v = []\n        for i_code in range(N_code+1):\n            i_s_co = max(0, i_code - self.n_code_context)\n            i_e_co = min(N_code, i_code + self.n_code_context)\n\n            i_s_pe = i_s_co - i_code + self.n_code_context\n            i_e_pe = i_e_co - i_code + self.n_code_context\n\n\n            y_code_pe = l2_norm(y_code[i_s_co:i_e_co] + self.rel_pe[i_s_pe:i_e_pe], axis=1, eps=self.eps).mean(axis=0)\n            y_code_pe_v.append(y_code_pe)\n\n\n        y_md_norm = l2_norm(y_md, axis=1, eps=self.eps)[:,None,:]    \n        y_code_pe = torch.vstack( y_code_pe_v )[None,:,:]\n\n\n        y_prod = y_md_norm * y_code_pe\n        y_sum = 0.5 * (y_md_norm + y_code_pe)\n\n#         y_code_norm_rep = y_code_pe + torch.zeros_like(y_md_norm)\n#         y_md_norm_rep   = torch.zeros_like(y_code_pe) + y_md_norm\n\n        y_cat = torch.cat( [y_prod, y_sum], axis=-1)\n        \n        ctx_mat = self.context_layer(y_cat)[:,:,0]\n        \n        return ctx_mat\n    \n    \n    def build_position_matrix(self, y_pred, f_code, f_md):\n        assert y_pred.shape[2] == self.embedding_dim\n        \n        y_code_ud_norm, y_md_ud_norm = self._norm_and_cat(y_pred, f_code, f_md)\n        \n        pos_mat_u = (y_md_ud_norm[:,0,:] @ y_code_ud_norm[:,0,:].T)\n        pos_mat_d = (y_md_ud_norm[:,1,:] @ y_code_ud_norm[:,1,:].T)\n        \n        #y_code_ud_norm = y_code_ud_norm.reshape( (-1, 2 * self.embedding_dim) )\n        #y_md_ud_norm   = y_md_ud_norm.reshape(   (-1, 2 * self.embedding_dim) )\n        \n        if self.use_max_pooling:\n            pos_mat = self.s * torch.maximum(pos_mat_u, pos_mat_d)\n        else:\n            pos_mat = (0.5 * self.s) * (pos_mat_u + pos_mat_d)\n        \n        if self.use_code_context:\n            ctx_mat = self.build_context_attention_matrix(y_pred, f_code, f_md)\n            pos_mat = ctx_mat * pos_mat\n\n        return pos_mat\n\n    \n    def forward(self, y_pred, target, f_code, f_md):\n        pos_mat = self.build_position_matrix(y_pred, f_code, f_md)\n        \n        if self.use_position_weights:\n            loss = self.criterion(pos_mat, target, f_code, f_md)\n            \n        else:\n            loss = self.criterion(pos_mat, target)\n        \n        return loss\n    \n    \n    def predict_md_pos(self, y_pred, f_code, f_md):\n        y_code_ud_norm, y_md_ud_norm = self._norm_and_cat(y_pred, f_code, f_md)\n        \n        y_code_ud_norm = y_code_ud_norm.reshape( (-1, 2 * self.embedding_dim) )\n        y_md_ud_norm   = y_md_ud_norm.reshape(   (-1, 2 * self.embedding_dim) )\n        \n        N_code = y_code_ud_norm.shape[0]\n        N_md   = y_md_ud_norm.shape[0]\n        \n        m = (y_md_ud_norm[...,None] * y_code_ud_norm.T[None]).reshape(\n            (N_md, 2, self.embedding_dim, N_code)\n        ).sum(2) # m's range = [-1 ,1]\n        \n        \n        m = 0.5 * (m + 1.0)\n        \n        if self.use_code_context:\n            ctx_mat = self.build_context_attention_matrix(y_pred, f_code, f_md)\n            m = m * ctx_mat[:,None,:]\n            \n        m_sum = m.sum(1)\n        \n        if self.use_max_pooling:\n            idx_v = torch.maximum(m[:,0], m[:,1]).argmax(1)\n            \n        else:\n            idx_v = m_sum.argmax(1)\n\n        p_up = m[torch.arange(N_md), 0, idx_v] / m_sum[torch.arange(N_md), idx_v]\n\n        return idx_v, p_up # position in between code cells, prob of being closer to the up-code-cell\n\n\n    def predict_rank_v(self, y_pred, f_code, f_md, return_order_v=False):\n        pred_md_idx_v, md_p_up = self.predict_md_pos(y_pred, f_code, f_md)\n\n\n        sample_idx_md = torch.argwhere( f_md ).T[0]\n        sample_idx_code = torch.argwhere( f_code ).T[0]\n        N_code = sample_idx_code.shape[0]\n\n        order_v = []\n        for i in range(N_code+1):\n            ii = torch.argwhere(pred_md_idx_v == i).T[0]\n            ip = md_p_up[ii]\n\n            if len(ii) > 0:\n                md_rank_v = sample_idx_md[ ii[ torch.argsort(ip, descending=True) ] ].detach().cpu().numpy()\n                order_v.extend( md_rank_v )\n\n            if i < N_code:\n                order_v.append( sample_idx_code[i].item() )\n\n        order_v = np.array( order_v )\n\n        rank_v = np.zeros_like(order_v)\n        rank_v[order_v] = np.arange(order_v.shape[0])\n\n        if return_order_v:\n            return rank_v, order_v\n\n        return rank_v\n    \n# self = PositionLoss(\n#     embedding_dim=256,\n#     use_position_weights=False,\n#     position_weights_factor=1/5000,\n#     use_abs_pos_enc=True\n# )","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:20.848765Z","iopub.execute_input":"2022-08-10T23:34:20.849640Z","iopub.status.idle":"2022-08-10T23:34:21.200871Z","shell.execute_reply.started":"2022-08-10T23:34:20.849603Z","shell.execute_reply":"2022-08-10T23:34:21.199652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ListMLELoss(nn.Module):\n    def __init__(self, eps=1e-8, padded_value_indicator=-1, reduction='mean'):\n        \"\"\"\n        ListMLE loss introduced in \"Listwise Approach to Learning to Rank - Theory and Algorithm\".\n        :param eps: epsilon value, used for numerical stability\n        :param padded_value_indicator: an indicator of the y_true index containing a padded item, e.g. -1\n        \"\"\"\n        super().__init__()\n        \n        self.eps = eps\n        self.padded_value_indicator = padded_value_indicator\n        self.reduction = reduction\n        \n        return None\n    \n    def forward(self, y_pred, y_true):\n        \"\"\"\n        :param y_pred: predictions from the model, shape [batch_size, slate_length]\n        :param y_true: ground truth labels, shape [batch_size, slate_length]\n\n        :return: loss value, a torch.Tensor\n        \"\"\"\n\n\n        # shuffle for randomised tie resolution\n        random_indices = torch.randperm(y_pred.shape[-1])\n        y_pred_shuffled = y_pred[:, random_indices]\n        y_true_shuffled = y_true[:, random_indices]\n\n        y_true_sorted, indices = y_true_shuffled.sort(descending=True, dim=-1)\n\n        mask = y_true_sorted == self.padded_value_indicator\n\n        preds_sorted_by_true = torch.gather(y_pred_shuffled, dim=1, index=indices)\n        preds_sorted_by_true[mask] = float(\"-inf\")\n\n        max_pred_values, _ = preds_sorted_by_true.max(dim=1, keepdim=True)\n\n        preds_sorted_by_true_minus_max = preds_sorted_by_true - max_pred_values\n\n        cumsums = torch.cumsum(preds_sorted_by_true_minus_max.exp().flip(dims=[1]), dim=1).flip(dims=[1])\n\n        observation_loss = torch.log(cumsums + self.eps) - preds_sorted_by_true_minus_max\n\n        observation_loss[mask] = 0.0\n        \n        loss_v = torch.sum(observation_loss, dim=1)\n        if self.reduction == 'mean':\n            loss = torch.mean(loss_v)\n            \n        elif self.reduction == 'sum':\n            loss = torch.sum(loss_v)\n            \n        elif self.reduction == 'none':\n            loss = loss_v\n            \n        else:\n            raise NotImplementedError(f'reduction = {self.reduction} ???')\n                \n        return loss\n","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.202202Z","iopub.execute_input":"2022-08-10T23:34:21.202601Z","iopub.status.idle":"2022-08-10T23:34:21.214090Z","shell.execute_reply.started":"2022-08-10T23:34:21.202558Z","shell.execute_reply":"2022-08-10T23:34:21.213121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SortingHead(nn.Module):\n    def __init__(\n        self,\n        embedding_dim=768,\n        n_hidden_layers=3,\n        pad_value=-1,\n        reduction='mean',\n        use_abs_pos_enc=True,\n        interpolate_pe=True,\n        use_pe_projection=False,\n        pe_lenght=200,\n        pe_dim=128,\n        eps=1e-08,\n    ):\n        super().__init__()\n        \n        self.embedding_dim = embedding_dim\n        self.n_hidden_layers = n_hidden_layers\n        self.pad_value = pad_value\n        self.reduction = reduction\n        \n        \n        self.use_abs_pos_enc = use_abs_pos_enc\n        self.interpolate_pe = interpolate_pe\n        self.use_pe_projection = use_pe_projection\n        self.pe_lenght = pe_lenght\n        self.pe_dim = pe_dim\n        self.eps = eps\n        \n        self._build_model()\n        return None\n    \n    \n    def _build_model(self):\n        self.config = BertConfig(\n            attention_probs_dropout_prob=0.1,\n            hidden_act=\"gelu\",\n            hidden_dropout_prob=0.1,\n            hidden_size=self.embedding_dim,\n            initializer_range=0.02,\n            intermediate_size=4*self.embedding_dim,\n            layer_norm_eps=1e-12,\n            num_attention_heads=max(2, self.embedding_dim//64),\n            num_hidden_layers=self.n_hidden_layers,\n        )\n\n        self.transformer = BertEncoder(\n            self.config\n        )\n        \n        self.output_layer = nn.Linear(\n            in_features=self.embedding_dim,\n            out_features=1,\n            bias=True\n        )\n        \n        if self.use_abs_pos_enc:\n            if self.use_pe_projection:\n                self.abs_pe_cte = torch.tensor(\n                    sinusoid_positional_encoding_ref(\n                        self.pe_lenght,\n                        self.pe_dim,\n                    ),\n                    dtype=torch.float32\n                )\n\n                self.abs_pe_proj = nn.Linear(\n                    self.pe_dim,\n                    self.embedding_dim,\n                    bias=False,\n                )\n            \n            else:  \n                self.abs_pe_cte = nn.Parameter(\n                    torch.Tensor(self.pe_lenght, self.embedding_dim)\n                )\n                self.abs_pe_cte.data.normal_(mean=0.0, std=0.02)\n        \n        \n        self.criterion = ListMLELoss(\n            eps=self.eps,\n            padded_value_indicator=self.pad_value,\n            reduction=self.reduction,\n        )\n            \n        return None\n    \n    \n    def _add_pe(self, embeddings, f_code=None, f_md=None):\n        \n        N_code = f_code.sum()\n        pe_cte = self.abs_pe_cte.to(embeddings.device).type(embeddings.dtype)\n\n        if self.interpolate_pe:\n            pe_idx_v = torch.round(\n                (self.pe_lenght - 1) * torch.linspace(0.0, 1.0, N_code),\n                decimals=0,\n            ).type(torch.int64)\n\n            pe_cte = pe_cte[pe_idx_v]\n\n        else:\n            if N_code > self.pe_lenght:\n                pe_cte = torch.cat(\n                    [pe_cte] + (N_code - self.pe_lenght) * [pe_cte[-1:]]\n                )\n\n            else:\n                pe_cte = pe_cte[:N_code]\n\n        if self.use_pe_projection:\n            pe = self.abs_pe_proj(pe_cte)\n\n        else:\n            pe = pe_cte\n\n        embeddings[f_code] = embeddings[f_code] + pe\n            \n        return embeddings\n    \n    \n    def forward_head(self, embeddings, f_code, f_md):\n        \"\"\"\n        embeddings.shape == (n_cells, n_embedding)\n        \"\"\"\n        \n        assert len(embeddings.shape) == 2\n        \n        if self.use_abs_pos_enc:\n            self._add_pe(embeddings, f_code=f_code, f_md=f_md)\n        \n        y = self.transformer(embeddings[None])['last_hidden_state']\n        y = self.output_layer(y)[:,:,0]\n        \n        # from IPython import embed; embed()\n        \n        if True:\n            # Positive constrain\n            y = torch.sigmoid(y)\n            \n            # Sum of code deltas\n            S = y[:,f_code].sum(axis=1)[:,None] + 1.0\n\n            y = torch.cat(\n                [\n                    torch.cumsum(y[:,f_code], dim=1), # code cells ordering\n                    S * y[:,f_md] # Scaling md cells\n                ],\n                dim=1,\n            )\n        \n        return y\n    \n    \n    def forward_criterion(self, embeddings, target_ranks, f_code, f_md):\n        if len(target_ranks.shape) == 1:\n            target_ranks = target_ranks[None]\n            \n        y_ranks = self.forward_head(embeddings, f_code, f_md)\n        loss = self.criterion(y_ranks, target_ranks)\n        return loss\n        \n    def forward(self, embeddings, target_ranks, f_code, f_md):\n        return self.forward_criterion(embeddings, target_ranks, f_code, f_md)\n        \n        \n    def predict_rank_v(self, embeddings, f_code, f_md, return_order_v=False):\n        \n        outputs = self.forward_head(embeddings, f_code, f_md).detach().cpu().numpy()[0]\n        idx_v = np.argsort(outputs)\n        ranks = np.zeros(idx_v.shape, dtype=np.int64)\n        ranks[idx_v] = np.arange(idx_v.shape[0])\n        \n#         outputs = self.forward_head(embeddings, f_code, f_md).detach().cpu().numpy()        \n#         idx_v = np.argsort(outputs, axis=1)\n#         ranks = np.zeros(idx_v.shape, dtype=np.int64)\n#         ranks[np.arange(idx_v.shape[0]), idx_v] = np.arange(idx_v.shape[1])\n        \n        if return_order_v:\n            return ranks, idx_v\n        else:\n            return ranks\n        ","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.215640Z","iopub.execute_input":"2022-08-10T23:34:21.216186Z","iopub.status.idle":"2022-08-10T23:34:21.240327Z","shell.execute_reply.started":"2022-08-10T23:34:21.216150Z","shell.execute_reply":"2022-08-10T23:34:21.239318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class NB_Model(nn.Module):\n    def __init__(self, cfg, config_path=None, pretrained=True):\n        super().__init__()\n        self.cfg = cfg\n        self.config_path = config_path\n        self.pretrained = pretrained\n        \n        self._build_model()\n        self._build_optimizer()\n        return None\n    \n    \n    def _build_optimizer(self):\n        param_optimizer = list( self.named_parameters() )\n\n        no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n        \n        encoder_v = []\n        decoder_v = []\n        for n, p in self.named_parameters():\n            if 'model' in n:\n                encoder_v.append( (n,p) )\n            else:\n                decoder_v.append( (n,p) )\n                \n                \n        optimizer_parameters = [\n            {\n                'params': [p for n, p in encoder_v if not any(nd in n for nd in no_decay)],\n                'lr': self.cfg.encoder_lr,\n                'weight_decay': self.cfg.weight_decay\n            },\n            {\n                'params': [p for n, p in encoder_v if any(nd in n for nd in no_decay)],\n                'lr': self.cfg.encoder_lr,\n                'weight_decay': 0.0\n            },\n            \n            {\n                'params': [p for n, p in decoder_v if not any(nd in n for nd in no_decay)],\n                'lr': self.cfg.decoder_lr,\n                'weight_decay': self.cfg.weight_decay\n            },\n            {\n                'params': [p for n, p in decoder_v if any(nd in n for nd in no_decay)],\n                'lr': self.cfg.decoder_lr,\n                'weight_decay': 0.0\n            }\n        ]\n\n        self.optimizer = AdamW(\n            optimizer_parameters, \n            lr=self.cfg.encoder_lr,\n            eps=self.cfg.eps, \n            betas=self.cfg.betas\n        )\n\n        return None\n    \n    \n    def _build_model(self):\n        \n        self.take_last_token_from_last_hidden_state = 'bloom' in self.cfg.model.lower()\n        \n        if self.config_path is None:\n            self.config = AutoConfig.from_pretrained(\n                self.cfg.model,\n                output_hidden_states=self.cfg.sum_hidden_states,\n            )\n            \n        else:\n            self.config = torch.load(\n                self.config_path\n            )\n            \n        if self.pretrained:\n            self.model = AutoModel.from_pretrained(\n                self.cfg.model,\n                config=self.config\n            )\n            \n        else:\n            self.model = AutoModel.from_config(\n                self.config\n            )\n            \n        if self.cfg.code_use_abs_pos_enc:\n            if self.cfg.code_use_pe_projection:\n                self.abs_pe_cte = torch.tensor(\n                    sinusoid_positional_encoding_ref(\n                        self.cfg.code_pe_lenght,\n                        self.cfg.code_pe_dim,\n                    ),\n                    dtype=torch.float32\n                )\n\n                self.abs_pe_proj = nn.Linear(\n                    self.cfg.code_pe_dim,\n                    self.config.hidden_size,\n                    bias=False,\n                )\n            \n            else:  \n                self.abs_pe_cte = nn.Parameter(\n                    torch.Tensor(self.cfg.code_pe_lenght, self.config.hidden_size)\n                )\n                self.abs_pe_cte.data.normal_(mean=0.0, std=0.02)\n\n        if self.cfg.sum_hidden_states:\n            self.cfg.sum_hidden_states_n = min(self.cfg.sum_hidden_states_n, self.config.num_hidden_layers + 1)\n            self.hs_alphas = nn.Parameter(\n                torch.Tensor( self.cfg.sum_hidden_states_n )\n            )\n            self.hs_alphas.data.normal_(mean=0.0, std=self.config.initializer_range)\n    \n        if self.cfg.fc_dropout > 0.0:\n            self.dropout = nn.Dropout(\n                self.cfg.fc_dropout\n            )\n            \n        else:\n            self.dropout = nn.Identity()\n            \n        \n        input_dim = self.config.hidden_size + (4 if self.cfg.model_type == 'abs' else 0)\n        \n        if self.cfg.use_batch_head:\n            head_config = BertConfig(\n                attention_probs_dropout_prob=0.1,\n                hidden_act=\"gelu\",\n                hidden_dropout_prob=0.1,\n                hidden_size=self.config.hidden_size,\n                initializer_range=0.02,\n                intermediate_size=4*self.config.hidden_size,\n                layer_norm_eps=1e-12,\n                num_attention_heads=max(2, self.config.hidden_size//64),\n                num_hidden_layers=self.cfg.batch_head_nhlayers,\n            )\n\n            self.batch_head = BertEncoder(\n                head_config\n            )\n            \n            \n        if self.cfg.output_l2_embedding:\n            self.fc = nn.Sequential(\n                nn.Linear(\n                    input_dim,\n                    3*self.cfg.output_dim,\n                    False,\n                ),\n                Reshape(-1, 3, self.cfg.output_dim),\n#                 L2NormLayer(\n#                     axis=-1,\n#                     eps=self.cfg.eps,\n#                 ),\n            )\n            \n        else:\n            if input_dim != self.cfg.output_dim:\n                self.fc = nn.Linear(\n                    input_dim,\n                    self.cfg.output_dim,\n                )\n\n                self._init_weights(\n                    self.fc\n                )\n                \n            else:\n                self.fc = nn.Identity()\n                \n            \n            \n        # Criteria \n        if self.cfg.criterion == 'MSELoss':\n            self.criterion = nn.MSELoss(\n                reduction=('none' if self.cfg.use_sample_weights else 'mean'),\n            )\n            \n        elif self.cfg.criterion == 'BCEWithLogitsLoss':\n            self.criterion = nn.BCEWithLogitsLoss(\n                reduction=('none' if self.cfg.use_sample_weights else 'mean'),\n            )\n\n        elif self.cfg.criterion == 'CrossEntropyLoss':\n            self.criterion = nn.CrossEntropyLoss(\n                reduction=('none' if self.cfg.use_sample_weights else 'mean'),\n            )\n            \n        elif self.cfg.criterion == 'FocalLoss':\n            self.criterion = FocalLoss(\n                gamma=self.cfg.fl_gamma,\n                positve_weighting=self.cfg.fl_positve_weighting,\n                reduction=('none' if self.cfg.use_sample_weights else 'mean'),\n            )\n        \n        elif self.cfg.criterion == 'PositionLoss':\n            self.criterion = PositionLoss(\n                embedding_dim=self.cfg.output_dim,\n                scale=self.cfg.ploss_scale,\n                reduction='mean',\n                use_position_weights=self.cfg.ploss_use_position_weights,\n                position_weights_factor=self.cfg.ploss_position_weights_factor,\n                position_weights_clamp_value=self.cfg.ploss_position_weights_clamp_value,\n                use_abs_pos_enc=self.cfg.ploss_use_abs_pos_enc,\n                interpolate_pe=self.cfg.ploss_interpolate_pe,\n                use_pe_projection=self.cfg.ploss_use_pe_projection,\n                pe_lenght=self.cfg.ploss_pe_lenght,\n                pe_dim=self.cfg.ploss_pe_dim,\n                \n                use_code_context=self.cfg.ploss_use_code_context,\n                n_code_context=self.cfg.ploss_n_code_context,\n                use_max_pooling=self.cfg.ploss_use_max_pooling,\n            )\n            \n        elif self.cfg.criterion == 'ListMLELoss':\n            self.criterion = SortingHead(\n                embedding_dim=self.cfg.output_dim,\n                n_hidden_layers=self.cfg.sort_n_hidden_layers,\n                pad_value=-1,\n                reduction='mean',\n                use_abs_pos_enc=self.cfg.sort_use_abs_pos_enc,\n                interpolate_pe=self.cfg.sort_interpolate_pe,\n                use_pe_projection=self.cfg.sort_use_pe_projection,\n                pe_lenght=self.cfg.sort_pe_lenght,\n                pe_dim=self.cfg.sort_pe_dim,\n            )\n            \n        else:\n            raise NotImplementedError(f'CFG.criterion = \"{CFG.criterion}\"')\n        \n        return None\n    \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                \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                \n        elif isinstance(module, nn.LayerNorm):\n            module.bias.data.zero_()\n            module.weight.data.fill_(1.0)\n            \n        return None\n    \n    \n    def _add_pe(self, embeddings, f_code, f_md):\n        N_code = f_code.sum()\n        pe_cte = self.abs_pe_cte.to(embeddings.device).type(embeddings.dtype)\n\n        if self.cfg.code_interpolate_pe:\n            pe_idx_v = torch.round(\n                (self.cfg.code_pe_lenght - 1) * torch.linspace(0.0, 1.0, N_code),\n                decimals=0,\n            ).type(torch.int64)\n\n            pe_cte = pe_cte[pe_idx_v]\n\n        else:\n            if N_code > self.cfg.code_pe_lenght:\n                pe_cte = torch.cat(\n                    [pe_cte] + (N_code - self.cfg.code_pe_lenght) * [pe_cte[-1:]]\n                )\n\n            else:\n                pe_cte = pe_cte[:N_code]\n\n        if self.cfg.code_use_pe_projection:\n            pe = self.cfg.code_abs_pe_proj(pe_cte)\n\n        else:\n            pe = pe_cte\n\n        embeddings[f_code] = embeddings[f_code] + pe\n            \n        return embeddings\n    \n        \n    def forward(self, inputs, return_features=False, f_code=None, f_md=None):\n        if 'features' in inputs.keys():\n            scalar_features = inputs['features']\n            \n        outputs = self.model( **{k:v for k,v in inputs.items() if k != 'features'} )\n        \n        if self.cfg.use_batch_head:\n            if self.take_last_token_from_last_hidden_state:\n                cells_last_hidden_state = outputs.last_hidden_state[:, -1, :]\n                \n            else:\n                cells_last_hidden_state = outputs.last_hidden_state[:, 0, :]\n        \n            if self.cfg.code_use_abs_pos_enc:\n                self._add_pe(cells_last_hidden_state, f_code=f_code, f_md=f_md)\n            \n            outputs = self.batch_head(cells_last_hidden_state[None])\n            \n            # Residual Connection\n            outputs['last_hidden_state'] = outputs['last_hidden_state'] + cells_last_hidden_state[None]\n            \n        \n        if self.cfg.sum_hidden_states:\n            if self.cfg.use_batch_head:\n                hidden_states = torch.stack([hs[0,:,:] for hs in outputs['hidden_states'][-self.cfg.sum_hidden_states_n:] ], dim=0)\n                \n            else:\n                hidden_states = torch.stack([hs[:,0,:] for hs in outputs['hidden_states'][-self.cfg.sum_hidden_states_n:] ], dim=0)\n            \n            features = (hidden_states * self.hs_alphas.softmax(axis=-1)[:, None, None]).sum(0)\n            \n        else:\n            if self.cfg.use_batch_head:\n                features = outputs.last_hidden_state[0, :, :]\n                \n            else:\n                features = outputs.last_hidden_state[:, 0, :]\n        \n        if self.cfg.model_type == 'abs':\n            features = torch.cat([features, scalar_features], axis=1)\n        \n        y_do = self.dropout(features)\n        output = self.fc(y_do)\n        \n        if len(output.shape) == 2 and output.shape[-1] == 1:\n            output = output[:,0]\n            \n        if return_features:\n            return output, features\n        \n        else:\n            return output\n\n        \n    def restore_model(\n        self,\n        checkpoint_path='./microsoft-deberta-base_fold0_best.pth',\n        restore_optimizer=False,\n        scheduler=None,\n    ):\n        state_dict = torch.load(\n            checkpoint_path,\n            map_location='cpu'\n        )\n        \n        model_sd = self.state_dict()\n        \n        for k in model_sd.keys():\n            if k in state_dict['model'].keys():\n                model_sd[k] = state_dict['model'][k]\n                \n            else:\n                print(f'WARNING: key not found in ckpt: \"{k}\"', file=sys.stderr)\n            \n\n        self.load_state_dict(\n            model_sd\n        )\n        \n        if restore_optimizer:\n            if 'optimizer' in state_dict.keys():\n                self.optimizer.load_state_dict(\n                    state_dict['optimizer']\n                )\n                \n            else:\n                print(f'WARNING: optimizer weights not found in ckpt.', file=sys.stderr)\n        \n        \n        if scheduler is not None:\n            if 'scheduler' in state_dict.keys():\n            \n                scheduler.load_state_dict(\n                    state_dict['scheduler']\n                )\n                \n            else:\n                print(f'WARNING: scheduler weights not found in ckpt.', file=sys.stderr)\n        \n        restored_epoch = state_dict.get('epoch', None)\n        \n        for k in list(state_dict.keys()):\n            del(state_dict[k])\n\n        del(state_dict)\n        del(model_sd)\n        gc.collect()\n        \n        return restored_epoch\n\n    \n# model = NB_Model(CFG, config_path=None, pretrained=True)","metadata":{"papermill":{"duration":0.05664,"end_time":"2022-02-08T03:59:08.586278","exception":false,"start_time":"2022-02-08T03:59:08.529638","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:34:21.242563Z","iopub.execute_input":"2022-08-10T23:34:21.243533Z","iopub.status.idle":"2022-08-10T23:34:21.291747Z","shell.execute_reply.started":"2022-08-10T23:34:21.243494Z","shell.execute_reply":"2022-08-10T23:34:21.290830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model testing","metadata":{}},{"cell_type":"code","source":"# model = NB_Model(CFG, config_path=None, pretrained=True)\n\n# data = ds_trn[3]\n\n# device = 'cuda:0'\n\n# _=model.to(device)\n\n# for i in range(100):\n#     data = ds_trn[i%2]\n    \n#     inputs = data['inputs'].to(device)\n#     target = data['target'].to(device)\n\n#     f_code = data['f_code'].to(device)\n#     f_md   = data['f_md'].to(device)\n    \n    \n#     y_pred = model(inputs)\n#     loss = model.criterion(y_pred, target, f_code, f_md)\n#     loss.backward()\n#     print(loss.item())\n#     model.optimizer.step()\n#     model.optimizer.zero_grad()\n\n\n\n# pos_mat = model.criterion.build_position_matrix(\n#     y_pred,\n#     f_code,\n#     f_md,\n# ).softmax(axis=-1).detach().cpu().numpy()\n\n# plt.imshow( pos_mat )\n\n# preds_d = {\n#     'sample_ids_v': [],\n#     'preds_v': [],\n# }\n\n# for data in [ds_trn[0], ds_trn[1]]:\n    \n#     inputs = data['inputs'].to(device)\n#     target = data['target'].to(device)\n\n#     f_code = data['f_code'].to(device)\n#     f_md   = data['f_md'].to(device)\n    \n    \n#     y_pred = model(inputs)\n    \n#     rank_v = model.criterion.pred_rank(\n#         y_pred,\n#         None,\n#         f_code,\n#         f_md,\n#         return_order_v=False,\n#     )\n    \n#     preds_d['sample_ids_v'].extend(data['sample_id'])\n#     preds_d['preds_v'].extend(rank_v)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.293516Z","iopub.execute_input":"2022-08-10T23:34:21.294197Z","iopub.status.idle":"2022-08-10T23:34:21.305153Z","shell.execute_reply.started":"2022-08-10T23:34:21.294159Z","shell.execute_reply":"2022-08-10T23:34:21.304143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helpler functions","metadata":{"papermill":{"duration":0.030307,"end_time":"2022-02-08T03:59:08.65031","exception":false,"start_time":"2022-02-08T03:59:08.620003","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Ensembling functions\n\ndef merge_ckpts(\n    ckpt_paths_v,\n    input_key='state_dict',\n    k_replace_v=None,\n    output_path=None,\n    output_key='state_dict',\n    return_statedict=False,\n):\n    \"\"\"\n    k_replace_v=['model.roberta', 'model' ]\n    \"\"\"\n    \n                 \n    ret_sd = OrderedDict()\n\n    for i_c, ckpt_path in enumerate(tqdm(ckpt_paths_v)):\n        ckpt_sd = torch.load(\n            ckpt_path,\n            map_location='cpu'\n        )\n        \n        if 'optimizer_states' in ckpt_sd.keys():\n            del ckpt_sd['optimizer_states']\n        \n        gc.collect()\n        \n        for k in ckpt_sd[input_key].keys():\n            if (k_replace_v is None) or (k_replace_v[0] in k):\n                \n                if k_replace_v is not None:\n                    new_k = k.replace(*k_replace_v)\n                else:\n                    new_k = k\n\n                if i_c == 0:\n                    ret_sd[new_k] = ckpt_sd[input_key][k]\n\n                else:\n                    ret_sd[new_k] = ret_sd[new_k] + ckpt_sd[input_key][k]\n                    \n                    \n        del(ckpt_sd[input_key])\n        gc.collect()\n        \n\n    del(ckpt_sd)\n    gc.collect()\n    \n    for k in ret_sd.keys():\n        ret_sd[k] = ret_sd[k] / len(ckpt_paths_v)\n    \n    \n    if output_path is not None:\n        print(f' - Saving: {output_path} ...')\n        torch.save(\n            {\n                output_key: ret_sd\n            },\n            output_path,\n        )\n        print(f'Done.')\n        \n    if return_statedict:\n        return ret_sd\n\n    for k in list( ret_sd.keys() ):\n        del ret_sd[k]\n        \n    del(ret_sd)\n    gc.collect()\n    \n    return None\n\n\ndef globa_rank2mk_rank(preds_v, f_code, f_md, output_down_probs=True):\n    assert preds_v.shape[0] == f_code.shape[0] == f_md.shape[0]\n    \n    md_rank_m = np.zeros( (f_md.shape[0], 2)) #rank and sub_rank\n    \n    sorted_cell_idx_v = np.argsort( preds_v )\n    \n    last_rank = np.zeros(2, dtype=np.int64)\n    for sorted_cell_idx in sorted_cell_idx_v:\n        is_code_cell = f_code[sorted_cell_idx]\n\n        if is_code_cell:\n            last_rank[0] = last_rank[0] + 1\n            last_rank[1] = 0\n\n        else:\n            md_rank_m[sorted_cell_idx] = last_rank\n            last_rank[1] += 1\n    \n    md_rank_m = md_rank_m[f_md]\n    \n    if output_down_probs:\n        for u, c in zip( *np.unique( md_rank_m[:,0], return_counts=True) ):\n            f = md_rank_m[:,0] == u\n            if c == 1:\n                md_rank_m[f,1] = 0.5\n            else:\n                md_rank_m[f,1] = md_rank_m[f,1] / (c - 1) \n            \n    return md_rank_m\n\n\n\ndef mk_rank2globa_rank(md_rank_m, f_code, f_md):\n    assert f_code.shape[0] == f_md.shape[0]\n    assert md_rank_m.shape[0] == f_md.sum()\n    \n    pred_md_idx_v, md_p_down = md_rank_m.T\n    pred_md_idx_v = np.round(pred_md_idx_v).astype(np.int64)\n    \n    sample_idx_md = np.argwhere( f_md ).T[0]\n    sample_idx_code = np.argwhere( f_code ).T[0]\n    N_code = sample_idx_code.shape[0]\n\n    order_v = []\n    for i in range(N_code+1):\n        ii = np.argwhere(pred_md_idx_v == i).T[0]\n        ip = md_p_down[ii]\n\n        if len(ii) > 0:\n            md_rank_v = sample_idx_md[ ii[ np.argsort(ip) ] ]\n            order_v.extend( md_rank_v )\n\n        if i < N_code:\n            order_v.append( sample_idx_code[i] )\n\n    order_v = np.array( order_v )\n\n    rank_v = np.zeros_like(order_v)\n    rank_v[order_v] = np.arange(order_v.shape[0])\n\n    return rank_v\n    \ndef ensemble_preds(pred_d_v, ds, w_v=None):\n    n_preds = len(pred_d_v)\n    if w_v == None:\n        w_v = np.ones(n_preds)\n    else:\n        w_v = np.array(w_v)\n        \n    w_v = w_v / w_v.sum()\n    \n    assert all([(pred_d_v[0]['sample_ids_v'] == pred_d_v[i]['sample_ids_v']).all() for i in range(1, n_preds)] )\n    all_sample_ids_v = pred_d_v[0]['sample_ids_v']\n\n    ens_preds_d = {'sample_ids_v':[], 'preds_v':[]}\n\n    last_nb_id = all_sample_ids_v[0].split('_')[0]\n    i_s_nb = 0\n    i_e_nb = 0\n    for i_sample, s_id in enumerate(tqdm(all_sample_ids_v, desc='Building Ensemble ...')):\n        nb_id, cell_id = s_id.split('_')\n        if (nb_id != last_nb_id) or i_sample+1 == len(all_sample_ids_v):\n            i_e_nb = i_sample if i_sample+1 < len(all_sample_ids_v) else i_sample+1\n\n            sample_ids_v = all_sample_ids_v[i_s_nb:i_e_nb]\n            preds_v_v = [pred_d_v[i]['preds_v'][i_s_nb:i_e_nb] for i in range(n_preds)]\n\n            cell_id_v = np.array( [s_id.split('_')[1] for s_id in sample_ids_v] )\n\n            nb_df = ds.df.loc[last_nb_id]\n            assert (nb_df.index == cell_id_v).all()\n\n            f_code = (nb_df.cell_type == 'code').values\n            f_md   = ~f_code\n\n            ens_md_rank_m = np.sum(\n                [\n                    w * globa_rank2mk_rank(preds_v, f_code, f_md)\n                    for w, preds_v in zip(w_v, preds_v_v)\n                ],\n                axis=0\n            )\n\n            ens_preds_v = mk_rank2globa_rank(ens_md_rank_m, f_code, f_md)\n\n            ens_preds_d['sample_ids_v'].append(sample_ids_v)\n            ens_preds_d['preds_v'].append(ens_preds_v)\n\n            if 0:\n                preds_v = preds_v_v[0]\n                md_rank_m = globa_rank2mk_rank(preds_v, f_code, f_md)\n                preds_v_2 = mk_rank2globa_rank(md_rank_m, f_code, f_md)\n                assert (preds_v_2 == preds_v).all()\n\n    #         break\n\n            i_s_nb = i_e_nb\n            last_nb_id = nb_id\n    \n    ens_preds_d['sample_ids_v'] = np.concatenate( ens_preds_d['sample_ids_v'], axis=0)\n    ens_preds_d['preds_v'] = np.concatenate( ens_preds_d['preds_v'], axis=0)\n    \n    return ens_preds_d","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.307195Z","iopub.execute_input":"2022-08-10T23:34:21.308002Z","iopub.status.idle":"2022-08-10T23:34:21.335801Z","shell.execute_reply.started":"2022-08-10T23:34:21.307965Z","shell.execute_reply":"2022-08-10T23:34:21.334805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def print_params(model, return_counts=False):\n    n_tot = 0\n    n_tot_opt = 0\n    for i, (n, p) in enumerate(model.named_parameters()):\n        print(f' {i:4d} {\"G\" if p.requires_grad else \"-\"} {n:60s}  {p.numel():8d}   {str(tuple(p.shape)).replace(\" \", \"\")}')\n        n_tot += p.numel()\n        \n        if p.requires_grad:\n            n_tot_opt += p.numel()\n    \n    def print_params(n_tot, s='Total params: '):\n        if n_tot < 1_000:\n            print(f'{s}{n_tot}')\n\n        elif n_tot < 1_000_000:\n            print(f'{s}{n_tot/1_000:0.02f} k')\n\n        else:\n            print(f'{s}{n_tot/1_000_000:0.02f} M')\n        \n        return None\n    \n    \n    print_params(n_tot,     s='Total params:      ')\n    print_params(n_tot_opt, s='Total opt. params: ')\n    \n    if return_counts:\n        return n_tot, n_tot_opt\n\n    return None","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.337369Z","iopub.execute_input":"2022-08-10T23:34:21.338040Z","iopub.status.idle":"2022-08-10T23:34:21.350680Z","shell.execute_reply.started":"2022-08-10T23:34:21.338005Z","shell.execute_reply":"2022-08-10T23:34:21.349650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_datasets(\n    df,\n    nb_feat_d,\n    i_fold,\n    cfg,\n    tokenizer,\n    ):\n    \n    if cfg.model_type == 'abs':\n        ds_trn = AbsTrainDataset(df, nb_feat_d, cfg, tokenizer, i_fold=i_fold, training=True,  only_md_samples=not cfg.abs_train_md_and_code_cells, md_sample_default_rank=cfg.abs_md_sample_default_rank)\n        ds_val = AbsTrainDataset(df, nb_feat_d, cfg, tokenizer, i_fold=i_fold, training=False, only_md_samples=not cfg.abs_train_md_and_code_cells, md_sample_default_rank=cfg.abs_md_sample_default_rank)\n        ds_tst = AbsTrainDataset(df, nb_feat_d, cfg, tokenizer, i_fold=-1,     training=False, only_md_samples=not cfg.abs_train_md_and_code_cells, md_sample_default_rank=cfg.abs_md_sample_default_rank)\n        \n    elif cfg.model_type == 'order':\n        ds_trn = OrderTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=True , p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, only_md_cells=cfg.order_train_only_md_cells)\n        ds_val = OrderTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=False, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, only_md_cells=cfg.order_train_only_md_cells)\n        ds_tst = OrderTrainDataset(df, cfg, tokenizer, i_fold=-1,     training=False, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, only_md_cells=cfg.order_train_only_md_cells)\n    \n    elif cfg.model_type == 'matrix':\n        ds_trn = MatrixTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=True , md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn,\n                                   use_metanotebooks=cfg.mtx_use_metanotebooks, use_rnd_source_cuts=cfg.mtx_use_rnd_source_cuts, max_cell_len=cfg.batch_size)\n        \n        ds_val = MatrixTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=False, md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn,\n                                   use_metanotebooks=cfg.mtx_use_metanotebooks, use_rnd_source_cuts=False, max_cell_len=cfg.batch_size)\n        \n        ds_tst = MatrixTrainDataset(df, cfg, tokenizer, i_fold=-1,     training=False, md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn,\n                                    use_metanotebooks=cfg.mtx_use_metanotebooks, use_rnd_source_cuts=False, max_cell_len=cfg.batch_size)\n        \n    elif cfg.model_type == 'sorting':\n        ds_trn = MatrixTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=True , md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, global_rank_target=True)\n        ds_val = MatrixTrainDataset(df, cfg, tokenizer, i_fold=i_fold, training=False, md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, global_rank_target=True)\n        ds_tst = MatrixTrainDataset(df, cfg, tokenizer, i_fold=-1,     training=False, md_sample_default_rank=cfg.abs_md_sample_default_rank, p_rnd_mask_tkn=cfg.p_rnd_mask_tkn, global_rank_target=True)\n        \n    else:\n        raise  NotImplementedError(f'model_type = {cfg.model_type}')\n        \n    \n    return ds_trn, ds_val, ds_tst\n\n\ndef get_dataloaders(\n    df,\n    nb_feat_d,\n    i_fold,\n    cfg,\n    tokenizer,\n    val_bs_mult=2,\n    ):\n    \n\n        \n    ds_trn, ds_val, ds_tst = get_datasets(\n        df,\n        nb_feat_d,\n        i_fold,\n        cfg,\n        tokenizer,\n    )\n    \n    if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n        batch_size_trn = batch_size_val = batch_size_tst = 1\n        \n    else:\n        batch_size_trn = cfg.batch_size\n        batch_size_val = cfg.batch_size * val_bs_mult\n        batch_size_tst = cfg.batch_size * val_bs_mult\n        \n    \n        \n    dl_trn = DataLoader(\n        ds_trn,\n        batch_size=batch_size_trn,\n        shuffle=True,\n        num_workers=cfg.num_workers,\n        pin_memory=cfg.pin_memory,\n        drop_last=True,\n        collate_fn=ds_trn.collate_fn,\n    )\n    \n    dl_val = DataLoader(\n        ds_val,\n        batch_size=batch_size_val,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=cfg.pin_memory,\n        drop_last=False,\n        collate_fn=ds_val.collate_fn,\n    )\n    \n    \n    dl_tst = DataLoader(\n        ds_tst,\n        batch_size=batch_size_tst,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=cfg.pin_memory,\n        drop_last=False,\n        collate_fn=ds_val.collate_fn,\n    )\n    \n    return ds_trn, ds_val, ds_tst, dl_trn, dl_val, dl_tst\n\n\n\ndef get_scheduler(cfg, optimizer, num_train_steps):\n    if cfg.scheduler_type == 'linear':\n        scheduler = get_linear_schedule_with_warmup(\n            optimizer,\n            num_warmup_steps=cfg.num_warmup_steps,\n            num_training_steps=num_train_steps,\n        )\n        \n    elif cfg.scheduler_type == 'cosine':\n        scheduler = get_cosine_schedule_with_warmup(\n            optimizer,\n            num_warmup_steps=cfg.num_warmup_steps,\n            num_training_steps=num_train_steps,\n            num_cycles=cfg.num_cycles,\n        )\n        \n    else:\n        raise NotImplementedError(f'Scheduler not implementend: {cfg.scheduler_type}')\n\n    return scheduler\n\n\ndef get_time_str(end=' '):\n    time_str = time.strftime('%H:%M:%S', time.localtime( time.time() ))\n    return f\"[{time_str}]\" + end","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.352325Z","iopub.execute_input":"2022-08-10T23:34:21.352791Z","iopub.status.idle":"2022-08-10T23:34:21.375596Z","shell.execute_reply.started":"2022-08-10T23:34:21.352758Z","shell.execute_reply":"2022-08-10T23:34:21.374612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class K_CkptSaver:\n    def __init__(self, k=5, max_score=True):\n        assert k > 0\n        self.k = k\n        self.max_score = max_score\n        self.history_v = []\n        \n        return None\n    \n    def is_better(self, score):\n        if len(self.history_v) < self.k:\n            return True\n        \n        if self.max_score:\n            return any( [np.isnan(l[0]) or score > l[0] for l in self.history_v] )\n        else:\n            return any( [np.isnan(l[0]) or score < l[0] for l in self.history_v] )\n        \n        \n    def is_the_best(self, score):\n        if len(self.history_v) == 0:\n            return True\n        \n        scores_v = [l[0] for l in self.history_v]\n        \n        if self.max_score:\n            return np.max( scores_v ) < score\n        else:\n            return np.min( scores_v ) > score\n        \n    \n    def get_best_score(self):\n        scores_v = [l[0] for l in self.history_v]\n        \n        if self.max_score:\n            return np.max( scores_v )\n        \n        else:\n            return np.min( scores_v )\n    \n    def save(self, score, data, path):\n        i_d = 0\n        new_path = path\n        while new_path in [l[1] for l in self.history_v]:\n            new_path = path+f'({i_d})'\n            i_d += 1\n            \n        path = new_path\n        \n        torch.save(\n            data,\n            path,\n        )\n        \n        self.history_v.append(\n            (score, path)\n        )\n        \n        self.history_v.sort(\n            reverse=self.max_score\n        )\n        \n        if len(self.history_v) > self.k:\n            to_del_score, to_del_path = self.history_v.pop()\n\n            if to_del_path is not None and os.path.exists(to_del_path):\n                os.remove(to_del_path)\n        \n        return None\n    \n# saver = K_CkptSaver(max_score=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.381925Z","iopub.execute_input":"2022-08-10T23:34:21.382768Z","iopub.status.idle":"2022-08-10T23:34:21.395600Z","shell.execute_reply.started":"2022-08-10T23:34:21.382732Z","shell.execute_reply":"2022-08-10T23:34:21.394636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self, n_sma=500, plt_ylim_v=(0.0, 10.0)):\n        self.n_sma = n_sma\n        self.plt_ylim_v = plt_ylim_v\n                \n        self.reset()\n        return None\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n        \n        self.val_v = []\n        self.avg_v = []\n        return None\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        self.val_v.append(self.val)\n        self.avg_v.append(self.avg)\n        \n        return None\n    \n    def moving_average(self, x):\n        x_sma = np.convolve(\n            x,\n            np.ones(self.n_sma)/self.n_sma,\n            'valid'\n        )\n        return x_sma\n        \n        \n    def plot(self, title='summary', save_dir=None):\n        plt.figure(0, figsize=(15,10))\n        plt.plot(\n            self.val_v,\n            label='val'\n        )\n        \n        plt.plot(\n            self.avg_v,\n            label='avg'\n        )\n        plt.plot(\n            self.moving_average(self.val_v),\n            label=f'sma {self.n_sma}'\n        )\n        \n        plt.ylim( self.plt_ylim_v)\n        plt.title(title)\n        plt.legend()\n        plt.grid()\n        if save_dir is not None:\n            save_dir.mkdir(exist_ok=True)\n            plt.savefig(save_dir / (title + '.png'))\n            \n        plt.show()\n        \n        return None\n    \n# am = AverageMeter()\n\n# for x in np.random.random(20000):\n#     am.update(x)\n    \n# am.plot()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.396977Z","iopub.execute_input":"2022-08-10T23:34:21.397552Z","iopub.status.idle":"2022-08-10T23:34:21.411340Z","shell.execute_reply.started":"2022-08-10T23:34:21.397516Z","shell.execute_reply":"2022-08-10T23:34:21.410370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def batch_predict(inputs, model, batch_size, device, f_code=None, f_md=None):\n    n_samples = inputs['input_ids'].shape[0]\n    \n    if n_samples > batch_size:\n        preds_v = []\n        \n        for i_s in range(0, n_samples, batch_size):\n            i_e = min(i_s + batch_size, n_samples)\n            if i_e > i_s:\n                batch_inputs = transformers.tokenization_utils_base.BatchEncoding(\n                    {\n                        k: inputs[k][i_s:i_e].to(device) for k in inputs.keys()\n                    }\n                )\n                \n                if (f_code is not None) and (f_md is not None):\n                    preds = model(\n                        batch_inputs,\n                        f_code=f_code[i_s:i_e].to(device),\n                        f_md=f_md[i_s:i_e].to(device),\n                    )\n                    \n                else:\n                    preds = model(batch_inputs)\n\n                preds_v.append(preds)\n\n        preds_v = torch.cat(preds_v, axis=0)\n        \n    else:\n        if (f_code is not None) and (f_md is not None):\n            preds_v = model(\n                inputs.to(device),\n                f_code=f_code.to(device),\n                f_md=f_md.to(device),\n            )\n            \n        else:\n            preds_v = model(\n                inputs.to(device)\n            )\n    \n    assert preds_v.shape[0] == n_samples, f'preds_v.shape={preds_v.shape}  n_samples={n_samples}'\n    \n    return preds_v\n\n\ndef train_fn(epoch, cfg, dl_trn, model, scheduler, device):\n    criterion = model.criterion\n    optimizer = model.optimizer\n    \n    model.train()\n    \n    scaler = torch.cuda.amp.GradScaler(enabled=cfg.apex)\n    losses = AverageMeter()\n    global_step = 0\n    grad_norm = -1\n    max_bs = 0\n    \n    itt = tqdm(dl_trn)\n    for step, data in enumerate(itt):\n        if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n            if data['n_cells'] > cfg.batch_size:\n                continue\n                \n            if  data['n_cells'] > max_bs:\n                max_bs = data['n_cells']\n                    \n        inputs = data['inputs']\n        target = data['target']\n        \n        target     = target.to(device)\n        batch_size = target.size(0)\n        \n        with torch.cuda.amp.autocast(enabled=cfg.apex):\n            if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n                if cfg.code_use_abs_pos_enc:\n                    y_preds = batch_predict(\n                        inputs,\n                        model,\n                        cfg.batch_size,\n                        device,\n                        f_code=data['f_code'],\n                        f_md=data['f_md'],\n                    )\n                    \n                else:\n                    y_preds = batch_predict(\n                        inputs,\n                        model,\n                        cfg.batch_size,\n                        device,\n                    )\n                \n            else:\n                for k, v in inputs.items():\n                    inputs[k] = v.to(device)\n                    \n                y_preds = model(inputs)\n            \n            if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n                f_code = data['f_code'].to(device)\n                f_md   = data['f_md'].to(device)\n                loss = model.criterion(y_preds, target, f_code, f_md)\n\n            else:\n                target = target.type(y_preds.dtype)\n                loss = criterion(y_preds, target)\n\n            if cfg.use_sample_weights:\n                assert loss.shape == data['w'].shape, f\"ERROR: loss.shape={loss.shape} != data['w'].shape={data['w'].shape}\"\n                loss = (loss * data['w'].to(device)).mean()\n            \n        loss_item = loss.item()\n        \n        if np.isnan( loss_item ):\n            #print(f\"WARNING, train_fn: loss is NaN, skipping sample_id={data['sample_id']}\", file=sys.stderr)\n            continue\n        \n        losses.update(loss_item, batch_size)\n        \n        \n        if cfg.gradient_accumulation_steps > 1:\n            loss = loss / cfg.gradient_accumulation_steps\n        \n        scaler.scale(loss).backward()\n        \n        if (step + 1) % cfg.gradient_accumulation_steps == 0:\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.max_grad_norm)\n            \n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            global_step += 1\n            \n            if cfg.use_scheduler:\n                scheduler.step()\n        \n        itt.set_description(f'T[E{epoch:2d} L{losses.avg:0.05f} lr{scheduler.get_lr()[0]:0.02e} G{grad_norm:0.1e} BS{max_bs}]')\n        \n        # if step > 1000:\n        #     break\n    \n    losses.plot(\n        title=f'TRN_loss_[E={epoch}]',\n        save_dir=cfg.checkpoints_dir / 'figs'\n    )\n    \n    return losses.avg\n\n\ndef valid_fn(epoch, cfg, dl_val, model, device):\n    criterion = model.criterion\n    \n    losses = AverageMeter()\n    model.eval()\n    \n    sample_ids_v = []\n    preds_v = []\n    targets_v = []\n    with torch.no_grad():\n        itt = tqdm(dl_val)\n        for step, data in enumerate(itt):\n            inputs = data['inputs']\n            target = data['target']\n            sample_ids = data['sample_id']\n            \n            target = target.to(device)\n            batch_size = target.size(0)\n            \n            with torch.cuda.amp.autocast(enabled=cfg.apex):\n                if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n                    if cfg.code_use_abs_pos_enc:\n                        y_preds = batch_predict(\n                            inputs,\n                            model,\n                            cfg.batch_size,\n                            device,\n                            f_code=data['f_code'],\n                            f_md=data['f_md'],\n                        )\n\n                    else:\n                        y_preds = batch_predict(\n                            inputs,\n                            model,\n                            cfg.batch_size,\n                            device,\n                        )\n\n                else:\n                    for k, v in inputs.items():\n                        inputs[k] = v.to(device)\n\n                    y_preds = model(inputs)\n\n                if cfg.model_type == 'matrix' or cfg.model_type == 'sorting':\n                    f_code = data['f_code'].to(device)\n                    f_md   = data['f_md'].to(device)\n                    \n                    loss = model.criterion(\n                        y_preds,\n                        target,\n                        f_code,\n                        f_md\n                    )\n\n                    y_preds = model.criterion.predict_rank_v(\n                        y_preds,\n                        f_code,\n                        f_md,\n                        return_order_v=False,\n                    )\n\n                    target = data['nb_df']['rank'].values\n\n                else:\n                    loss = criterion(y_preds, target)\n\n                    y_preds = y_preds.detach().cpu().numpy()\n\n                    target = target.detach().cpu().numpy()\n\n                if cfg.use_sample_weights:\n                    assert loss.shape == data['w'].shape, f\"ERROR: loss.shape={loss.shape} != data['w'].shape={data['w'].shape}\"\n                    loss = (loss * data['w'].to(device)).mean()\n            \n            \n            loss_item = loss.item()\n            \n            if np.isnan( loss_item ):\n                print(f\"WARNING, valid_fn: loss is NaN, skipping sample_id={data['sample_id']}\", file=sys.stderr)\n                continue\n            \n            losses.update(loss_item, batch_size)\n\n            preds_v.append(y_preds)\n            targets_v.append(target)\n            sample_ids_v.append(sample_ids)\n            \n            itt.set_description(f'V[E{epoch:2d} L{losses.avg:0.05f}]')\n    \n    \n#     from IPython import embed; embed()\n    \n    preds_v = np.concatenate(preds_v, axis=0)\n    targets_v = np.concatenate(targets_v, axis=0)\n    sample_ids_v = np.concatenate(sample_ids_v, axis=0)\n    \n    preds_d = {'sample_ids_v':sample_ids_v, 'preds_v':preds_v, 'targets_v':targets_v}\n    \n    losses.plot(\n        title=f'VAL_loss_[E={epoch}]',\n        save_dir=cfg.checkpoints_dir / 'figs'\n    )\n    \n    return losses.avg, preds_d\n\n\ndef model_inference(model, dl, device='cuda:0', add_targets=False, add_features=False):\n    model.eval()\n    model.to(device)\n    \n    sample_ids_v = []\n    preds_v      = []\n    features_v   = []\n    targets_v    = []\n    with torch.no_grad():\n        for i_d, data in enumerate(tqdm(dl)):\n            inputs = data['inputs']\n            sample_ids = data['sample_id']\n            \n                \n            if model.cfg.model_type == 'matrix' or model.cfg.model_type == 'sorting':\n                if model.cfg.code_use_abs_pos_enc:\n                    y_pred = batch_predict(\n                        inputs,\n                        model,\n                        model.cfg.batch_size,\n                        device,\n                        f_code=data['f_code'],\n                        f_md=data['f_md'],\n                    )\n\n                else:\n                    y_pred = batch_predict(\n                        inputs,\n                        model,\n                        model.cfg.batch_size,\n                        device,\n                    )\n            else:\n                for k, v in inputs.items():\n                    inputs[k] = v.to(device)\n                \n                if add_features:\n                    y_pred, y_features = model(inputs, return_features=True)\n                    features_v.append(y_features.detach().cpu().numpy())\n                    \n                else:\n                    y_pred = model(inputs, return_features=False)\n                \n                \n            if model.cfg.model_type == 'matrix' or model.cfg.model_type == 'sorting':\n                y_pred = model.criterion.predict_rank_v(\n                    y_pred,\n                    data['f_code'].to(device),\n                    data['f_md'].to(device),\n                    return_order_v=False,\n                )\n                \n                if add_targets:\n                    target = data['nb_df']['rank'].values\n                \n            else:\n                y_pred = y_pred.detach().cpu().numpy()\n                \n                if add_targets:\n                    target = data['target'].detach().cpu().numpy()\n                \n\n            sample_ids_v.append( sample_ids )\n            preds_v.append( y_pred )\n            \n            if add_targets:\n                targets_v.append(target)\n            \n#             if i_d == 100:\n#                 break\n                \n    sample_ids_v = np.concatenate(sample_ids_v)\n    \n#     from IPython import embed; embed()\n    preds_v = np.concatenate(preds_v)\n    \n    preds_d = {'sample_ids_v':sample_ids_v, 'preds_v':preds_v}\n    \n    if add_targets:\n        preds_d['targets_v'] = np.concatenate(targets_v)\n        \n        \n    if add_features:\n        if model.cfg.model_type == 'matrix' or model.cfg.model_type == 'sorting':\n            preds_d['features_v'] = None\n            \n        else:\n            preds_d['features_v'] = np.concatenate(features_v)\n    \n    return preds_d","metadata":{"papermill":{"duration":0.066091,"end_time":"2022-02-08T03:59:08.747312","exception":false,"start_time":"2022-02-08T03:59:08.681221","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:34:21.412882Z","iopub.execute_input":"2022-08-10T23:34:21.413844Z","iopub.status.idle":"2022-08-10T23:34:21.455208Z","shell.execute_reply.started":"2022-08-10T23:34:21.413807Z","shell.execute_reply":"2022-08-10T23:34:21.454152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def save_last_epoch(model, scheduler, epoch, filename=\"last.ckpt\", verbose=True):\n    \n#     save_path = os.path.join(\n#         model.cfg.checkpoints_dir,\n#         filename\n#     )\n    \n#     if verbose:\n#         print(f'Saving model: {save_path}')\n        \n#     # Saving last epoch\n#     torch.save(\n#         {\n#             'model': model.state_dict(),\n#             'optimizer': model.optimizer.state_dict(),\n#             'scheduler': scheduler.state_dict(),\n#             'CFG': model.cfg.get_argsdict(),\n#             'epoch':epoch,\n#         },\n#         save_path\n#     )\n#     return None\n\n\ndef save_last_epoch(\n    model,\n    scheduler,\n    epoch,\n    filename=\"last.ckpt\",\n    save_optimizer=True,\n    verbose=True\n):\n    \n    save_path = os.path.join(\n        model.cfg.checkpoints_dir,\n        filename\n    )\n    \n    if verbose:\n        print(f'Saving model: {save_path}')\n        \n    # Saving last epoch\n    torch.save(\n        {\n            'model': model.state_dict(),\n            'scheduler': scheduler.state_dict(),\n            'CFG': model.cfg.get_argsdict(),\n            'epoch':epoch,\n        },\n        save_path\n    )\n    \n    if save_optimizer:\n        torch.save(\n            {\n                'optimizer': model.optimizer.state_dict(),\n                'scheduler': scheduler.state_dict(),\n                'CFG': model.cfg.get_argsdict(),\n                'epoch':epoch,\n            },\n            save_path.replace('.ckpt', '') + '.opt'\n        )\n    \n    return None\n\n\ndef optimizer_to(optim, device):\n    for param in optim.state.values():\n        # Not sure there are any global tensors in the state dict\n        if isinstance(param, torch.Tensor):\n            param.data = param.data.to(device)\n            if param._grad is not None:\n                param._grad.data = param._grad.data.to(device)\n                \n        elif isinstance(param, dict):\n            for subparam in param.values():\n                if isinstance(subparam, torch.Tensor):\n                    subparam.data = subparam.data.to(device)\n                    if subparam._grad is not None:\n                        subparam._grad.data = subparam._grad.data.to(device)\n                        \n    return None","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:21.458481Z","iopub.execute_input":"2022-08-10T23:34:21.458761Z","iopub.status.idle":"2022-08-10T23:34:21.471623Z","shell.execute_reply.started":"2022-08-10T23:34:21.458736Z","shell.execute_reply":"2022-08-10T23:34:21.470683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(df, nb_feat_d, i_fold, cfg, tokenizer, logger):\n    logger.info(get_time_str() + f\"========== fold: {i_fold} training ==========\")\n    \n    device = torch.device(cfg.device)\n    # Dataloaders\n    ds_trn, ds_val, ds_tst, dl_trn, dl_val, dl_tst = get_dataloaders(\n        df,\n        nb_feat_d,\n        i_fold,\n        cfg,\n        tokenizer,\n        val_bs_mult=cfg.val_bs_mult,\n    )\n    \n#     ds_val = ds_trn\n#     dl_val = dl_trn\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = NB_Model(cfg, config_path=None, pretrained=True)\n    \n    scheduler = get_scheduler(\n        cfg,\n        model.optimizer,\n        num_train_steps = cfg.epochs * len(dl_trn) / cfg.gradient_accumulation_steps\n    )\n    \n    torch.save(\n        model.config,\n        os.path.join(cfg.checkpoints_dir, 'config.pth')\n    )\n    \n    torch.save(\n        model.cfg.get_argsdict(),\n        os.path.join(cfg.checkpoints_dir, 'model_cfg.pth')\n    )\n    \n    load_opt_ckpt = False\n    if cfg.initial_ckpt_path is None:\n        initial_epoch = 0\n        \n    else:\n        print(f'Restoring: \"{cfg.initial_ckpt_path}\"')\n        \n        if (\"last.ckpt\" in cfg.initial_ckpt_path) or (\"error_recovery\" in cfg.initial_ckpt_path):\n            opt_ckpt_path = cfg.initial_ckpt_path.replace('.ckpt', '') + '.opt'\n            load_opt_ckpt = cfg.restore_optimizer and os.path.exists(opt_ckpt_path)\n            \n            restored_epoch = model.restore_model(\n                cfg.initial_ckpt_path,\n                restore_optimizer=cfg.restore_optimizer and (not load_opt_ckpt),\n                scheduler=scheduler if cfg.restore_scheduler else None,\n            )\n            \n        else:\n            restored_epoch = model.restore_model(cfg.initial_ckpt_path)\n        \n        if restored_epoch is None:\n            try:\n                initial_epoch = [int( l.replace('E', '') ) for l in os.path.basename(cfg.initial_ckpt_path).split('_') if len(l) > 0 and l[0] == 'E'][0] + 1\n\n            except Exception as e:\n                print(f\"WARNING: Unable to retrieve Last Epoch Number from checkpoint name: {os.path.basename(cfg.initial_ckpt_path)}.\", file=sys.stderr)\n                initial_epoch = 0 \n        else:\n            initial_epoch = restored_epoch + 1\n            \n        seed_everything(seed=cfg.train_seed + initial_epoch + 1)\n            \n        \n    model.to(device)\n    if load_opt_ckpt:\n        print(f'Restoring optimizer: \"{opt_ckpt_path}\"')\n        model.optimizer.load_state_dict(\n            torch.load(\n                opt_ckpt_path,\n                map_location=device,\n            )['optimizer']\n        )\n        \n    elif cfg.restore_optimizer:\n        optimizer_to(model.optimizer, device)\n    \n    saver = K_CkptSaver(k=cfg.save_top_k, max_score=True)\n    best_scores_d = None\n    for epoch in range(initial_epoch, cfg.epochs):\n\n        start_time = time.time()\n        \n        try:\n            # train\n            trn_loss = train_fn(epoch, cfg, dl_trn, model, scheduler, device)\n            \n        except Exception as e:\n            print('ERROR detected. trying to save last epoch')\n            save_last_epoch(model, scheduler, epoch, filename=f\"error_recovery_{int(time.time())}.ckpt\")\n            raise e\n        \n        # Saving last epoch\n        save_last_epoch(model, scheduler, epoch, filename=\"last.ckpt\", verbose=True)\n\n        # eval\n        val_loss, val_preds_d = valid_fn(epoch, cfg, dl_val, model, device)\n        \n        \n        # scoring ...\n        comp_score = calc_score(\n            val_preds_d,\n            ds_val,\n            cfg,\n        )\n\n        elapsed = time.time() - start_time\n\n        logger.info(f'Epoch {epoch} - trn_loss: {trn_loss:.4f}  val_loss: {val_loss:.4f}  time: {get_time_str(\"\")}  epoch_time: {elapsed:.0f}s')\n        logger.info(f'Epoch {epoch} - comp_score: {comp_score:0.04f}')\n        \n\n        if saver.is_better(comp_score):\n            logger.info(f'Epoch {epoch} - New Top k Score: {comp_score:.4f}')\n            \n            best_ckpt_path = os.path.join(\n                cfg.checkpoints_dir,\n                f\"{cfg.model.replace('/', '-')}_F{i_fold}_E{epoch}_TL{trn_loss:0.04f}_VL{val_loss:0.04f}_VS{comp_score:0.04f}.ckpt\"\n            )\n            \n            data = {\n                'model': model.state_dict(),\n                #'optimizer': model.optimizer.state_dict(),\n                #'scheduler': scheduler.state_dict(),\n                'CFG': model.cfg.get_argsdict(),\n                #'val_preds_d': val_preds_d,\n                'scores_d': {'comp_score':comp_score, 'val_loss':val_loss, 'trn_loss':trn_loss},\n                'epoch':epoch,\n            }\n            \n            if saver.is_the_best(comp_score):\n                logger.info(f'Epoch {epoch} - New Best Score!!!: {comp_score:.4f}')\n                best_scores_d = data['scores_d']\n            \n            \n            saver.save(\n                comp_score,\n                data,\n                best_ckpt_path,\n            )\n            \n            logger.info(f'Epoch {epoch} - Save path: \"{best_ckpt_path}\"')\n            \n            \n\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return best_scores_d\n","metadata":{"papermill":{"duration":0.058285,"end_time":"2022-02-08T03:59:08.837273","exception":false,"start_time":"2022-02-08T03:59:08.778988","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:34:21.473252Z","iopub.execute_input":"2022-08-10T23:34:21.473583Z","iopub.status.idle":"2022-08-10T23:34:21.495857Z","shell.execute_reply.started":"2022-08-10T23:34:21.473550Z","shell.execute_reply":"2022-08-10T23:34:21.494812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"if CFG.do_training:\n    seed_everything(CFG.train_seed)\n    \n    if CFG.do_training:\n        LOGGER.info(get_time_str() + '==================== NEW RUN ====================')\n        \n        LOGGER.info(f'CFG: {CFG.get_argsdict()}')\n        LOGGER.info(f\"Train Dataset: n_nb = {CFG.max_dataset_size}\")\n        LOGGER.info(f'batch_size: {CFG.batch_size}')\n        LOGGER.info(f'model_type: {CFG.model_type}')\n        LOGGER.info(f'model: {CFG.model}')\n        LOGGER.info(f\"max_md_tkn_len: {CFG.max_md_tkn_len}\")\n        LOGGER.info(f\"max_code_tkn_len: {CFG.max_code_tkn_len}\")\n        LOGGER.info(f'max_abs_tkn_len: {CFG.max_abs_tkn_len}')\n    \n    \n    for i_fold in range(CFG.n_folds):\n        if i_fold in CFG.trn_fold:\n            best_scores_d = train_loop(\n                df=df_nb,\n                nb_feat_d=nb_feat_d,\n                i_fold=i_fold,\n                cfg=CFG,\n                tokenizer=tokenizer,\n                logger=LOGGER,\n            )\n\n    sys.exit(0)","metadata":{"papermill":{"duration":26802.974645,"end_time":"2022-02-08T11:25:51.850092","exception":false,"start_time":"2022-02-08T03:59:08.875447","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-09T03:48:38.776774Z","iopub.execute_input":"2022-08-09T03:48:38.779628Z","iopub.status.idle":"2022-08-09T03:48:38.789739Z","shell.execute_reply.started":"2022-08-09T03:48:38.779590Z","shell.execute_reply":"2022-08-09T03:48:38.788784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model inference","metadata":{}},{"cell_type":"code","source":"if CFG.do_inference:\n    pred_d_d = {}\n    pred_d_d = load_obj('../input/some-predictions/pred_d_d.pickle')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:40:50.688692Z","iopub.execute_input":"2022-08-10T23:40:50.689758Z","iopub.status.idle":"2022-08-10T23:40:57.634360Z","shell.execute_reply.started":"2022-08-10T23:40:50.689721Z","shell.execute_reply":"2022-08-10T23:40:57.632292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    ds_trn, ds_val, ds_tst, dl_trn, dl_val, dl_tst = get_dataloaders(df_nb, None, 0, CFG, tokenizer, 2)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:28.185595Z","iopub.execute_input":"2022-08-10T23:34:28.186189Z","iopub.status.idle":"2022-08-10T23:34:40.004015Z","shell.execute_reply.started":"2022-08-10T23:34:28.186150Z","shell.execute_reply":"2022-08-10T23:34:40.002903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    model = NB_Model(\n        CFG,\n        config_path=CFG.checkpoints_dir / \"config.pth\",\n        pretrained=False\n    )\n    # model.restore_model('../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.1955_VL1.0152_VS0.8787.ckpt')\n    model.eval()\n    pass","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:40.005753Z","iopub.execute_input":"2022-08-10T23:34:40.006185Z","iopub.status.idle":"2022-08-10T23:34:42.290907Z","shell.execute_reply.started":"2022-08-10T23:34:40.006144Z","shell.execute_reply":"2022-08-10T23:34:42.289774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    model = NB_Model(\n        Args(\n            torch.load('../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/model_cfg.pth')\n        ),\n        config_path='../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/config.pth',\n        pretrained=False\n    )\n    model.eval()","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:34:42.292435Z","iopub.execute_input":"2022-08-10T23:34:42.293098Z","iopub.status.idle":"2022-08-10T23:34:44.767149Z","shell.execute_reply.started":"2022-08-10T23:34:42.293051Z","shell.execute_reply":"2022-08-10T23:34:44.766171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    ckpt_v = [\n        # VAST Entrenado con Metanotebooks + GitHub (semilla R43E18)\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.1955_VL1.0152_VS0.8787.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/error_recover.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL1.3469_VL1.1779_VS0.8796.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL1.3333_VL1.1726_VS0.8798.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL1.3232_VL1.1854_VS0.8799.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL1.2866_VL1.1460_VS0.8799.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL1.2736_VL1.1491_VS0.8810.ckpt\",\n\n        \n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26E27.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26E27.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E25E26E27.ckpt\",\n        \n        \n        # Este se utilizó como semilla para entrenar los modelos de Vast (25 folds , Sin MNB, entrnando a partir de R42E15, GitHubDS_85kNB )\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E17_TL0.9883_VL1.0027_VS0.8653.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E16_TL1.0187_VL1.0012_VS0.8711.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E18_TL0.9703_VL1.0359_VS0.8694.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TL0.9550_VL1.0279_VS0.8666.ckpt\",\n        \n        # Nuevo entrenamiento 25 folds (Sin MNB, solamente notebooks originales de kaggle)\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TLnan_VL1.0943_VS0.8684.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E15_TLnan_VL1.0837_VS0.8692.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E13_TLnan_VL1.0509_VS0.8688.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E12_TLnan_VL1.0464_VS0.8680.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E11_TLnan_VL1.0481_VS0.8683.ckpt\",\n        \n        # VAST Entrendo con todos lo NBs y los notebooks extra de kaggle (Entrenado a partir del run_41, sinMetaNB)\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E9_TL1.0432_VL2.4768_VS0.8614.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659290334.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E11_TL0.7875_VL2.5361_VS0.8666.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E12_TL0.7707_VL2.6240_VS0.8683.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E13_TL0.7615_VL2.6432_VS0.8668.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E14_TL0.7530_VL2.6928_VS0.8701.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E15_TL0.7460_VL2.7321_VS0.8674.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E16_TL0.7405_VL2.6780_VS0.8700.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E17_TL0.7317_VL2.7011_VS0.8690.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659808200.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL0.6828_VL2.7560_VS0.8683(1).ckpt\",\n        \n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E12E13E14E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14ER34E11E12E13E14E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E17ER00E19.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E16E17ER00E19.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16E17ER00E19.ckpt\",\n        \n        \n        # VAST Fork de run-48 (Entrenado a aprtir de run-48-E13, conMetaNB usando GitHubNB)\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/error_recovery_1659541967.ckpt\", # No usa GHnb\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E14_TL1.3146_VL1.1827_VS0.8690.ckpt\",\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E15_TL1.4038_VL1.1817_VS0.8627.ckpt\",\n        \n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14.ckpt\",\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14E15.ckpt\",\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoE14E15.ckpt\",\n        \n        # Este era el mejor modelo LB=0.7014 (entrenado conMetaNB)\n        \"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E8_TLnan_VL1.0432_VS0.8685.ckpt\",\n        \"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E5_TLnan_VL1.0372_VS0.8682.ckpt\",\n        \"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E6_TLnan_VL1.0529_VS0.8677.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E7_TL1.1609_VL1.0455_VS0.8672.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E3_TL1.2291_VL1.0381_VS0.8662.ckpt\",\n        \n    ]","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:26:25.904860Z","iopub.execute_input":"2022-08-11T00:26:25.905332Z","iopub.status.idle":"2022-08-11T00:26:25.922978Z","shell.execute_reply.started":"2022-08-11T00:26:25.905296Z","shell.execute_reply":"2022-08-11T00:26:25.922050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_run_and_epoch(path):\n    ret_str = \"\"\n    \n    if \"run-\" in path:\n        run_num = path.split('-')[1]\n        ret_str += f'R{run_num}' \n        \n    if \"_E\" in path:\n        try:\n            epoch_num = path.split('_E')[1].split('_')[0]\n            ret_str += f'-E{int(epoch_num):02d}'\n            \n        except Exception as e:\n            epoch_num = path.split('_E')[1].split('.')[0]\n            ret_str += f'-E{epoch_num}' \n        \n    if 'Fo' in path:\n        try:\n            epoch_num = path.split('Fo')[1].split('_')[0]\n            ret_str += f'-Fo{int(epoch_num):02d}'\n        except Exception as e:\n            epoch_num = path.split('_Fo')[1].split('.')[0]\n            ret_str += f'-Fo{epoch_num}' \n        \n        \n    if \"error_recovery_\" in path:\n        ret_str += \"-\" + os.path.basename(path).split('.')[0].replace('error_recovery_', 'ER')\n        \n    elif \"error_recover\" in path:\n        ret_str += \"-\" + os.path.basename(path).split('.')[0].replace('error_recover', 'ER')\n    \n    return ret_str","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:20:46.076447Z","iopub.execute_input":"2022-08-11T00:20:46.076819Z","iopub.status.idle":"2022-08-11T00:20:46.088589Z","shell.execute_reply.started":"2022-08-11T00:20:46.076787Z","shell.execute_reply":"2022-08-11T00:20:46.087488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    gt_orders = ds_val.get_gt_orders()\n\n    for ckpt_path in ckpt_v:\n        #ds_val.use_rnd_source_cuts = not ds_val.use_rnd_source_cuts\n        #if ds_val.use_rnd_source_cuts == False:\n        #    pred_d = pred_d_d.get(ckpt_path, None)\n        #else:\n        #    pred_d = None\n        \n        pred_d = pred_d_d.get(ckpt_path, None)\n        if pred_d is None:\n            print('Inference:', os.path.basename(ckpt_path))\n            model.restore_model(ckpt_path)\n            pred_d = model_inference(model, dl_val, device='cuda:0', add_features=False, add_targets=False)\n\n            pred_d_d[ckpt_path] = pred_d\n        \n        pred_orders, _ = get_orders_from_prediction(pred_d, ds_val, code_rank_column=None)\n\n        comp_score = kendall_tau(\n            gt_orders.loc[pred_orders.index.values],\n            pred_orders\n        )\n\n        print(get_run_and_epoch(ckpt_path), ':' , comp_score, ds_val.use_rnd_source_cuts)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T01:16:09.467385Z","iopub.execute_input":"2022-08-11T01:16:09.468072Z","iopub.status.idle":"2022-08-11T01:34:45.969010Z","shell.execute_reply.started":"2022-08-11T01:16:09.468034Z","shell.execute_reply":"2022-08-11T01:34:45.967953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Merge Ckpts \nif True and CFG.do_inference:\n    merge_ckpts(\n        ckpt_paths_v=[\n            \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/error_recovery_1659541967.ckpt\", # No usa GHnb\n            \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E14_TL1.3146_VL1.1827_VS0.8690.ckpt\",\n            \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E15_TL1.4038_VL1.1817_VS0.8627.ckpt\",\n        ],\n        input_key='model',\n        k_replace_v=None,\n        output_path='./run-48-MERGE_FoER67E14E15.ckpt',\n        output_key='model',\n        return_statedict=False,\n    )","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:23:25.420673Z","iopub.execute_input":"2022-08-11T00:23:25.422787Z","iopub.status.idle":"2022-08-11T00:23:34.638249Z","shell.execute_reply.started":"2022-08-11T00:23:25.422748Z","shell.execute_reply":"2022-08-11T00:23:34.637291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False and CFG.do_inference:\n    for k,v in list( pred_d_d.items() ):\n        if './run-45' == k[:8]:\n            new_k = \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/\" + os.path.basename(k)\n\n            if new_k not in pred_d_d.keys():\n                pred_d_d[new_k] = v\n\n    for k in pred_d_d.keys():\n        print(k)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:38.894350Z","iopub.execute_input":"2022-08-09T03:48:38.895477Z","iopub.status.idle":"2022-08-09T03:48:38.906668Z","shell.execute_reply.started":"2022-08-09T03:48:38.895441Z","shell.execute_reply":"2022-08-09T03:48:38.905729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    model = NB_Model(\n        Args(\n            torch.load('../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/model_cfg.pth')\n        ),\n        config_path='../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/config.pth',\n        pretrained=False\n    )\n    model.eval()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T22:57:05.419695Z","iopub.execute_input":"2022-08-09T22:57:05.420115Z","iopub.status.idle":"2022-08-09T22:57:07.826318Z","shell.execute_reply.started":"2022-08-09T22:57:05.420073Z","shell.execute_reply":"2022-08-09T22:57:07.825302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    ckpt_v = [\n        # GitHub Training + KaggleNB (conMetaNB)\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E18_TL1.1382_VL0.9904_VS0.8766.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL1.1069_VL0.9743_VS0.8782.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E20_TL1.2697_VL1.1128_VS0.8759.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.2412_VL1.1213_VS0.8753.ckpt\",\n        \n        '../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20.ckpt',\n        '../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20E21.ckpt',\n        '../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E18E19E20E21.ckpt',\n        \n        \n        # New Training + KaggleNB (sinMetaNB a partir de run-45 E21)\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E22_TL0.7064_VL1.8204_VS0.8831.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E23_TL0.6845_VL1.7250_VS0.8847.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL0.6709_VL1.7212_VS0.8876.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL0.6590_VL1.7388_VS0.8875.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL0.6480_VL1.7243_VS0.8874.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL0.6400_VL1.7032_VS0.8906.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL0.6327_VL1.7068_VS0.8894.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E29_TL0.6240_VL1.6679_VS0.8902.ckpt\",\n        \n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23E24E25.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E23E24E25.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25.ckpt\",\n        \n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E27E28.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E26E27E28.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E25E26E27E28.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25E26E27E28.ckpt\",\n        \n    ]","metadata":{"execution":{"iopub.status.busy":"2022-08-09T22:57:05.409410Z","iopub.execute_input":"2022-08-09T22:57:05.409970Z","iopub.status.idle":"2022-08-09T22:57:05.417791Z","shell.execute_reply.started":"2022-08-09T22:57:05.409931Z","shell.execute_reply":"2022-08-09T22:57:05.416558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    gt_orders = ds_val.get_gt_orders()\n\n    for ckpt_path in ckpt_v:\n        #ds_val.use_rnd_source_cuts = not ds_val.use_rnd_source_cuts\n        #if ds_val.use_rnd_source_cuts == False:\n        #    pred_d = pred_d_d.get(ckpt_path, None)\n        #else:\n        #    pred_d = None\n        \n        pred_d = pred_d_d.get(ckpt_path, None)\n        if pred_d is None:\n            print('Inference:', os.path.basename(ckpt_path))\n            model.restore_model(ckpt_path)\n            pred_d = model_inference(model, dl_val, device='cuda:0', add_features=False, add_targets=False)\n\n            pred_d_d[ckpt_path] = pred_d\n        \n        pred_orders, _ = get_orders_from_prediction(pred_d, ds_val, code_rank_column=None)\n\n        comp_score = kendall_tau(\n            gt_orders.loc[pred_orders.index.values],\n            pred_orders\n        )\n\n        print(get_run_and_epoch(ckpt_path), ':' , comp_score, ds_val.use_rnd_source_cuts)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T22:57:07.828110Z","iopub.execute_input":"2022-08-09T22:57:07.829620Z","iopub.status.idle":"2022-08-09T23:16:30.681097Z","shell.execute_reply.started":"2022-08-09T22:57:07.829581Z","shell.execute_reply":"2022-08-09T23:16:30.679772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    save_obj( pred_d_d, 'pred_d_d.pickle' )","metadata":{"execution":{"iopub.status.busy":"2022-08-11T01:46:36.141613Z","iopub.execute_input":"2022-08-11T01:46:36.142685Z","iopub.status.idle":"2022-08-11T01:46:44.094023Z","shell.execute_reply.started":"2022-08-11T01:46:36.142644Z","shell.execute_reply":"2022-08-11T01:46:44.092199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    gt_orders = ds_val.get_gt_orders()\n    \n    pred_d_to_ensemble_v = []\n    for ckpt_path in [\n        # Model Type 44\n        # VAST Entrenado con Metanotebooks + GitHub (semilla R43E18)\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.1955_VL1.0152_VS0.8787.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/error_recover.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL1.3469_VL1.1779_VS0.8796.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL1.3333_VL1.1726_VS0.8798.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL1.3232_VL1.1854_VS0.8799.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL1.2866_VL1.1460_VS0.8799.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL1.2736_VL1.1491_VS0.8810.ckpt\",\n\n        \n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E25E26E27.ckpt\",\n        \n        \n        # Este se utilizó como semilla para entrenar los modelos de Vast (25 folds , Sin MNB, entrnando a partir de R42E15, GitHubDS_85kNB )\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E17_TL0.9883_VL1.0027_VS0.8653.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E16_TL1.0187_VL1.0012_VS0.8711.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E18_TL0.9703_VL1.0359_VS0.8694.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TL0.9550_VL1.0279_VS0.8666.ckpt\",\n        \n        # Nuevo entrenamiento 25 folds (Sin MNB, solamente notebooks originales de kaggle)\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TLnan_VL1.0943_VS0.8684.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E15_TLnan_VL1.0837_VS0.8692.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E13_TLnan_VL1.0509_VS0.8688.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E12_TLnan_VL1.0464_VS0.8680.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E11_TLnan_VL1.0481_VS0.8683.ckpt\",\n        \n        # VAST Entrendo con todos lo NBs y los notebooks extra de kaggle (Entrenado a partir del run_41, sinMetaNB)\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E9_TL1.0432_VL2.4768_VS0.8614.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659290334.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E11_TL0.7875_VL2.5361_VS0.8666.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E12_TL0.7707_VL2.6240_VS0.8683.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E13_TL0.7615_VL2.6432_VS0.8668.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E14_TL0.7530_VL2.6928_VS0.8701.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E15_TL0.7460_VL2.7321_VS0.8674.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E16_TL0.7405_VL2.6780_VS0.8700.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E17_TL0.7317_VL2.7011_VS0.8690.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659808200.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL0.6828_VL2.7560_VS0.8683(1).ckpt\",\n        \n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E12E13E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14ER34E11E12E13E14E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E16E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16E17ER00E19.ckpt\",\n        \n        \n        # VAST Fork de run-48 (Entrenado a aprtir de run-48-E13, conMetaNB usando GitHubNB)\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/error_recovery_1659541967.ckpt\", # No usa GHnb\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E14_TL1.3146_VL1.1827_VS0.8690.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E15_TL1.4038_VL1.1817_VS0.8627.ckpt\",\n        \n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14E15.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoE14E15.ckpt\",\n        \n        # Este era el mejor modelo LB=0.7014 (entrenado conMetaNB)\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E8_TLnan_VL1.0432_VS0.8685.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E5_TLnan_VL1.0372_VS0.8682.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E6_TLnan_VL1.0529_VS0.8677.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E7_TL1.1609_VL1.0455_VS0.8672.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E3_TL1.2291_VL1.0381_VS0.8662.ckpt\",\n        \n        \n        # Model Type 45\n        # GitHub Training + KaggleNB (conMetaNB)\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E18_TL1.1382_VL0.9904_VS0.8766.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL1.1069_VL0.9743_VS0.8782.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E20_TL1.2697_VL1.1128_VS0.8759.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.2412_VL1.1213_VS0.8753.ckpt\",\n        \n        #'../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20.ckpt',\n        #'../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20E21.ckpt',\n        '../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E18E19E20E21.ckpt',\n        \n        \n        # New Training + KaggleNB (sinMetaNB a partir de run-45 E21)\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E22_TL0.7064_VL1.8204_VS0.8831.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E23_TL0.6845_VL1.7250_VS0.8847.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL0.6709_VL1.7212_VS0.8876.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL0.6590_VL1.7388_VS0.8875.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL0.6480_VL1.7243_VS0.8874.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL0.6400_VL1.7032_VS0.8906.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL0.6327_VL1.7068_VS0.8894.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E29_TL0.6240_VL1.6679_VS0.8902.ckpt\",\n        \n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23E24E25.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E23E24E25.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25.ckpt\",\n        \n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E26E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E25E26E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25E26E27E28.ckpt\",\n    \n    ]:\n        pred_d_to_ensemble_v.append(\n            pred_d_d[ckpt_path]\n        )\n        \n        print(get_run_and_epoch(ckpt_path))\n        \n    ens_preds_d = ensemble_preds(pred_d_to_ensemble_v, ds_val, w_v=None)\n    pred_orders, _ = get_orders_from_prediction(ens_preds_d, ds_val, code_rank_column=None)\n\n    comp_score = kendall_tau(\n        gt_orders.loc[pred_orders.index.values],\n        pred_orders\n    )\n    \n\n    print('Ensemble Score:', comp_score)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T01:44:27.551851Z","iopub.execute_input":"2022-08-11T01:44:27.552425Z","iopub.status.idle":"2022-08-11T01:44:42.536572Z","shell.execute_reply.started":"2022-08-11T01:44:27.552375Z","shell.execute_reply":"2022-08-11T01:44:42.535496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Ensemble summary\n```\nrun_44\nEnsemble Score: 0.8808954705419908\n\nrun_43\nEnsemble Score: 0.8729411279550664\n    \nrun_44 + run_43[0]\nEnsemble Score: 0.8825536246425979\n    \nrun_45\nEnsemble Score: 0.8803968250378114\n    \n    \nrun_44 + run_45\nEnsemble Score: 0.8854586029026008\n    \nrun_44[:2] + run_45[:2]\nEnsemble Score: 0.8854588133015476\n    \nrun_44[:2] + run_45[0]\nEnsemble Score: 0.8849058848690651\n    \n    \nrun_44 + run_45 + run_43[0]\nEnsemble Score: 0.8861764841094617\n    \n    \nrun_44[0] + run_45[0] + run_43[0]\nEnsemble Score: 0.884984153277316\n    \n\n\n\nrun_41[:2]\nEnsemble Score: 0.8963259189095106\n\nrun_41[:3]\nEnsemble Score: 0.8973726536703939 # Entrenados con los datos\n    \n    \nrun_44[0] + run_41[:2] + run_45[:2]\nEnsemble Score: 0.8936761545721111\n\nrun_44[:2] + run_41[:2] + run_45[:2]\nEnsemble Score: 0.8929494366094799\n    \nLB Score:0.8750\nrun_44[:2] + run_41[:3] + run_45[:2]\nEnsemble Score: 0.8944954480713834\n\n\nrun_44[:2] + run_43[0] + run_41[:2] + run_45[:2]\nEnsemble Score: 0.8925694561113583\n\n\n\nrun_48[:4]\nEnsemble Score: 0.9048525466320336\n\nrun_48[2:]\nEnsemble Score: 0.9059007541855452\n\n\nR44-E24:\nR44-E25:\nEnsemble Score: 0.8817151848391569\n\n##################################\nR48-E12:\nR41-E08:\nEnsemble Score: 0.904095952018941\n\nR44-E24:\nerror_recover:\nR44-E25:\nEnsemble Score: 0.8830446957846887\n\nR45-E19:\nR45-E20:\nEnsemble Score: 0.8804839302018326\n\nerror_recover:\nR44-E24:\nR44-E25:\nR48-E12:\nR41-E08:\nR45-E19:\nR45-E20:\nEnsemble Score: 0.8950776219574865\n\n\nerror_recover:\nR44-E24:\nR44-E25:\nR48-E12:\nR41-E08:\nR45-E19:\nR45-E20:\nEnsemble Score: 0.8950776219574865\n\n\n\n\nerror_recover:\nR44-E21:\nR41-E05:\nR41-E06:\nR41-E08:\nR45-E18:\nR45-E19:\nEnsemble Score: 0.8944981832576931\nLB Score:0.8750\n\n\n\nerror_recover:\nR44-E24:\nR44-E25:\nR48-E11:\nR48-E12:\nR41-E08:\nR45-E19:\nEnsemble Score: 0.8983137681599276\nLB_Score: 0.8772\n\n\nR44-ERE24E25E26\nR48-E14E15E16\nR48-ER1659541967\nR45-E18E19E20E21\nR45-E22E23\nEnsemble Score: 0.9022989346133725\n\nR44-ERE24E25E26\nR48-E16\nR48-E14E15E16\nR48-ER1659541967\nR45-E18E19E20E21\nR45-E22E23\nEnsemble Score: 0.9056068268567102\n\nR44-ERE24E25E26\nR48-E16\nR48-E14E15E16\nR48-ER1659541967\nR41-E06\nR45-E18E19E20E21\nR45-E23\nEnsemble Score: 0.905763574072159\n\nR48-E16\nR48-E14E15E16\nR48-ER1659541967\nR45-E18E19E20E21\nR45-E23\nEnsemble Score: 0.9076220279702256\n\n\nR44-ERE24E25E26\nR48-E15\nR48-E16\nR48-E14E15E16\nR48-ER1659541967\nR45-E18E19E20E21\nR45-E23\nEnsemble Score: 0.9075854185534631\nLB_Score: 0.8776\n\n\n\nR44-E24E25E26E27\nR48-E19\nR48-E17ER00E19\nR48\nR45-E18E19E20E21\nR45-E27E28\nEnsemble Score: 0.907985807749435\n\n\nR44-E24E25E26E27\nR48-ER1659808200\nR48-E19\nR48-E17ER00E19\nR48\nR45-E18E19E20E21\nR45-E27\nEnsemble Score: 0.9100677053291213\n\n\nR44-E24E25E26E27\nR48-E17\nR48-ER1659808200\nR48-E17ER00E19\nR48\nR45-E19E20E21\nR45-E27\nEnsemble Score: 0.9100881140269717\nLB_Score: 0.8774\n\n\nR44-E28\nR45-E27\nR45-E29\nEnsemble Score: 0.8949238203272945\nLB_Score: 0.8751\n\n\nR48-E17\nR48-E19\nR48-E14\nEnsemble Score: 0.9147263588116289\nLB_Score: 0.8723\n\n\nR44-E28\nR48-E14\nR48-E17ER00E19\nR48-FoER67E14\nR45-E18E19E20E21\nR45-E27\nR45-E29\nEnsemble Score: 0.9070213389767942\n\n\n```","metadata":{}},{"cell_type":"markdown","source":"### TRN Inference","metadata":{}},{"cell_type":"code","source":"if CFG.do_inference:\n    pred_d = model_inference(model, dl_trn, device='cuda:0', add_features=False)\n\n    pred_orders, _ = get_orders_from_prediction(pred_d, ds_trn, code_rank_column=None)\n\n    gt_orders = ds_trn.get_gt_orders()\n\n    comp_score = kendall_tau(\n        gt_orders.loc[pred_orders.index.values],\n        pred_orders\n    )\n    \n    print(comp_score)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:38.985028Z","iopub.execute_input":"2022-08-09T03:48:38.987428Z","iopub.status.idle":"2022-08-09T03:48:38.994884Z","shell.execute_reply.started":"2022-08-09T03:48:38.987389Z","shell.execute_reply":"2022-08-09T03:48:38.993976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### VAL Inference","metadata":{}},{"cell_type":"code","source":"if CFG.do_inference:\n    pred_d = model_inference(model, dl_val, device='cuda:0', add_features=False)\n\n    pred_orders, _ = get_orders_from_prediction(pred_d, ds_val, code_rank_column=None)\n\n    gt_orders = ds_val.get_gt_orders()\n\n    comp_score = kendall_tau(\n        gt_orders.loc[pred_orders.index.values],\n        pred_orders\n    )\n    \n    print(comp_score)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:38.999543Z","iopub.execute_input":"2022-08-09T03:48:39.000405Z","iopub.status.idle":"2022-08-09T03:48:39.009390Z","shell.execute_reply.started":"2022-08-09T03:48:39.000368Z","shell.execute_reply":"2022-08-09T03:48:39.008418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### TST Inference","metadata":{}},{"cell_type":"code","source":"if CFG.do_inference:\n    pred_d = model_inference(model, dl_tst, device='cuda:0', add_features=False)\n\n    pred_orders, _ = get_orders_from_prediction(pred_d, ds_tst, code_rank_column=None)\n\n    gt_orders = ds_tst.get_gt_orders()\n\n    comp_score = kendall_tau(\n        gt_orders.loc[pred_orders.index.values],\n        pred_orders\n    )\n    \n    print(comp_score)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.012873Z","iopub.execute_input":"2022-08-09T03:48:39.014283Z","iopub.status.idle":"2022-08-09T03:48:39.024292Z","shell.execute_reply.started":"2022-08-09T03:48:39.014246Z","shell.execute_reply":"2022-08-09T03:48:39.023335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eval_p_matrix(dl, device='cuda:0'):\n    sample_ids_v = []\n    preds_v      = []\n    targets_v    = []\n    target_ranks_v = []\n    \n    sample_w_v = []\n    \n    p_matrix_v = []\n    pred_idx_v = []\n    target_v = []\n    nb_id_v = []\n    \n    n_code_v = []\n    n_md_v = []\n    \n    \n    loss0_v = []\n    loss1_v = []\n    loss2_v = []\n    loss3_v = []\n    \n    \n    wloss0_v = []\n    wloss1_v = []\n    wloss2_v = []\n    wloss3_v = []\n    \n    err_abs_v = []\n    err_square_v = []\n    \n    with torch.no_grad():\n        for data in tqdm(dl):\n            \n            target_ranks = data['nb_df'].sort_values('rank').index.values\n            \n            \n            inputs = data['inputs'].to(device)\n            target = data['target'].to(device)\n            f_code = data['f_code'].to(device)\n            f_md = data['f_md'].to(device)\n\n            nb_id = data['nb_id']\n            sample_ids = np.array([str(nb_id) + \"_\" + sid.split('_')[1] for sid in data['sample_id']])\n\n            y_pred = model.forward(\n                inputs,\n                return_features=False,\n                f_code=f_code,\n                f_md=f_md,\n            )\n            \n            \n            y_rank_pred = model.criterion.predict_rank_v(\n                y_pred,\n                data['f_code'].to(device),\n                data['f_md'].to(device),\n                return_order_v=False,\n            )\n\n\n            p_matrix = model.criterion.build_position_matrix(y_pred, f_code, f_md).softmax(axis=-1).detach().cpu().numpy()\n            pred_idx = p_matrix.argmax(axis=-1)\n            target = data['target'].detach().cpu().numpy()\n            targets = data['nb_df']['rank'].values\n            \n            f_code = f_code.detach().cpu().numpy()\n            f_md = f_md.detach().cpu().numpy()\n            \n            \n            if True:\n                n_code = f_code.sum()\n                n_md = f_md.sum()\n                \n                w = (\n                    np.abs(\n                        np.arange(\n                            0, n_code + 1,\n                        )[:,None] - target\n                    ) + 1.0\n                ).T\n                \n                \n                                \n\n                target_oh = np.zeros_like(p_matrix)\n                target_oh[np.arange(target.shape[0]), target] = 1.0\n\n                loss0 = (w * np.square(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                loss1 = (w * np.abs(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                \n                loss2 = (np.square(w) * np.square(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                loss3 = (np.square(w) * np.abs(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                \n                \n                for i in range(n_md):\n                    for j in range(n_md):\n                        if i != j:\n                            if target[i] < target[j]:\n                                w[i] += (np.arange(0, n_code + 1) > target[j]).astype(w.dtype)\n\n                            elif target[i] > target[j]:\n                                w[i] += (np.arange(0, n_code + 1) < target[j]).astype(w.dtype)\n                                \n                \n                wloss0 = (w * np.square(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                wloss1 = (w * np.abs(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                \n                wloss2 = (np.square(w) * np.square(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n                wloss3 = (np.square(w) * np.abs(p_matrix - target_oh)).sum(axis=-1).mean(axis=-1)\n            \n            \n            sample_w_v.append(data['w'].detach().cpu().numpy())\n            p_matrix_v.append(p_matrix)\n            pred_idx_v.append(pred_idx)\n            target_v.append(target)\n            targets_v.append(targets)\n            target_ranks_v.append(target_ranks)\n            \n            nb_id_v.append(nb_id)\n            sample_ids_v.append( sample_ids )\n            preds_v.append( y_rank_pred )\n            \n            erros_v = target - pred_idx\n            \n            err_abs_v.append( np.abs(erros_v).mean() )\n            err_square_v.append( np.square(erros_v).mean() )\n            \n            n_code_v.append( f_code.sum() )\n            n_md_v.append( f_md.sum() )\n            \n            loss0_v.append(loss0)\n            loss1_v.append(loss1)\n            loss2_v.append(loss2)\n            loss3_v.append(loss3)\n            \n            wloss0_v.append(wloss0)\n            wloss1_v.append(wloss1)\n            wloss2_v.append(wloss2)\n            wloss3_v.append(wloss3)\n            \n    \n    ret_d = {\n        'p_matrix': np.array(p_matrix_v, dtype=np.object_),\n        'pred_idx': np.array(pred_idx_v, dtype=np.object_),\n        'target': np.array(target_v, dtype=np.object_),\n        'target_ranks_v': np.array(target_ranks_v, dtype=np.object_),\n        'nb_id': np.array(nb_id_v),\n        'err_abs': np.array(err_abs_v),\n        'err_square': np.array(err_square_v),\n        'n_code':n_code_v,\n        'n_md':  n_md_v,\n        \n        'sample_w_v': np.array(sample_w_v),\n        \n        'sample_ids_v': np.concatenate(sample_ids_v),\n        'targets_v': np.concatenate(targets_v),\n        'preds_v': np.concatenate(preds_v),\n        \n        'loss0': loss0_v,\n        'loss1': loss1_v,\n        'loss2': loss2_v,\n        'loss3': loss3_v,\n        \n        'wloss0': wloss0_v,\n        'wloss1': wloss1_v,\n        'wloss2': wloss2_v,\n        'wloss3': wloss3_v,\n    }\n    \n    return ret_d","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.026870Z","iopub.execute_input":"2022-08-09T03:48:39.028540Z","iopub.status.idle":"2022-08-09T03:48:39.069130Z","shell.execute_reply.started":"2022-08-09T03:48:39.028503Z","shell.execute_reply":"2022-08-09T03:48:39.067853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    device = 'cuda:0'\n\n    torch.set_grad_enabled(False)\n    model.to(device)\n    model.eval()\n    pass","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.073888Z","iopub.execute_input":"2022-08-09T03:48:39.076033Z","iopub.status.idle":"2022-08-09T03:48:39.083447Z","shell.execute_reply.started":"2022-08-09T03:48:39.075994Z","shell.execute_reply":"2022-08-09T03:48:39.082402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    p_matrix_d = eval_p_matrix(dl_tst, device='cuda:0')\n    \n    pred_orders, _ = get_orders_from_prediction(p_matrix_d, None, code_rank_column=None)\n\n    gt_orders = pd.Series(\n        [p_matrix_d['target_ranks_v'][i].tolist() for i in range(p_matrix_d['target_ranks_v'].shape[0])],\n        [p_matrix_d['nb_id'][i] for i in range(p_matrix_d['nb_id'].shape[0])],\n    )\n\n\n    comp_score = kendall_tau(\n        gt_orders, #.loc[pred_orders.index.values],\n        pred_orders\n    )\n\n    print('comp_score:', comp_score)\n\n    scores_v = []\n    len_v = []\n\n    nb_ids = pred_orders.index.values\n    for nb_id in tqdm( nb_ids ):\n        pred_orders_df = pred_orders.loc[[nb_id]]\n\n        comp_score = kendall_tau(\n            gt_orders.loc[[nb_id]],\n            pred_orders_df,\n        )\n\n        scores_v.append(comp_score)\n        len_v.append( len(pred_orders_df.iloc[0]) )\n\n    scores_v = np.array(scores_v)\n    len_v = np.array(len_v)\n\n    results_df = pd.DataFrame(\n        {\n            'nb_ids': nb_ids,\n            'scores': scores_v,\n            'len': len_v,\n        }\n    )\n\n\n    len_squared = results_df.len.values ** 2\n    results_df['w'] = len(results_df)* len_squared/len_squared.sum()\n\n    \nif CFG.do_inference:\n    dd = {\n            'nb_ids': p_matrix_d['nb_id'],\n            'err_abs': p_matrix_d['err_abs'],\n            'err_square': p_matrix_d['err_square'],\n        \n            'sample_w_v': p_matrix_d['sample_w_v'],\n\n            'n_code': p_matrix_d['n_code'],\n            'n_md': p_matrix_d['n_md'],\n        }\n    for k in p_matrix_d.keys():\n        if 'loss' in k:\n            dd[k] = p_matrix_d[k]\n    \n    error_df = pd.DataFrame(dd)\n        \n    results_and_errors_df = results_df.merge(error_df, on='nb_ids')","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.088425Z","iopub.execute_input":"2022-08-09T03:48:39.090640Z","iopub.status.idle":"2022-08-09T03:48:39.106117Z","shell.execute_reply.started":"2022-08-09T03:48:39.090602Z","shell.execute_reply":"2022-08-09T03:48:39.105097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    save_obj(results_and_errors_df, 'sin_usar_meta_nw_results_and_errors_df.pickle')","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.111516Z","iopub.execute_input":"2022-08-09T03:48:39.114238Z","iopub.status.idle":"2022-08-09T03:48:39.120209Z","shell.execute_reply.started":"2022-08-09T03:48:39.114200Z","shell.execute_reply":"2022-08-09T03:48:39.119187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    from sklearn.linear_model import LinearRegression\n\n    x = results_and_errors_df.wloss0.values * (results_and_errors_df.sample_w_v ** 2)\n    y = results_and_errors_df.scores.values  #* results_and_errors_df.w.values\n\n\n    reg = LinearRegression().fit(x[:,None], y[:,None])\n    r2 = reg.score(x[:,None], y[:,None])\n\n\n    plt.figure(0, figsize=(15,8))\n    plt.plot(\n        x,\n        y,\n        'o',\n        label=f'r2={r2:0.02f}',\n        alpha=0.15,\n    )\n\n    xx = np.unique( x )\n    yy = reg.predict(xx[:,None])[:,0]\n    plt.plot(\n        xx,\n        yy,\n        'r-o'\n    )\n\n\n    plt.legend()\n    plt.grid()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.125441Z","iopub.execute_input":"2022-08-09T03:48:39.127671Z","iopub.status.idle":"2022-08-09T03:48:39.138511Z","shell.execute_reply.started":"2022-08-09T03:48:39.127624Z","shell.execute_reply":"2022-08-09T03:48:39.137340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    target = np.array([2, 4, 5, 0,  8])\n    n_code = 8\n    n_md = 5\n\n    w = (\n        np.abs(\n            np.arange(\n                0, n_code + 1,\n            )[:,None] - target\n        ) + 1.0\n    ).T\n\n    plt.imshow(w, cmap='Oranges')\n    plt.show()\n\n    plt.plot( wt[2], '-o')\n    plt.show()\n\n    wt = np.zeros_like(w)\n    for i in range(n_md):\n        for j in range(n_md):\n            if i != j:\n                if target[i] < target[j]:\n                    wt[i] += (np.arange(0, n_code + 1) > target[j]).astype(wt.dtype)\n\n                elif target[i] > target[j]:\n                    wt[i] += (np.arange(0, n_code + 1) < target[j]).astype(wt.dtype)\n\n\n    plt.imshow(wt, cmap='Oranges')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.140901Z","iopub.execute_input":"2022-08-09T03:48:39.141884Z","iopub.status.idle":"2022-08-09T03:48:39.151435Z","shell.execute_reply.started":"2022-08-09T03:48:39.141847Z","shell.execute_reply":"2022-08-09T03:48:39.150464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    x = results_and_errors_df.scores.values * results_and_errors_df.w.values\n    y = results_and_errors_df.err_abs.values * results_and_errors_df.len.values /500\n\n    reg = LinearRegression().fit(x[:,None], y[:,None])\n    r2 = reg.score(x[:,None], y[:,None])\n\n\n    plt.figure(0, figsize=(15,8))\n    plt.semilogy(\n        x,\n        y,\n        'o',\n        label=f'r2={r2:0.02f}',\n        alpha=0.15,\n    )\n\n    xx = np.unique( x )\n    yy = reg.predict(xx[:,None])[:,0]\n    plt.semilogy(\n        xx,\n        yy,\n        'r-o'\n    )\n\n\n    plt.legend()\n    plt.grid()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.152803Z","iopub.execute_input":"2022-08-09T03:48:39.153525Z","iopub.status.idle":"2022-08-09T03:48:39.162975Z","shell.execute_reply.started":"2022-08-09T03:48:39.153487Z","shell.execute_reply":"2022-08-09T03:48:39.162068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Models ensemble","metadata":{}},{"cell_type":"code","source":"if CFG.do_inference:\n    ds_tst = MatrixTrainDataset(\n        df=df_nb,\n        cfg=CFG,\n        tokenizer=tokenizer,\n        i_fold=None,\n        training=False,\n        md_sample_default_rank=0.5,\n        p_rnd_mask_tkn=0.0,\n    )","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.164599Z","iopub.execute_input":"2022-08-09T03:48:39.165238Z","iopub.status.idle":"2022-08-09T03:48:39.175914Z","shell.execute_reply.started":"2022-08-09T03:48:39.165202Z","shell.execute_reply":"2022-08-09T03:48:39.174734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    ens_preds_d = ensemble_preds(pred_d_v, w_v=None)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.179972Z","iopub.execute_input":"2022-08-09T03:48:39.181326Z","iopub.status.idle":"2022-08-09T03:48:39.189084Z","shell.execute_reply.started":"2022-08-09T03:48:39.181289Z","shell.execute_reply":"2022-08-09T03:48:39.187885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.do_inference:\n    pred_d = ens_preds_d\n    pred_orders, _ = get_orders_from_prediction(pred_d, ds_tst, code_rank_column=None)\n\n    gt_orders = ds_tst.get_gt_orders()\n\n    kendall_tau(\n        gt_orders.loc[pred_orders.index.values],\n        pred_orders\n    )","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.190834Z","iopub.execute_input":"2022-08-09T03:48:39.191776Z","iopub.status.idle":"2022-08-09T03:48:39.203849Z","shell.execute_reply.started":"2022-08-09T03:48:39.191740Z","shell.execute_reply":"2022-08-09T03:48:39.202668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# RUN 25\n# F0 E4 TST_score = 0.8170729055004826\n# F1 E3 TST_score = 0.817120667923368\n# F0 E6 TST_score = 0.8136965570827053\n# F1 E6 TST_score = 0.8174447700786616\n\n# Ensemble[:2] = 0.8256805185748205\n# Ensemble[:4] = 0.8301697598758404","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:48:39.205279Z","iopub.execute_input":"2022-08-09T03:48:39.205788Z","iopub.status.idle":"2022-08-09T03:48:39.211069Z","shell.execute_reply.started":"2022-08-09T03:48:39.205753Z","shell.execute_reply":"2022-08-09T03:48:39.209907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Matrix Kaggle Inference","metadata":{}},{"cell_type":"code","source":"# Matrix Kaggle Inference (Multi-Model)\nif True and CFG.do_kaggle_inference:\n    ds_tst = MatrixTrainDataset(\n        df=df_nb_tst,\n        cfg=CFG,\n        tokenizer=tokenizer,\n        i_fold=None,\n        training=False,\n        md_sample_default_rank=CFG.abs_md_sample_default_rank,\n    )\n    \n    pred_d_v = []\n    \n    # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #\n    # MODEL 44\n    model = NB_Model(\n        Args(\n            torch.load('../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/model_cfg.pth')\n        ),\n        config_path='../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/config.pth',\n        pretrained=False\n    )\n    model.eval()\n    \n    for ckpt_path in [\n\n        # Model Type 44\n        # VAST Entrenado con Metanotebooks + GitHub (semilla R43E18)\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.1955_VL1.0152_VS0.8787.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/error_recover.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL1.3469_VL1.1779_VS0.8796.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL1.3333_VL1.1726_VS0.8798.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL1.3232_VL1.1854_VS0.8799.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL1.2866_VL1.1460_VS0.8799.ckpt\",\n        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL1.2736_VL1.1491_VS0.8810.ckpt\",\n\n        \n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E25E26E27.ckpt\",\n        \n        \n        # Este se utilizó como semilla para entrenar los modelos de Vast (25 folds , Sin MNB, entrnando a partir de R42E15, GitHubDS_85kNB )\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E17_TL0.9883_VL1.0027_VS0.8653.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E16_TL1.0187_VL1.0012_VS0.8711.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E18_TL0.9703_VL1.0359_VS0.8694.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TL0.9550_VL1.0279_VS0.8666.ckpt\",\n        \n        # Nuevo entrenamiento 25 folds (Sin MNB, solamente notebooks originales de kaggle)\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TLnan_VL1.0943_VS0.8684.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E15_TLnan_VL1.0837_VS0.8692.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E13_TLnan_VL1.0509_VS0.8688.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E12_TLnan_VL1.0464_VS0.8680.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E11_TLnan_VL1.0481_VS0.8683.ckpt\",\n        \n        # VAST Entrendo con todos lo NBs y los notebooks extra de kaggle (Entrenado a partir del run_41, sinMetaNB)\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E9_TL1.0432_VL2.4768_VS0.8614.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659290334.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E11_TL0.7875_VL2.5361_VS0.8666.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E12_TL0.7707_VL2.6240_VS0.8683.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E13_TL0.7615_VL2.6432_VS0.8668.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E14_TL0.7530_VL2.6928_VS0.8701.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E15_TL0.7460_VL2.7321_VS0.8674.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E16_TL0.7405_VL2.6780_VS0.8700.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E17_TL0.7317_VL2.7011_VS0.8690.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659808200.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL0.6828_VL2.7560_VS0.8683(1).ckpt\",\n        \n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E12E13E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14ER34E11E12E13E14E15E16.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E16E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16E17ER00E19.ckpt\",\n        \n        \n        # VAST Fork de run-48 (Entrenado a aprtir de run-48-E13, conMetaNB usando GitHubNB)\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/error_recovery_1659541967.ckpt\", # No usa GHnb\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E14_TL1.3146_VL1.1827_VS0.8690.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E15_TL1.4038_VL1.1817_VS0.8627.ckpt\",\n        \n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14E15.ckpt\",\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoE14E15.ckpt\",\n        \n        # Este era el mejor modelo LB=0.7014 (entrenado conMetaNB)\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E8_TLnan_VL1.0432_VS0.8685.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E5_TLnan_VL1.0372_VS0.8682.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E6_TLnan_VL1.0529_VS0.8677.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E7_TL1.1609_VL1.0455_VS0.8672.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E3_TL1.2291_VL1.0381_VS0.8662.ckpt\",\n        \n        \n    ]:\n        print('Inference:', os.path.basename(ckpt_path))\n        model.restore_model(ckpt_path)\n        pred_d = model_inference(model, ds_tst, device='cuda:0', add_targets=False, add_features=False)    \n        pred_d_v.append(pred_d)\n        \n    # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #\n    # MODEL 45\n    model = NB_Model(\n        Args(\n            torch.load('../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/model_cfg.pth')\n        ),\n        config_path='../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/config.pth',\n        pretrained=False\n    )\n    model.eval()\n    \n    for ckpt_path in [\n        \n        # Model Type 45\n        # GitHub Training + KaggleNB (conMetaNB)\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E18_TL1.1382_VL0.9904_VS0.8766.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL1.1069_VL0.9743_VS0.8782.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E20_TL1.2697_VL1.1128_VS0.8759.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.2412_VL1.1213_VS0.8753.ckpt\",\n        \n        #'../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20.ckpt',\n        #'../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E19E20E21.ckpt',\n        '../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E18E19E20E21.ckpt',\n        \n        \n        # New Training + KaggleNB (sinMetaNB a partir de run-45 E21)\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E22_TL0.7064_VL1.8204_VS0.8831.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E23_TL0.6845_VL1.7250_VS0.8847.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL0.6709_VL1.7212_VS0.8876.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL0.6590_VL1.7388_VS0.8875.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL0.6480_VL1.7243_VS0.8874.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL0.6400_VL1.7032_VS0.8906.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL0.6327_VL1.7068_VS0.8894.ckpt\",\n        \"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E29_TL0.6240_VL1.6679_VS0.8902.ckpt\",\n        \n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E22E23E24E25.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E23E24E25.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25.ckpt\",\n        \n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E26E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E25E26E27E28.ckpt\",\n        #\"../input/run-45-matrixmdco-microsoft-codebert-base-pl-vast/run-45-MERGE_E24E25E26E27E28.ckpt\",\n        \n    ]:\n        print('Inference:', os.path.basename(ckpt_path))\n        model.restore_model(ckpt_path)\n        pred_d = model_inference(model, ds_tst, device='cuda:0', add_targets=False, add_features=False)    \n        pred_d_v.append(pred_d)\n    # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #\n        \n    \n    ens_preds_d = ensemble_preds(pred_d_v, ds_tst, w_v=None)\n    pred_orders, _ = get_orders_from_prediction(ens_preds_d, ds_tst, code_rank_column=None)\n    \n    sub_df = pd.DataFrame( pred_orders.apply(lambda x: \" \".join(x)) ).reset_index().rename(columns={\"cell_id\": \"cell_order\"})\n    sub_df.to_csv(\"submission.csv\", index=False)\n    display( sub_df )","metadata":{"execution":{"iopub.status.busy":"2022-08-10T00:04:52.982999Z","iopub.execute_input":"2022-08-10T00:04:52.986854Z","iopub.status.idle":"2022-08-10T00:05:17.644753Z","shell.execute_reply.started":"2022-08-10T00:04:52.986794Z","shell.execute_reply":"2022-08-10T00:05:17.643438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Matrix Kaggle Inference (Single model)\nif False and CFG.do_kaggle_inference:\n    ds_tst = MatrixTrainDataset(\n        df=df_nb_tst,\n        cfg=CFG,\n        tokenizer=tokenizer,\n        i_fold=None,\n        training=False,\n        md_sample_default_rank=CFG.abs_md_sample_default_rank,\n    )\n\n    model = NB_Model(\n        CFG,\n        config_path=CFG.checkpoints_dir / \"config.pth\",\n        pretrained=False\n    )\n    \n    pred_d_v = []\n    for ckpt_path in [\n        \n       # VAST Entrenado con Metanotebooks + GitHub (semilla R43E18)\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E21_TL1.1955_VL1.0152_VS0.8787.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/error_recover.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E24_TL1.3469_VL1.1779_VS0.8796.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E25_TL1.3333_VL1.1726_VS0.8798.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E26_TL1.3232_VL1.1854_VS0.8799.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E27_TL1.2866_VL1.1460_VS0.8799.ckpt\",\n#        \"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E28_TL1.2736_VL1.1491_VS0.8810.ckpt\",\n\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_ERE24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E24E25E26E27.ckpt\",\n        #\"../input/run-44-matrixmdco-microsoft-codebert-base-pl-vast/run-44-MERGE_E25E26E27.ckpt\",\n        \n        \n        # Este se utilizó como semilla para entrenar los modelos de Vast (25 folds , Sin MNB, entrnando a partir de R42E15, GitHubDS_85kNB )\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E17_TL0.9883_VL1.0027_VS0.8653.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E16_TL1.0187_VL1.0012_VS0.8711.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E18_TL0.9703_VL1.0359_VS0.8694.ckpt\",\n        #\"../input/run-43-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TL0.9550_VL1.0279_VS0.8666.ckpt\",\n        \n        # Nuevo entrenamiento 25 folds (Sin MNB, solamente notebooks originales de kaggle)\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E19_TLnan_VL1.0943_VS0.8684.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E15_TLnan_VL1.0837_VS0.8692.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E13_TLnan_VL1.0509_VS0.8688.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E12_TLnan_VL1.0464_VS0.8680.ckpt\",\n        #\"../input/run-42-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E11_TLnan_VL1.0481_VS0.8683.ckpt\",\n        \n        # VAST Entrendo con todos lo NBs y los notebooks extra de kaggle (Entrenado a partir del run_41, sinMetaNB)\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E9_TL1.0432_VL2.4768_VS0.8614.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659290334.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E11_TL0.7875_VL2.5361_VS0.8666.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E12_TL0.7707_VL2.6240_VS0.8683.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E13_TL0.7615_VL2.6432_VS0.8668.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E14_TL0.7530_VL2.6928_VS0.8701.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E15_TL0.7460_VL2.7321_VS0.8674.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E16_TL0.7405_VL2.6780_VS0.8700.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E17_TL0.7317_VL2.7011_VS0.8690.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/error_recovery_1659808200.ckpt\",\n        \"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/microsoft-codebert-base_F0_E19_TL0.6828_VL2.7560_VS0.8683(1).ckpt\",\n        \n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E12E13E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E14ER34E11E12E13E14E15E16.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E16E17ER00E19.ckpt\",\n        #\"../input/run-48-matrixmdco-microsoft-codebert-base-pl-vast/run-48-MERGE_E15E16E17ER00E19.ckpt\",\n        \n        \n        # VAST Fork de run-48 (Entrenado a aprtir de run-48-E13, conMetaNB usando GitHubNB)\n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/error_recovery_1659541967.ckpt\", # No usa GHnb\n        \"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/microsoft-codebert-base_F0_E14_TL1.3146_VL1.1827_VS0.8690.ckpt\",\n        \n        #\"../input/run-48-matrixmdco-codebert-pl-forkmn-ghnb-vast/run-48-MERGE_FoER67E14.ckpt\",\n        \n        # Este era el mejor modelo LB=0.7014 (entrenado conMetaNB)\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E8_TLnan_VL1.0432_VS0.8685.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E5_TLnan_VL1.0372_VS0.8682.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E6_TLnan_VL1.0529_VS0.8677.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E7_TL1.1609_VL1.0455_VS0.8672.ckpt\",\n        #\"../input/run-41-matrixmdco-microsoft-codebert-base-pl/microsoft-codebert-base_F0_E3_TL1.2291_VL1.0381_VS0.8662.ckpt\",\n    ]:\n        print('Inference:', os.path.basename(ckpt_path))\n        model.restore_model(ckpt_path)\n        pred_d = model_inference(model, ds_tst, device='cuda:0', add_targets=False, add_features=False)    \n        pred_d_v.append(pred_d)\n        \n    \n    ens_preds_d = ensemble_preds(pred_d_v, ds_tst, w_v=None)\n    pred_orders, _ = get_orders_from_prediction(ens_preds_d, ds_tst, code_rank_column=None)\n    \n    sub_df = pd.DataFrame( pred_orders.apply(lambda x: \" \".join(x)) ).reset_index().rename(columns={\"cell_id\": \"cell_order\"})\n    sub_df.to_csv(\"submission.csv\", index=False)\n    display( sub_df )","metadata":{"execution":{"iopub.status.busy":"2022-08-10T04:26:36.974660Z","iopub.execute_input":"2022-08-10T04:26:36.975120Z","iopub.status.idle":"2022-08-10T04:27:01.927631Z","shell.execute_reply.started":"2022-08-10T04:26:36.975077Z","shell.execute_reply":"2022-08-10T04:27:01.926542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}