Skip to content
Merged
Show file tree
Hide file tree
Changes from 6 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
376 changes: 376 additions & 0 deletions examples/bert4rec.ipynb
Original file line number Diff line number Diff line change
@@ -0,0 +1,376 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import sys\n",
"sys.path.append(\"/data/home/maspirina1/tasks/repo/RecTools/\")"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"import os\n",
"import torch\n",
"import threadpoolctl\n",
"from pathlib import Path\n",
"from lightning_fabric import seed_everything\n",
"\n",
"import numpy as np\n",
"import pandas as pd\n",
"from rectools import Columns\n",
"\n",
"\n",
"from rectools.dataset import Dataset\n",
"from rectools.metrics import MAP, calc_metrics, MeanInvUserFreq, Serendipity\n",
"from rectools.models.bert4rec import IdEmbeddingsItemNet, BERT4RecModel"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"<threadpoolctl.threadpool_limits at 0x7fb340f9bfa0>"
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n",
"os.environ[\"OPENBLAS_NUM_THREADS\"] = \"1\"\n",
"threadpoolctl.threadpool_limits(1, \"blas\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Prepare data"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [],
"source": [
"# %%time\n",
"# !wget -q https://github.com/irsafilo/KION_DATASET/raw/f69775be31fa5779907cf0a92ddedb70037fb5ae/data_original.zip -O data_original.zip\n",
"# !unzip -o data_original.zip\n",
"# !rm data_original.zip"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [],
"source": [
"DATA_PATH = Path(\"data_original\")\n",
"\n",
"interactions = (\n",
" pd.read_csv(DATA_PATH / 'interactions.csv', parse_dates=[\"last_watch_dt\"])\n",
" .rename(columns={\"last_watch_dt\": \"datetime\"})\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [],
"source": [
"interactions[Columns.Weight] = np.where(interactions['watched_pct'] > 10, 3, 1)\n",
"\n",
"# Split to train / test\n",
"max_date = interactions[Columns.Datetime].max()\n",
"train = interactions[interactions[Columns.Datetime] < max_date - pd.Timedelta(days=7)].copy()\n",
"test = interactions[interactions[Columns.Datetime] >= max_date - pd.Timedelta(days=7)].copy()\n",
"train.drop(train.query(\"total_dur < 300\").index, inplace=True)\n",
"\n",
"# drop items with less than 20 interactions in train\n",
"items = train[\"item_id\"].value_counts()\n",
"items = items[items >= 20]\n",
"items = items.index.to_list()\n",
"train = train[train[\"item_id\"].isin(items)]\n",
" \n",
"# drop users with less than 2 interactions in train\n",
"users = train[\"user_id\"].value_counts()\n",
"users = users[users >= 2]\n",
"users = users.index.to_list()\n",
"train = train[(train[\"user_id\"].isin(users))]\n",
"\n",
"users = train[\"user_id\"].drop_duplicates().to_list()\n",
"\n",
"# drop cold users from test\n",
"test_users_sasrec = test[Columns.User].unique()\n",
"cold_users = set(test[Columns.User]) - set(train[Columns.User])\n",
"test.drop(test[test[Columns.User].isin(cold_users)].index, inplace=True)\n",
"test_users = test[Columns.User].unique()\n"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [],
"source": [
"items = pd.read_csv(DATA_PATH / 'items.csv')"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [],
"source": [
"# Process item features to the form of a flatten dataframe\n",
"items = items.loc[items[Columns.Item].isin(train[Columns.Item])].copy()\n",
"items[\"genre\"] = items[\"genres\"].str.lower().str.replace(\", \", \",\", regex=False).str.split(\",\")\n",
"genre_feature = items[[\"item_id\", \"genre\"]].explode(\"genre\")\n",
"genre_feature.columns = [\"id\", \"value\"]\n",
"genre_feature[\"feature\"] = \"genre\"\n",
"content_feature = items.reindex(columns=[Columns.Item, \"content_type\"])\n",
"content_feature.columns = [\"id\", \"value\"]\n",
"content_feature[\"feature\"] = \"content_type\"\n",
"item_features = pd.concat((genre_feature, content_feature))\n",
"\n",
"candidate_items = interactions['item_id'].drop_duplicates().astype(int)\n",
"test[\"user_id\"] = test[\"user_id\"].astype(int)\n",
"test[\"item_id\"] = test[\"item_id\"].astype(int)\n",
"\n",
"catalog=train[Columns.Item].unique()"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [],
"source": [
"dataset_no_features = Dataset.construct(\n",
" interactions_df=train,\n",
")\n",
"\n",
"dataset_item_features = Dataset.construct(\n",
" interactions_df=train,\n",
" item_features_df=item_features,\n",
" cat_item_features=[\"genre\", \"content_type\"],\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [],
"source": [
"metrics_name = {\n",
" 'MAP': MAP,\n",
" 'MIUF': MeanInvUserFreq,\n",
" 'Serendipity': Serendipity\n",
" \n",
"\n",
"}\n",
"metrics = {}\n",
"for metric_name, metric in metrics_name.items():\n",
" for k in (1, 5, 10):\n",
" metrics[f'{metric_name}@{k}'] = metric(k=k)\n",
"\n",
"# list with metrics results of all models\n",
"features_results = []\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# BERT4Rec"
]
},
{
"cell_type": "code",
"execution_count": 10,
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Seed set to 32\n"
]
},
{
"data": {
"text/plain": [
"32"
]
},
"execution_count": 10,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"RANDOM_SEED = 32\n",
"torch.use_deterministic_algorithms(True)\n",
"seed_everything(RANDOM_SEED, workers=True)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### BERT4Rec with item ids embeddings in ItemNetBlock"
]
},
{
"cell_type": "code",
"execution_count": 11,
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Trainer will use only 1 of 2 GPUs because it is running inside an interactive / notebook environment. You may try to set `Trainer(devices=2)` but please note that multi-GPU inside interactive / notebook environments is considered experimental and unstable. Your mileage may vary.\n",
"GPU available: True (cuda), used: True\n",
"TPU available: False, using: 0 TPU cores\n",
"IPU available: False, using: 0 IPUs\n",
"HPU available: False, using: 0 HPUs\n"
]
}
],
"source": [
"model = BERT4RecModel(\n",
" n_blocks=3,\n",
" n_heads=4,\n",
" dropout_rate=0.2,\n",
" session_max_len=32,\n",
" lr=1e-3,\n",
" epochs=5,\n",
" verbose=1,\n",
" mask_prob=0.5,\n",
" deterministic=True,\n",
" item_net_block_types=(IdEmbeddingsItemNet, ),\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 18,
"metadata": {},
"outputs": [],
"source": [
"%%time\n",
"model.fit(dataset_no_features)"
]
},
{
"cell_type": "code",
"execution_count": 17,
"metadata": {},
"outputs": [],
"source": [
"%%time\n",
"recos = model.recommend(\n",
" users=test_users_sasrec, \n",
" dataset=dataset_item_features,\n",
" k=10,\n",
" filter_viewed=True,\n",
" on_unsupported_targets=\"warn\"\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 14,
"metadata": {},
"outputs": [],
"source": [
"recos[\"item_id\"] = recos[\"item_id\"].apply(str)\n",
"test[\"item_id\"] = test[\"item_id\"].astype(str)\n",
"metric_values = calc_metrics(metrics, recos[[\"user_id\", \"item_id\", \"rank\"]], test, train, catalog)\n",
"metric_values[\"model\"] = \"bert4rec_ids\"\n",
"features_results.append(metric_values)"
]
},
{
"cell_type": "code",
"execution_count": 17,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[{'MAP@1': 0.04153170911358253,\n",
" 'MAP@5': 0.07096106984411608,\n",
" 'MAP@10': 0.07874644762957389,\n",
" 'MIUF@1': 18.824620072061013,\n",
" 'MIUF@5': 18.824620072061013,\n",
" 'MIUF@10': 18.824620072061013,\n",
" 'Serendipity@1': 0.08494799640990444,\n",
" 'Serendipity@5': 0.05316937762913509,\n",
" 'Serendipity@10': 0.03892074762532452,\n",
" 'model': 'bert4rec_ids'}]"
]
},
"execution_count": 17,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"features_results"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.8.2"
}
},
"nbformat": 4,
"nbformat_minor": 2
}
Loading