{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"I am referring to [@ragnar123](https://www.kaggle.com/ragnar123)'s awesome [notebook](https://www.kaggle.com/code/ragnar123/unsupervised-baseline-arcface) on [shopee competition](https://www.kaggle.com/competitions/shopee-product-matching).\n\nSince the design of this competition was similar to that of the shopee competition, I tried metric learning, which is the top solution in \nthe shopee competition.\n\nHowever, unlike the shopee competition, the most important information in this competition is `latitude` and `longitude`, not textual information, so I tried metric learning model in the 1st stage and GBDT model with 1st stage features in the 2nd stage.\n\n**about this notebook**\n\nOn both Kaggle and Colab, training and inference can be run on this single notebook!\n\n1. Training\n    \n    Set `CFG.train = True` and run.\n\n2. Inference\n\n    Set `CFG.train = False` and run.\n\n\n","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"id":"97e12bc7","outputId":"f9132382-3b28-43d4-e7af-97806eb7e771","papermill":{"duration":0.727074,"end_time":"2022-07-03T16:49:24.482219","exception":false,"start_time":"2022-07-03T16:49:23.755145","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:40:53.562358Z","iopub.execute_input":"2022-07-09T18:40:53.562759Z","iopub.status.idle":"2022-07-09T18:40:54.298787Z","shell.execute_reply.started":"2022-07-09T18:40:53.562684Z","shell.execute_reply":"2022-07-09T18:40:54.297610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{"id":"1e0d427f","papermill":{"duration":0.010656,"end_time":"2022-07-03T16:49:24.504141","exception":false,"start_time":"2022-07-03T16:49:24.493485","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# import libraries1\n# ====================================================\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport os\nimport sys\nimport math\nimport random\nimport time\nimport numpy as np\nimport pandas as pd\nimport gc\nimport json\nimport joblib\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport itertools\nimport collections\nfrom collections import Counter\n\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau, _LRScheduler\nfrom torch.nn import Parameter\n\nimport datetime\nfrom datetime import timedelta\nimport hashlib\nimport difflib\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom requests import get\nfrom PIL import Image\nimport pickle\nfrom contextlib import contextmanager\nimport multiprocessing\n\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\nfrom sklearn.neighbors import KNeighborsRegressor, NearestNeighbors\nfrom sklearn.metrics import mean_squared_error, f1_score\nfrom sklearn.linear_model import RidgeCV\nfrom sklearn.base import BaseEstimator, TransformerMixin\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.utils.class_weight import compute_sample_weight\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom sklearn.decomposition import TruncatedSVD\nfrom sklearn.feature_extraction.text import TfidfVectorizer\nfrom gensim.models import word2vec\n\nimport lightgbm as lgb\nimport typing as tp\n\nfrom logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n\ntqdm.pandas()\npd.set_option('display.max_rows', 500)\npd.set_option('display.max_columns', 500)\n\nif torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n    \nprint(f'Using device: {device}')","metadata":{"id":"6f9da346","outputId":"6a7d26d5-c812-4c3b-decd-df8546a7e044","papermill":{"duration":5.13333,"end_time":"2022-07-03T16:49:29.648189","exception":false,"start_time":"2022-07-03T16:49:24.514859","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:40:54.301084Z","iopub.execute_input":"2022-07-09T18:40:54.301881Z","iopub.status.idle":"2022-07-09T18:40:58.898921Z","shell.execute_reply.started":"2022-07-09T18:40:54.301839Z","shell.execute_reply":"2022-07-09T18:40:58.898004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"id":"88f91350","papermill":{"duration":0.011348,"end_time":"2022-07-03T16:49:29.671345","exception":false,"start_time":"2022-07-03T16:49:29.659997","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG:\n    # ====================================================\n    # Basic Setting\n    # ====================================================\n    colab = \"google.colab\" in sys.modules\n    exp = \"130\"\n    train = False\n    api_path = '/content/drive/My Drive/kaggle.json'\n    seed = 42\n    n_neighbors = 50\n    threshold = 0.15\n    # ====================================================\n    # Model\n    # ====================================================\n    model = \"sentence-transformers/paraphrase-multilingual-mpnet-base-v2\"\n    epochs = 15\n    num_workers = 8\n    batch_size = 32\n    max_length = 32\n    lr = 1e-5\n    scheduler= 'linear' # ['linear', 'cosine', 'ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts', 'CosineAnnealingWarmupRestarts']\n\nif not CFG.colab:\n    CFG.model = \"../input/sbert-models/paraphrase-multilingual-mpnet-base-v2\"","metadata":{"id":"c601dce5","papermill":{"duration":0.021127,"end_time":"2022-07-03T16:49:29.703876","exception":false,"start_time":"2022-07-03T16:49:29.682749","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:40:58.900369Z","iopub.execute_input":"2022-07-09T18:40:58.900971Z","iopub.status.idle":"2022-07-09T18:40:58.908076Z","shell.execute_reply.started":"2022-07-09T18:40:58.900913Z","shell.execute_reply":"2022-07-09T18:40:58.907062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.colab:\n    print(\"==============================================\")\n    print(\"This environment is Google Colab\")\n    print(\"==============================================\")\n\n    # Google Drive\n    from google.colab import drive, files\n    drive.mount('/content/drive')\n    %cd \"drive/My Drive/foursquare/\"\n\n    # Kaggle API\n    f = open(CFG.api_path, 'r')\n    json_data = json.load(f) \n    os.environ[\"KAGGLE_USERNAME\"] = json_data[\"username\"]\n    os.environ[\"KAGGLE_KEY\"] = json_data[\"key\"]\n\n    # Directory Setting\n    if not os.path.exists(f\"output/exp{CFG.exp}/\"):\n        os.makedirs(f\"output/exp{CFG.exp}/\")\n    \n    DATA_DIR = \"input/\"\n    OUTPUT_DIR = f\"output/exp{CFG.exp}/\"\n    MODEL_DIR = OUTPUT_DIR\n    \n    # Data Loading\n    if not os.path.isfile(os.path.join(DATA_DIR, \"foursquare-location-matching.zip\")):\n        !kaggle competitions download -c foursquare-location-matching -p $DATA_DIR\n\n    # Libraries\n    #!pip install -q catboost\n    #!pip install -q Levenshtein\n    #!pip install -q textdistance==4.2.2\n    #!pip install -q pylcs==0.0.6\n    #!pip install -q fasttext\n    !pip install -q reverse_geocode\n    !pip install -q transformers\n    !pip install -q sentence_transformers==2.2.0\n\nelse:\n    print(\"==============================================\")\n    print(\" This environment is Kaggle Notebook\")\n    print(\"==============================================\")\n\n    # Directory Setting\n    DATA_DIR = \"../input/foursquare-location-matching/\"\n    OUTPUT_DIR = \"./\"\n    MODEL_DIR = f\"../input/foursquare-stage1-exp{CFG.exp}/\"\n\n    # Libraries\n    !pip install /kaggle/input/reversegeocode/reverse_geocode-1.4.1-py3-none-any.whl\n    #!pip install ../input/textdistance-install/textdistance-4.2.2-py3-none-any.whl\n    #!pip install --force-reinstall ../input/pylcs-install/pybind11-2.9.2-py2.py3-none-any.whl\n\n    #!rm -r mypip\n    #!mkdir mypip\n    #!tar -czvf mypip/pylcs-0.0.6.tar.gz -C ../input/pylcs-install/pylcs-0.0.6/pylcs-0.0.6 .\n    #!ls -l mypip\n\n    #!pip install --no-index mypip/pylcs-0.0.6.tar.gz\n    \n    sys.path.append(\"../input/sentencetransformersinstall/sentence-transformers-2.2.0\")\n\n# ====================================================\n# import libraries2\n# ====================================================\n\n#from catboost import CatBoost, Pool\n#import Levenshtein\n#import pylcs\n#import textdistance\nimport reverse_geocode\n#from fasttext import load_model\nfrom transformers import DistilBertModel, DistilBertTokenizer, AutoTokenizer, AutoModel, AutoConfig\nfrom transformers import AdamW\nfrom transformers import get_linear_schedule_with_warmup,get_cosine_schedule_with_warmup\nfrom transformers import get_cosine_with_hard_restarts_schedule_with_warmup","metadata":{"id":"43929680","outputId":"1a889400-65dc-4e31-f51e-be821bb27cb4","papermill":{"duration":88.780878,"end_time":"2022-07-03T16:50:58.495751","exception":false,"start_time":"2022-07-03T16:49:29.714873","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:40:58.911210Z","iopub.execute_input":"2022-07-09T18:40:58.911642Z","iopub.status.idle":"2022-07-09T18:41:34.534505Z","shell.execute_reply.started":"2022-07-09T18:40:58.911606Z","shell.execute_reply":"2022-07-09T18:41:34.533650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper Functions","metadata":{"id":"7850cddb","papermill":{"duration":0.013064,"end_time":"2022-07-03T16:50:58.521788","exception":false,"start_time":"2022-07-03T16:50:58.508724","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def reduce_mem_usage(df, verbose=True):\n    numerics = ['int16', 'int32', 'int64', 'float16', 'float32', 'float64']\n    start_mem = df.memory_usage().sum() / 1024**2    \n    for col in df.columns:\n        col_type = df[col].dtypes\n        if col_type in numerics:\n            c_min = df[col].min()\n            c_max = df[col].max()\n            if str(col_type)[:3] == 'int':\n                if c_min > np.iinfo(np.int8).min and c_max < np.iinfo(np.int8).max:\n                    df[col] = df[col].astype(np.int8)\n                elif c_min > np.iinfo(np.int16).min and c_max < np.iinfo(np.int16).max:\n                    df[col] = df[col].astype(np.int16)\n                elif c_min > np.iinfo(np.int32).min and c_max < np.iinfo(np.int32).max:\n                    df[col] = df[col].astype(np.int32)\n                elif c_min > np.iinfo(np.int64).min and c_max < np.iinfo(np.int64).max:\n                    df[col] = df[col].astype(np.int64)  \n            else:\n                if c_min > np.finfo(np.float16).min and c_max < np.finfo(np.float16).max:\n                    df[col] = df[col].astype(np.float16)\n                elif c_min > np.finfo(np.float32).min and c_max < np.finfo(np.float32).max:\n                    df[col] = df[col].astype(np.float32)\n                else:\n                    df[col] = df[col].astype(np.float64)    \n    end_mem = df.memory_usage().sum() / 1024**2\n    if verbose: print('Memory usage decreased to {:5.2f} Mb ({:.1f}% reduction)'.format(end_mem, 100 * (start_mem - end_mem) / start_mem))\n    return df","metadata":{"id":"0593c4ad","papermill":{"duration":0.027586,"end_time":"2022-07-03T16:50:58.562367","exception":false,"start_time":"2022-07-03T16:50:58.534781","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.536189Z","iopub.execute_input":"2022-07-09T18:41:34.536575Z","iopub.status.idle":"2022-07-09T18:41:34.550606Z","shell.execute_reply.started":"2022-07-09T18:41:34.536536Z","shell.execute_reply":"2022-07-09T18:41:34.549621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@contextmanager\ndef timer(name: str):\n    t0 = time.time()\n    print(f\"[{name}] start\")\n    yield\n    msg = f\"[{name}] done in {time.time() - t0:.0f} s\"\n    print(msg)","metadata":{"id":"376681d2","papermill":{"duration":0.019259,"end_time":"2022-07-03T16:50:58.594017","exception":false,"start_time":"2022-07-03T16:50:58.574758","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.552516Z","iopub.execute_input":"2022-07-09T18:41:34.553218Z","iopub.status.idle":"2022-07-09T18:41:34.567056Z","shell.execute_reply.started":"2022-07-09T18:41:34.553181Z","shell.execute_reply":"2022-07-09T18:41:34.566254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nseed_everything(CFG.seed)","metadata":{"id":"dcd5b647","papermill":{"duration":0.027127,"end_time":"2022-07-03T16:50:58.634063","exception":false,"start_time":"2022-07-03T16:50:58.606936","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:44:09.049214Z","iopub.execute_input":"2022-07-09T18:44:09.049776Z","iopub.status.idle":"2022-07-09T18:44:09.059297Z","shell.execute_reply.started":"2022-07-09T18:44:09.049734Z","shell.execute_reply":"2022-07-09T18:44:09.058453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_worker(worker_id):\n    worker_seed = torch.initial_seed() % 2**32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\ng = torch.Generator()\ng.manual_seed(CFG.seed)","metadata":{"id":"6a61b154","outputId":"1a3bb24e-ce6e-478e-9647-a0dda94e224b","papermill":{"duration":0.023754,"end_time":"2022-07-03T16:50:58.67242","exception":false,"start_time":"2022-07-03T16:50:58.648666","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.581330Z","iopub.execute_input":"2022-07-09T18:41:34.581843Z","iopub.status.idle":"2022-07-09T18:41:34.591736Z","shell.execute_reply.started":"2022-07-09T18:41:34.581806Z","shell.execute_reply":"2022-07-09T18:41:34.590859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_logger(log_file=OUTPUT_DIR+'train.log'):\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    logger.hasHandlers()\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()","metadata":{"id":"5bda4708","papermill":{"duration":0.021293,"end_time":"2022-07-03T16:50:58.706216","exception":false,"start_time":"2022-07-03T16:50:58.684923","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.593310Z","iopub.execute_input":"2022-07-09T18:41:34.593771Z","iopub.status.idle":"2022-07-09T18:41:34.601436Z","shell.execute_reply.started":"2022-07-09T18:41:34.593679Z","shell.execute_reply":"2022-07-09T18:41:34.600480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_id2poi(input_df: pd.DataFrame) -> dict:\n    return dict(zip(input_df['id'], input_df['point_of_interest']))\n\ndef get_poi2ids(input_df: pd.DataFrame) -> dict:\n    return input_df.groupby('point_of_interest')['id'].apply(set).to_dict()\n\ndef get_score(input_df: pd.DataFrame):\n    scores = []\n    for id_str, matches in zip(input_df['id'].to_numpy(), input_df['matches'].to_numpy()):\n        targets = poi2ids[id2poi[id_str]]\n        preds = set(matches.split())\n        score = len((targets & preds)) / len((targets | preds))\n        scores.append(score)\n    scores = np.array(scores)\n    return scores.mean()\n\ndef analysis(df):\n    print('Num of data: %s' % len(df))\n    print('Num of unique id: %s' % df['id'].nunique())\n    print('Num of unique poi: %s' % df['point_of_interest'].nunique())\n    \n    poi_grouped = df.groupby('point_of_interest')['id'].count().reset_index()\n    print('Mean num of unique poi: %s' % poi_grouped['id'].mean())","metadata":{"id":"f7e541ae","papermill":{"duration":0.023644,"end_time":"2022-07-03T16:50:58.743273","exception":false,"start_time":"2022-07-03T16:50:58.719629","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.605708Z","iopub.execute_input":"2022-07-09T18:41:34.606298Z","iopub.status.idle":"2022-07-09T18:41:34.616204Z","shell.execute_reply.started":"2022-07-09T18:41:34.606267Z","shell.execute_reply":"2022-07-09T18:41:34.615287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cos_sim(v1, v2):\n    return np.dot(v1, v2) / (np.linalg.norm(v1) * np.linalg.norm(v2))","metadata":{"id":"262b5dce","papermill":{"duration":0.019401,"end_time":"2022-07-03T16:50:58.77505","exception":false,"start_time":"2022-07-03T16:50:58.755649","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.618085Z","iopub.execute_input":"2022-07-09T18:41:34.618542Z","iopub.status.idle":"2022-07-09T18:41:34.624822Z","shell.execute_reply.started":"2022-07-09T18:41:34.618500Z","shell.execute_reply":"2022-07-09T18:41:34.624098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"id":"5084ce59","papermill":{"duration":0.014484,"end_time":"2022-07-03T16:50:58.837003","exception":false,"start_time":"2022-07-03T16:50:58.822519","status":"completed"},"tags":[]}},{"cell_type":"code","source":"with timer(\"Data Loading\"):\n    if CFG.train:\n        original_df = pd.read_csv(DATA_DIR + \"train.csv\")\n    else:\n        original_df = pd.read_csv(DATA_DIR + \"test.csv\")\n        original_df[\"point_of_interest\"] = \"match\"\ndisplay(original_df)","metadata":{"id":"792f43b4","outputId":"7062f96b-5789-4001-bd31-8467c4eef701","papermill":{"duration":0.05167,"end_time":"2022-07-03T16:50:58.901998","exception":false,"start_time":"2022-07-03T16:50:58.850328","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.626260Z","iopub.execute_input":"2022-07-09T18:41:34.626780Z","iopub.status.idle":"2022-07-09T18:41:34.664960Z","shell.execute_reply.started":"2022-07-09T18:41:34.626742Z","shell.execute_reply":"2022-07-09T18:41:34.663995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id2poi = get_id2poi(original_df)\npoi2ids = get_poi2ids(original_df)","metadata":{"id":"be34346d","papermill":{"duration":0.02529,"end_time":"2022-07-03T16:50:58.972273","exception":false,"start_time":"2022-07-03T16:50:58.946983","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.666372Z","iopub.execute_input":"2022-07-09T18:41:34.667014Z","iopub.status.idle":"2022-07-09T18:41:34.675779Z","shell.execute_reply.started":"2022-07-09T18:41:34.666975Z","shell.execute_reply":"2022-07-09T18:41:34.674816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"le = LabelEncoder()\noriginal_df['point_of_interest'] = le.fit_transform(original_df['point_of_interest'])\noriginal_df['point_of_interest'] = original_df['point_of_interest'].astype(\"int32\")\n\nn_classes = original_df[\"point_of_interest\"].nunique()\nprint(f\"n_classes: {n_classes}\")","metadata":{"id":"f8a84799","outputId":"16c88d9d-55f4-4b44-a392-c3b94eca3c12","papermill":{"duration":0.029765,"end_time":"2022-07-03T16:50:59.016518","exception":false,"start_time":"2022-07-03T16:50:58.986753","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.678622Z","iopub.execute_input":"2022-07-09T18:41:34.678950Z","iopub.status.idle":"2022-07-09T18:41:34.689939Z","shell.execute_reply.started":"2022-07-09T18:41:34.678894Z","shell.execute_reply":"2022-07-09T18:41:34.689174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocess","metadata":{"id":"63d3079f","papermill":{"duration":0.012665,"end_time":"2022-07-03T16:50:59.043229","exception":false,"start_time":"2022-07-03T16:50:59.030564","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"`reverse_geocode` library is used to determine city information from `latitude` and `longitude`.\n\nBy including this city information in the model, it is intended that `latitude` and `longitude` information is also taken into account for BERT model.","metadata":{"id":"dec9f22f","papermill":{"duration":0.012516,"end_time":"2022-07-03T16:50:59.068666","exception":false,"start_time":"2022-07-03T16:50:59.05615","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_geo_info(coords):\n    data = reverse_geocode.search(coords)\n    return [v['country_code'] for v in data], [v['city'] for v in data]\n\noriginal_df['city2'] = get_geo_info(original_df[['latitude', 'longitude']])[1]","metadata":{"id":"59c041f8","papermill":{"duration":0.471752,"end_time":"2022-07-03T16:50:59.553406","exception":false,"start_time":"2022-07-03T16:50:59.081654","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:34.691194Z","iopub.execute_input":"2022-07-09T18:41:34.691609Z","iopub.status.idle":"2022-07-09T18:41:35.130644Z","shell.execute_reply.started":"2022-07-09T18:41:34.691573Z","shell.execute_reply":"2022-07-09T18:41:35.129829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"original_df[\"text\"] = original_df[\"name\"].fillna(\"\") + \" \" +\\\n                        original_df[\"city2\"].fillna(\"\") + \" \" +\\\n                        original_df[\"address\"].fillna(\"\") + \" \" +\\\n                        original_df[\"categories\"].fillna(\"\")\n\n#original_df[\"text\"] = original_df[\"name\"].fillna(\"\") + \"[SEP]\" +\\\n#                        original_df[\"city2\"].fillna(\"\") + \"[SEP]\" +\\\n#                        original_df[\"address\"].fillna(\"\") + \"[SEP]\" +\\\n#                        original_df[\"categories\"].fillna(\"\")","metadata":{"id":"b9cde7b9","papermill":{"duration":0.023332,"end_time":"2022-07-03T16:50:59.590193","exception":false,"start_time":"2022-07-03T16:50:59.566861","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.132093Z","iopub.execute_input":"2022-07-09T18:41:35.132422Z","iopub.status.idle":"2022-07-09T18:41:35.140569Z","shell.execute_reply.started":"2022-07-09T18:41:35.132385Z","shell.execute_reply":"2022-07-09T18:41:35.139612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ==============================================\n# Degree to radian\n# ==============================================\n\noriginal_df[\"latitude\"] = original_df[\"latitude\"] * np.pi / 180\noriginal_df[\"longitude\"] = original_df[\"longitude\"] * np.pi / 180","metadata":{"id":"7b0d052a","papermill":{"duration":0.023988,"end_time":"2022-07-03T16:50:59.629589","exception":false,"start_time":"2022-07-03T16:50:59.605601","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.142466Z","iopub.execute_input":"2022-07-09T18:41:35.142796Z","iopub.status.idle":"2022-07-09T18:41:35.151039Z","shell.execute_reply.started":"2022-07-09T18:41:35.142764Z","shell.execute_reply":"2022-07-09T18:41:35.150199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"699ff55e","papermill":{"duration":0.013029,"end_time":"2022-07-03T16:50:59.658293","exception":false,"start_time":"2022-07-03T16:50:59.645264","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n#  Dataset\n# ====================================================\n\nclass FoursquareDataset(Dataset):\n    def __init__(self, df, include_labels=True):\n        tokenizer = AutoTokenizer.from_pretrained(CFG.model)\n\n        self.df = df\n        self.include_labels = include_labels\n\n        self.text = df['text'].tolist()\n        self.lat = df['latitude'].values\n        self.lon = df['longitude'].values\n        self.labels = df['point_of_interest'].values\n\n        self.encoded = tokenizer.batch_encode_plus(\n            self.text,\n            padding = 'max_length',            \n            max_length = CFG.max_length,\n            truncation = True,\n            return_attention_mask=True\n        )\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n\n        input_ids = torch.tensor(self.encoded['input_ids'][idx], dtype=torch.long)\n        attention_mask = torch.tensor(self.encoded['attention_mask'][idx], dtype=torch.long)\n        lat = torch.tensor(self.lat[idx], dtype=torch.float)\n        lon = torch.tensor(self.lon[idx], dtype=torch.float)\n\n        if self.include_labels:\n            label = torch.tensor(self.labels[idx], dtype=torch.long)\n            return input_ids, attention_mask, lat, lon, label\n\n        return input_ids, attention_mask, lat, lon","metadata":{"id":"7e45a1a5","papermill":{"duration":0.028595,"end_time":"2022-07-03T16:50:59.702132","exception":false,"start_time":"2022-07-03T16:50:59.673537","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.152108Z","iopub.execute_input":"2022-07-09T18:41:35.152546Z","iopub.status.idle":"2022-07-09T18:41:35.165488Z","shell.execute_reply.started":"2022-07-09T18:41:35.152510Z","shell.execute_reply":"2022-07-09T18:41:35.164798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metric Learning","metadata":{"id":"0bb852a6","papermill":{"duration":0.013392,"end_time":"2022-07-03T16:50:59.728924","exception":false,"start_time":"2022-07-03T16:50:59.715532","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n#  CurricularFace\n# ====================================================   \n\ndef l2_norm(input, axis = 1):\n    norm = torch.norm(input, 2, axis, True)\n    output = torch.div(input, norm)\n\n    return output\n\nclass CurricularFace(nn.Module):\n    def __init__(self, in_features, out_features, s = 5, m = 0.050):\n        super(CurricularFace, self).__init__()\n\n        self.in_features = in_features\n        self.out_features = out_features\n        self.m = m\n        self.s = s\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.threshold = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n        self.kernel = nn.Parameter(torch.Tensor(in_features, out_features))\n        self.register_buffer('t', torch.zeros(1))\n        nn.init.normal_(self.kernel, std=0.01)\n\n    def forward(self, embbedings, label):\n        embbedings = l2_norm(embbedings, axis = 1)\n        kernel_norm = l2_norm(self.kernel, axis = 0)\n        cos_theta = torch.mm(embbedings, kernel_norm)\n        cos_theta = cos_theta.clamp(-1, 1)\n        with torch.no_grad():\n            origin_cos = cos_theta.clone()\n        target_logit = cos_theta[torch.arange(0, embbedings.size(0)), label].view(-1, 1)\n\n        sin_theta = torch.sqrt(1.0 - torch.pow(target_logit, 2))\n        cos_theta_m = target_logit * self.cos_m - sin_theta * self.sin_m\n        mask = cos_theta > cos_theta_m\n        final_target_logit = torch.where(target_logit > self.threshold, cos_theta_m, target_logit - self.mm)\n\n        hard_example = cos_theta[mask]\n        with torch.no_grad():\n            self.t = target_logit.mean() * 0.01 + (1 - 0.01) * self.t\n        cos_theta[mask] = hard_example * (self.t + hard_example)\n        cos_theta.scatter_(1, label.view(-1, 1).long(), final_target_logit)\n        output = cos_theta * self.s\n        return output","metadata":{"id":"edb59b62","papermill":{"duration":0.02847,"end_time":"2022-07-03T16:50:59.770409","exception":false,"start_time":"2022-07-03T16:50:59.741939","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.168904Z","iopub.execute_input":"2022-07-09T18:41:35.169187Z","iopub.status.idle":"2022-07-09T18:41:35.185515Z","shell.execute_reply.started":"2022-07-09T18:41:35.169146Z","shell.execute_reply":"2022-07-09T18:41:35.184799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n#  ArcFace\n# ====================================================\n\nclass ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, s=10.0, m=0.050, easy_margin=True, ls_eps=0.0):\n        super(ArcMarginProduct, self).__init__()\n\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps\n        self.weight = Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, input, label):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        one_hot = torch.zeros(cosine.size(), device=device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n\n        return output","metadata":{"id":"d4c25274","papermill":{"duration":0.026176,"end_time":"2022-07-03T16:50:59.809374","exception":false,"start_time":"2022-07-03T16:50:59.783198","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.186616Z","iopub.execute_input":"2022-07-09T18:41:35.186987Z","iopub.status.idle":"2022-07-09T18:41:35.201064Z","shell.execute_reply.started":"2022-07-09T18:41:35.186950Z","shell.execute_reply":"2022-07-09T18:41:35.200222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{"id":"8c2623dc","papermill":{"duration":0.012599,"end_time":"2022-07-03T16:50:59.834927","exception":false,"start_time":"2022-07-03T16:50:59.822328","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n#  Focal Loss\n# ====================================================    \n\nclass FocalLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(FocalLoss, self).__init__()\n\n    def forward(self, inputs, targets, alpha=0.8, gamma=2, smooth=1):\n        \n        inputs = F.sigmoid(inputs)       \n        \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        BCE = F.binary_cross_entropy(inputs, targets, reduction='mean')\n        BCE_EXP = torch.exp(-BCE)\n        focal_loss = alpha * (1-BCE_EXP)**gamma * BCE\n                       \n        return focal_loss","metadata":{"id":"00337d45","papermill":{"duration":0.02277,"end_time":"2022-07-03T16:50:59.870501","exception":false,"start_time":"2022-07-03T16:50:59.847731","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.202155Z","iopub.execute_input":"2022-07-09T18:41:35.202537Z","iopub.status.idle":"2022-07-09T18:41:35.220785Z","shell.execute_reply.started":"2022-07-09T18:41:35.202492Z","shell.execute_reply":"2022-07-09T18:41:35.219904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"01491cb7","papermill":{"duration":0.013187,"end_time":"2022-07-03T16:50:59.954083","exception":false,"start_time":"2022-07-03T16:50:59.940896","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n#  Model\n# ==================================================== \n\nclass CustomModel(nn.Module):\n    def __init__(self, model_name, embedding_size=128):                 \n        super(CustomModel, self).__init__()\n\n        self.config = AutoConfig.from_pretrained(model_name)\n        self.model = AutoModel.from_pretrained(model_name, \n                                               config=self.config)\n        #self.fc = ArcMarginProduct(embedding_size, n_classes)\n        self.fc = CurricularFace(embedding_size, 739972)\n        self.head = nn.Sequential(\n            nn.Linear(self.config.hidden_size + 2, embedding_size),\n            nn.BatchNorm1d(embedding_size),\n        )\n\n    def forward(self, ids, mask, lat, lon, labels):\n        embedding = self.extract(ids=ids, mask=mask, lat=lat, lon=lon)\n        output = self.fc(embedding, labels)\n        return output\n    \n    def extract(self, ids, mask, lat, lon):\n        lat, lon = lat.view(-1, 1), lon.view(-1, 1)\n        out = self.model(input_ids=ids, attention_mask=mask)\n        embedding = out[0][:, 0, :] # CLS Token\n        embedding = torch.cat([embedding, lat, lon], axis=1)\n        embedding = self.head(embedding)\n        return embedding\n    \nprint(CustomModel(CFG.model))","metadata":{"id":"4ac89dbd","outputId":"1102042d-1d3e-4b40-d214-5544accdc5a1","papermill":{"duration":12.217573,"end_time":"2022-07-03T16:51:12.184583","exception":false,"start_time":"2022-07-03T16:50:59.96701","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:41:35.223840Z","iopub.execute_input":"2022-07-09T18:41:35.224186Z","iopub.status.idle":"2022-07-09T18:42:03.545758Z","shell.execute_reply.started":"2022-07-09T18:41:35.224158Z","shell.execute_reply":"2022-07-09T18:42:03.544901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{"id":"da91323f","papermill":{"duration":0.01373,"end_time":"2022-07-03T16:51:12.213339","exception":false,"start_time":"2022-07-03T16:51:12.199609","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_accuracy(preds, targets):\n    preds = preds.argmax(dim=1)\n    acc = (preds == targets).float().mean()\n    return acc","metadata":{"id":"72ef6f3d","papermill":{"duration":0.020797,"end_time":"2022-07-03T16:51:12.247376","exception":false,"start_time":"2022-07-03T16:51:12.226579","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.547103Z","iopub.execute_input":"2022-07-09T18:42:03.547627Z","iopub.status.idle":"2022-07-09T18:42:03.552391Z","shell.execute_reply.started":"2022-07-09T18:42:03.547588Z","shell.execute_reply":"2022-07-09T18:42:03.551556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    \n    start = end = time.time()\n    losses = AverageMeter()\n    scores = AverageMeter()\n\n    model.train()\n\n    for step, (input_ids, attention_mask, lat, lon, label) in enumerate(train_loader):\n        input_ids = input_ids.to(device, dtype=torch.long)\n        attention_mask = attention_mask.to(device, dtype=torch.long)\n        lat = lat.to(device)\n        lon = lon.to(device)\n        label = label.to(device, dtype=torch.long)\n\n        batch_size = label.size(0)\n\n        output = model(input_ids, attention_mask, lat, lon, label)\n        loss = criterion(output, label)\n\n        # record loss\n        losses.update(loss.item(), batch_size)\n        loss.backward()\n\n        optimizer.step()\n        optimizer.zero_grad()\n\n        # score\n        score = get_accuracy(output.detach(), label)\n        scores.update(score.item(), batch_size)\n\n        # step\n        scheduler.step()\n                \n        if CFG.scheduler=='ReduceLROnPlateau':\n            lr = optimizer.param_groups[0]['lr']\n        else:\n            lr = scheduler.get_lr()[0]\n\n        if step % 1000 == 0 or step == (len(train_loader) - 1):\n            LOGGER.info(\n                f\"Epoch: [{epoch + 1}][{step}/{len(train_loader)}] \"\n                f\"Elapsed {timeSince(start, float(step + 1) / len(train_loader)):s} \"\n                f\"Loss: {losses.avg:.6f} \"\n                f\"Score: {scores.avg:.6f} \"\n                f\"LR: {lr:.8f} \"\n            )\n\n    return losses.avg","metadata":{"id":"a99c010e","papermill":{"duration":0.026185,"end_time":"2022-07-03T16:51:12.286823","exception":false,"start_time":"2022-07-03T16:51:12.260638","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.553832Z","iopub.execute_input":"2022-07-09T18:42:03.554312Z","iopub.status.idle":"2022-07-09T18:42:03.567359Z","shell.execute_reply.started":"2022-07-09T18:42:03.554278Z","shell.execute_reply":"2022-07-09T18:42:03.566609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# helper function\n# ====================================================\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return \"%dm %ds\" % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return \"%s (remain %s)\" % (asMinutes(s), asMinutes(rs))","metadata":{"id":"4359af1e","papermill":{"duration":0.02342,"end_time":"2022-07-03T16:51:12.323385","exception":false,"start_time":"2022-07-03T16:51:12.299965","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.568679Z","iopub.execute_input":"2022-07-09T18:42:03.569251Z","iopub.status.idle":"2022-07-09T18:42:03.578864Z","shell.execute_reply.started":"2022-07-09T18:42:03.569204Z","shell.execute_reply":"2022-07-09T18:42:03.578073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run(train):\n\n    LOGGER.info(f\"==============================================\")\n    LOGGER.info(f\"▶︎ Start Training\")\n    LOGGER.info(f\"==============================================\")\n\n    # ====================================================\n    #  Data Loader\n    # ====================================================\n\n    train_dataset = FoursquareDataset(train)\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=True,\n        num_workers=CFG.num_workers,\n        pin_memory=True,\n        drop_last=False,\n        worker_init_fn=seed_worker,\n        generator=g\n    )\n\n    # ====================================================\n    #  Model\n    # ====================================================\n    model = CustomModel(CFG.model)\n    model.to(device)\n\n    optimizer = AdamW(model.parameters(), lr=CFG.lr)\n    num_train_steps = len(train_loader)* CFG.epochs\n\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='linear':\n            scheduler = get_linear_schedule_with_warmup(\n                optimizer, num_warmup_steps=num_train_steps*0.05, num_training_steps=num_train_steps\n            )\n        elif CFG.scheduler=='cosine':\n            scheduler = get_cosine_schedule_with_warmup(\n                optimizer, num_warmup_steps=CFG.num_warmup_steps, num_training_steps=num_train_steps, num_cycles=CFG.num_cycles\n            )\n        elif CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=CFG.T_mult, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmupRestarts':\n            scheduler = CosineAnnealingWarmupRestarts(optimizer, first_cycle_steps=CFG.first_cycle_steps, cycle_mult=CFG.cycle_mult, max_lr=CFG.lr, min_lr=CFG.min_lr, gamma=CFG.gamma, last_epoch=-1)\n        return scheduler\n\n    scheduler = get_scheduler(optimizer)\n\n    criterion = nn.CrossEntropyLoss()\n\n    # ====================================================\n    #  Loop\n    # ====================================================\n\n    for epoch in range(CFG.epochs):\n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        elapsed = time.time() - start_time\n        LOGGER.info(\n            f\"Epoch {epoch+1} - avg_train_loss: {avg_loss:.6f} time: {elapsed:.0f}s\"\n        )\n\n        # Save model\n        torch.save(\n            model.state_dict(), OUTPUT_DIR + f\"bert_epoch{epoch+1}.pth\"\n        )\n\n        # ====================================================\n        #  Judge if the best score or not\n        # ====================================================\n        best_loss = np.inf\n\n        if avg_loss < best_loss:\n            best_loss = avg_loss\n            LOGGER.info(f\"Epoch {epoch+1} - Best Loss Model\")\n    \n    # ==============================================\n    # Create Kaggle Dataset\n    # ==============================================\n\n    !kaggle datasets init -p $OUTPUT_DIR\n\n    metadata = {\"id\": f\"shkanda/foursquare-stage1-exp{CFG.exp}\",\n                    \"title\": f\"foursquare-stage1-exp{CFG.exp}\",\n                    \"licenses\": [{\"name\": \"CC0-1.0\"}]}\n\n    with open(OUTPUT_DIR+'dataset-metadata.json', 'w') as fp:\n        json.dump(metadata, fp)\n\n    !kaggle datasets create -p $OUTPUT_DIR","metadata":{"id":"d2c2ee75","papermill":{"duration":0.090401,"end_time":"2022-07-03T16:51:12.427481","exception":false,"start_time":"2022-07-03T16:51:12.33708","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.580215Z","iopub.execute_input":"2022-07-09T18:42:03.580731Z","iopub.status.idle":"2022-07-09T18:42:03.652742Z","shell.execute_reply.started":"2022-07-09T18:42:03.580692Z","shell.execute_reply":"2022-07-09T18:42:03.651705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"98318f53","papermill":{"duration":0.01321,"end_time":"2022-07-03T16:51:12.454116","exception":false,"start_time":"2022-07-03T16:51:12.440906","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if CFG.train:\n    run(original_df)","metadata":{"id":"d915bd2e","outputId":"18dc049c-8408-46d1-f58d-2486369ded25","papermill":{"duration":0.019167,"end_time":"2022-07-03T16:51:12.486858","exception":false,"start_time":"2022-07-03T16:51:12.467691","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.654182Z","iopub.execute_input":"2022-07-09T18:42:03.654582Z","iopub.status.idle":"2022-07-09T18:42:03.664818Z","shell.execute_reply.started":"2022-07-09T18:42:03.654544Z","shell.execute_reply":"2022-07-09T18:42:03.663974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{"id":"50f6ae5a","papermill":{"duration":0.012973,"end_time":"2022-07-03T16:51:12.513386","exception":false,"start_time":"2022-07-03T16:51:12.500413","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def inference_fn(test):\n\n    test_dataset = FoursquareDataset(test, include_labels=False)\n\n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=CFG.num_workers,\n        pin_memory=True,\n        drop_last=False,\n        worker_init_fn=seed_worker,\n        generator=g\n    )\n\n    model = CustomModel(CFG.model)\n    path = MODEL_DIR + \"bert_epoch15.pth\"\n    state = torch.load(path, map_location=torch.device('cpu'))\n    model.load_state_dict(state)\n    model.to(device)\n    model.eval()\n\n    preds = []\n    for step, (input_ids, attention_mask, lat, lon) in tqdm(enumerate(test_loader), total=len(test_loader)):\n        input_ids = input_ids.to(device)\n        attention_mask = attention_mask.to(device)\n        lat = lat.to(device)\n        lon = lon.to(device)\n        \n        with torch.no_grad():\n            pred = model.extract(input_ids, attention_mask, lat, lon)\n        preds.append(pred.detach().cpu().numpy())\n        \n    preds = np.concatenate(preds)\n    \n    return preds","metadata":{"id":"ddb48ecf","papermill":{"duration":0.024257,"end_time":"2022-07-03T16:51:12.551036","exception":false,"start_time":"2022-07-03T16:51:12.526779","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:42:03.670056Z","iopub.execute_input":"2022-07-09T18:42:03.672010Z","iopub.status.idle":"2022-07-09T18:42:03.681541Z","shell.execute_reply.started":"2022-07-09T18:42:03.671976Z","shell.execute_reply":"2022-07-09T18:42:03.680806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(embedding):\n\n    df = []\n\n    for country, country_df in tqdm(original_df.groupby('country')):\n        \n        country_embedding = embedding[country_df.index]\n        country_df = country_df.reset_index(drop=True)\n        country_df[\"index\"] = country_df.index\n        \n        neighbors = min(len(country_df), CFG.n_neighbors)\n\n        knn = NearestNeighbors(n_neighbors = neighbors,\n                                             metric = 'cosine',\n                                             algorithm='brute',\n                                             n_jobs = -1)\n        \n        knn.fit(country_embedding, country_df.index)\n\n        dists, nears = knn.kneighbors(country_embedding, return_distance = True)\n\n        for k in range(neighbors):            \n            cur_df = country_df[['id']]\n            cur_df['match_id'] = country_df['id'].values[nears[:, k]]\n            cur_df['cos_dist'] = dists[:, k]\n\n            cur_df = cur_df[cur_df[\"cos_dist\"]<CFG.threshold]\n\n            df.append(cur_df)\n\n    df = pd.concat(df).reset_index(drop = True)\n\n    return df","metadata":{"id":"69d68292","papermill":{"duration":0.024608,"end_time":"2022-07-03T16:51:12.588949","exception":false,"start_time":"2022-07-03T16:51:12.564341","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:44:30.653386Z","iopub.execute_input":"2022-07-09T18:44:30.653743Z","iopub.status.idle":"2022-07-09T18:44:30.665665Z","shell.execute_reply.started":"2022-07-09T18:44:30.653713Z","shell.execute_reply":"2022-07-09T18:44:30.664814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def post_process(df):\n\n    id2match = dict(zip(df['id'].values, df['matches'].str.split()))\n\n    for base, match in df[['id', 'matches']].values:\n        match = match.split()\n        if len(match) == 1:        \n            continue\n\n        for m in match:\n            if base not in id2match[m]:\n                id2match[m].append(base)\n    df['matches'] = df['id'].map(id2match).map(' '.join)\n    \n    return df ","metadata":{"id":"bac45403","papermill":{"duration":0.021537,"end_time":"2022-07-03T16:51:12.623613","exception":false,"start_time":"2022-07-03T16:51:12.602076","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:44:31.570644Z","iopub.execute_input":"2022-07-09T18:44:31.571230Z","iopub.status.idle":"2022-07-09T18:44:31.579827Z","shell.execute_reply.started":"2022-07-09T18:44:31.571191Z","shell.execute_reply":"2022-07-09T18:44:31.578991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not CFG.train:\n    \n    embedding = inference_fn(original_df)\n    df = predict(embedding)\n    \n    # post process1\n    tmp_df = original_df[[\"id\"]]\n    tmp_df[\"match_id\"] = tmp_df[\"id\"]\n    df = pd.concat([df, tmp_df]).drop_duplicates([\"id\", \"match_id\"]).reset_index(drop=True)\n    \n    sub = df.groupby([\"id\"])[\"match_id\"].apply(list).reset_index()\n    sub.columns = [\"id\", \"matches\"]\n    sub[\"matches\"] = sub[\"matches\"].map(\" \".join)\n\n    # post process2\n    sub = post_process(sub)\n\n    sub.to_csv(\"submission.csv\", index=False)\n    display(sub)","metadata":{"id":"d08e65b0","papermill":{"duration":26.44037,"end_time":"2022-07-03T16:51:39.077033","exception":false,"start_time":"2022-07-03T16:51:12.636663","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-09T18:44:31.885879Z","iopub.execute_input":"2022-07-09T18:44:31.886257Z","iopub.status.idle":"2022-07-09T18:44:40.548644Z","shell.execute_reply.started":"2022-07-09T18:44:31.886227Z","shell.execute_reply":"2022-07-09T18:44:40.547764Z"},"trusted":true},"execution_count":null,"outputs":[]}]}