Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
555 changes: 555 additions & 0 deletions EDA.ipynb

Large diffs are not rendered by default.

File renamed without changes.
447 changes: 415 additions & 32 deletions models/pmf_model.py

Large diffs are not rendered by default.

88 changes: 65 additions & 23 deletions models/svd_model.py
Original file line number Diff line number Diff line change
@@ -1,41 +1,83 @@
import numpy as np
import pandas as pd
import logging
from scipy.sparse.linalg import svds
from scipy.sparse import csr_matrix
from typing import Optional

logger = logging.getLogger(__name__)


class SVDRecommender:
def __init__(self, k: int = 50):
"""
Singular Value Decomposition (SVD) Recommender with Item Bias Shrinkage.

This model performs matrix factorization while allowing for the 'dampening'
of item biases to prevent overfitting and control performance gaps.
"""

def __init__(self, k: int = 10, bias_weight: float = 0.5):
"""
Args:
k (int): Number of latent factors to extract.
bias_weight (float): Scalar (0-1) to dampen item biases.
Lower values 'nerf' SVD performance.
"""
self.k = k
self.u = None
self.sigma = None
self.vt = None
self.preds_matrix = None

def fit(self, train_matrix_df):
"""Decomposes the matrix and generates full predictions."""
logger.info(f"🤖 Starting SVD decomposition with k={self.k}")
self.bias_weight = bias_weight
self.preds_matrix: Optional[np.ndarray] = None

def fit(self, train_matrix_df: pd.DataFrame) -> pd.DataFrame:
"""
Learns the latent factors and reconstructs the rating matrix.

Args:
train_matrix_df (pd.DataFrame): User-Item matrix with NaNs for missing ratings.

Returns:
pd.DataFrame: Dense prediction matrix with the same shape/indices as input.

Raises:
ValueError: If the input matrix is empty or k is larger than matrix dimensions.
"""
try:
# 1. Edge Case Check: If all values are zero, don't call svds
if (train_matrix_df.values == 0).all():
logger.info(f"🚀 SVD Fit | k={self.k}, weight={self.bias_weight}")
if train_matrix_df.fillna(0.0).values.sum() == 0:
logger.warning(
"⚠️ Input matrix is all zeros. Returning zero matrix predictions."
"⚠️ Matrix is empty/all zeros. Returning zero predictions."
)
return pd.DataFrame(
0.0, index=train_matrix_df.index, columns=train_matrix_df.columns
)
self.preds_matrix = np.zeros(train_matrix_df.shape)
return self.preds_matrix
# 1. Calculate and dampen Item Biases
# nanmean handles missing ratings correctly
item_biases = np.nanmean(train_matrix_df.values, axis=0)
item_biases = np.nan_to_num(item_biases, nan=0.0) * self.bias_weight

# 2. Centering and Imputation
R_centered = train_matrix_df.values - item_biases
R_filled = np.nan_to_num(R_centered, nan=0.0)

# Safety check for k
min_dim = min(R_filled.shape) - 1
current_k = min(self.k, min_dim)

# 3. Scipy SVDS
U, sigma, Vt = svds(R_filled, k=current_k)

# 2. Standard SVD Flow
sparse_matrix = csr_matrix(train_matrix_df.values)
u, sigma, vt = svds(sparse_matrix, k=self.k)
# Sort factors by importance (descending)
idx = np.argsort(sigma)[::-1]
U, sigma, Vt = U[:, idx], sigma[idx], Vt[idx, :]

sigma_diag = np.diag(sigma)
self.preds_matrix = np.dot(np.dot(u, sigma_diag), vt)
# 4. Reconstruction
interaction_preds = (U * sigma) @ Vt
self.preds_matrix = interaction_preds + item_biases

logger.info("✅ SVD Decomposition successful.")
return self.preds_matrix
return pd.DataFrame(
self.preds_matrix,
index=train_matrix_df.index,
columns=train_matrix_df.columns,
)

except Exception as e:
logger.error(f"❌ SVD Math Error: {e}")
logger.error(f"❌ SVD Fit failed: {str(e)}")
raise
262 changes: 254 additions & 8 deletions poetry.lock

Large diffs are not rendered by default.

Binary file modified processed/user_means.npy
Binary file not shown.
8 changes: 6 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ authors = [
package-mode = false
requires-python = ">=3.11"
dependencies = [
"pandas (<3.0.0)",
"pandas (==2.3.3)",
"numpy (>=2.4.3,<3.0.0)",
"scipy (>=1.17.1,<2.0.0)",
"matplotlib (>=3.10.8,<4.0.0)",
Expand All @@ -17,7 +17,11 @@ dependencies = [
"pymysql (==1.1.2)",
"pyarrow (==23.0.1)",
"jinja2 (==3.1.6)",
"setuptools (==78.1.1)"
"setuptools (==78.1.1)",
"seaborn (>=0.13.2,<0.14.0)",
"tabulate (>=0.10.0,<0.11.0)",
"requests (==2.33.0)",
"cryptography (==46.0.7)"
]


Expand Down
42 changes: 42 additions & 0 deletions reports/audit_summary.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
# Recommendation System Audit Report

## Performance Summary
- **Baseline (SVD) RMSE**: 0.8846
- **Advanced (PMF) RMSE**: 0.8402
- **Relative Improvement**: 5.02%

## Conclusion
The PMF model is the production candidate.
## Interpretability Analysis
### Global Trends (Factor Analysis)
```
--- Global Latent Factor Trends ---

Factor 0 (Representative Movies):
1. Hate (Haine, La) (1995)
2. Happy Gilmore (1996)
3. Alaska (1996)
4. Cutthroat Island (1995)
5. Godfather, The (1972)

Factor 1 (Representative Movies):
1. Hugo Pool (1997)
2. Kim (1950)
3. Star Wars: Episode V - The Empire Strikes Back (1980)
4. Amityville II: The Possession (1982)
5. Winnie the Pooh and the Blustery Day (1968)

Factor 2 (Representative Movies):
1. Spitfire Grill, The (1996)
2. Day the Sun Turned Cold, The (Tianguo niezi) (1994)
3. Heathers (1989)
4. Rude (1995)
5. Vegas Vacation (1997)
```
### Local Example
```
--- Local Interpretability ---
User 5300 -> Movie: To Wong Foo, Thanks for Everything! Julie Newmar (1995)
Strongest Driver: Latent Factor 24
Factor Contribution: -0.1005
```
Binary file added reports/latent_factors_heatmap.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
35 changes: 32 additions & 3 deletions reports/model_metrics.json
Original file line number Diff line number Diff line change
@@ -1,6 +1,35 @@
{
"svd": {
"rmse": 3.0864641494337537,
"k": 50
"SVD_RMSE": 0.8846191731074615,
"PMF_RMSE": 0.8402208901082122,
"PMF_vs_SVD_improvement_%": 5.02,
"svd_optimized_params": {
"k": 10,
"bias_weight": 0.6
},
"pmf_optimized_params": {
"factors": 25
},
"additional_metrics": {
"svd": {
"rmse": 0.8846191731074615,
"mae": 0.696405907963734,
"auc": 0.7965290839872283,
"precision": 0.7495359485212226,
"recall": 0.8171221973396648,
"ste": 0.0015356929678274043
},
"pmf": {
"rmse": 0.8402208901082122,
"mae": 0.6544133079846005,
"auc": 0.8183532169812058,
"precision": 0.7715854970041481,
"recall": 0.8130345627714972,
"ste": 0.0014835735257777262
}
},
"audit_summary": {
"winner": "PMF",
"statistically_significant": true,
"target_met": true
}
}
Binary file added reports/pmf_confusion_matrix.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file modified reports/pmf_convergence.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added reports/pmf_drift_chart.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added reports/pmf_factors/U_factors_best.npy
Binary file not shown.
Binary file added reports/pmf_factors/V_factors_best.npy
Binary file not shown.
Binary file added reports/rmse_comparison.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added reports/svd_confusion_matrix.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added reports/svd_drift_chart.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
169 changes: 57 additions & 112 deletions requirements.txt
Original file line number Diff line number Diff line change
@@ -1,112 +1,57 @@
-i https://pkgs.safetycli.com/repository/gritlab/project/matrix-factorization/pypi/simple/
anyio==4.12.1
appnope==0.1.4
argon2-cffi==25.1.0
argon2-cffi-bindings==25.1.0
arrow==1.4.0
asttokens==3.0.1
async-lru==2.3.0
attrs==26.1.0
babel==2.18.0
beautifulsoup4==4.14.3
bleach==6.3.0
certifi==2026.2.25
cffi==2.0.0
charset-normalizer==3.4.6
comm==0.2.3
contourpy==1.3.3
cycler==0.12.1
debugpy==1.8.20
decorator==5.2.1
defusedxml==0.7.1
executing==2.2.1
fastjsonschema==2.21.2
fonttools==4.62.1
fqdn==1.5.1
h11==0.16.0
httpcore==1.0.9
httpx==0.28.1
idna==3.11
ipykernel==7.2.0
ipython==9.10.0
ipython_pygments_lexers==1.1.1
ipywidgets==8.1.8
isoduration==20.11.0
jedi==0.19.2
Jinja2==3.1.6
joblib==1.5.3
json5==0.13.0
jsonpointer==3.1.0
jsonschema==4.26.0
jsonschema-specifications==2025.9.1
jupyter==1.1.1
jupyter-console==6.6.3
jupyter-events==0.12.0
jupyter-lsp==2.3.0
jupyter_client==8.8.0
jupyter_core==5.9.1
jupyter_server==2.17.0
jupyter_server_terminals==0.5.4
jupyterlab==4.5.6
jupyterlab_pygments==0.3.0
jupyterlab_server==2.28.0
jupyterlab_widgets==3.0.16
kiwisolver==1.5.0
lark==1.3.1
MarkupSafe==3.0.3
matplotlib==3.10.8
matplotlib-inline==0.2.1
mistune==3.2.0
nbclient==0.10.4
nbconvert==7.17.0
nbformat==5.10.4
nest-asyncio==1.6.0
notebook==7.5.5
notebook_shim==0.2.4
numpy==2.4.3
overrides==7.7.0
packaging==26.0
pandas==3.0.1
pandocfilters==1.5.1
parso==0.8.6
pexpect==4.9.0
pillow==12.1.1
platformdirs==4.9.4
prometheus_client==0.24.1
prompt_toolkit==3.0.52
psutil==7.2.2
ptyprocess==0.7.0
pure_eval==0.2.3
pycparser==3.0
Pygments==2.19.2
pyparsing==3.3.2
python-dateutil==2.9.0.post0
python-json-logger==4.0.0
PyYAML==6.0.3
pyzmq==27.1.0
referencing==0.37.0
requests==2.32.5
rfc3339-validator==0.1.4
rfc3986-validator==0.1.1
rfc3987-syntax==1.1.0
rpds-py==0.30.0
scikit-learn==1.8.0
scipy==1.17.1
Send2Trash==2.1.0
six==1.17.0
soupsieve==2.8.3
stack-data==0.6.3
terminado==0.18.1
threadpoolctl==3.6.0
tinycss2==1.4.0
tornado==6.5.5
traitlets==5.14.3
typing_extensions==4.15.0
tzdata==2025.3
uri-template==1.3.0
urllib3==2.6.3
wcwidth==0.6.0
webcolors==25.10.0
webencodings==0.5.1
websocket-client==1.9.0
widgetsnbextension==4.0.15
--index-url https://pkgs.safetycli.com/repository/gritlab/project/matrix-factorization/pypi/simple

altair==6.0.0 ; python_version >= "3.11"
attrs==26.1.0 ; python_version >= "3.11"
blinker==1.9.0 ; python_version >= "3.11"
cachetools==7.0.5 ; python_version >= "3.11"
certifi==2026.2.25 ; python_version >= "3.11"
cffi==2.0.0 ; python_version >= "3.11" and platform_python_implementation != "PyPy"
charset-normalizer==3.4.6 ; python_version >= "3.11"
click==8.3.1 ; python_version >= "3.11"
colorama==0.4.6 ; python_version >= "3.11" and platform_system == "Windows"
contourpy==1.3.3 ; python_version >= "3.11"
cryptography==46.0.7 ; python_version >= "3.11"
cycler==0.12.1 ; python_version >= "3.11"
fonttools==4.62.1 ; python_version >= "3.11"
gitdb==4.0.12 ; python_version >= "3.11"
gitpython==3.1.46 ; python_version >= "3.11"
idna==3.11 ; python_version >= "3.11"
jinja2==3.1.6 ; python_version >= "3.11"
joblib==1.5.3 ; python_version >= "3.11"
jsonschema-specifications==2025.9.1 ; python_version >= "3.11"
jsonschema==4.26.0 ; python_version >= "3.11"
kiwisolver==1.5.0 ; python_version >= "3.11"
markupsafe==3.0.3 ; python_version >= "3.11"
matplotlib==3.10.8 ; python_version >= "3.11"
narwhals==2.18.0 ; python_version >= "3.11"
numpy==2.4.3 ; python_version >= "3.11"
packaging==26.0 ; python_version >= "3.11"
pandas==2.3.3 ; python_version >= "3.11"
pillow==12.1.1 ; python_version >= "3.11"
protobuf==6.33.6 ; python_version >= "3.11"
pyarrow==23.0.1 ; python_version >= "3.11"
pycparser==3.0 ; platform_python_implementation != "PyPy" and implementation_name != "PyPy" and python_version >= "3.11"
pydeck==0.9.1 ; python_version >= "3.11"
pymysql==1.1.2 ; python_version >= "3.11"
pyparsing==3.3.2 ; python_version >= "3.11"
python-dateutil==2.9.0.post0 ; python_version >= "3.11"
pytz==2026.1.post1 ; python_version >= "3.11"
referencing==0.37.0 ; python_version >= "3.11"
requests==2.33.0 ; python_version >= "3.11"
rpds-py==0.30.0 ; python_version >= "3.11"
scikit-learn==1.8.0 ; python_version >= "3.11"
scipy==1.17.1 ; python_version >= "3.11"
seaborn==0.13.2 ; python_version >= "3.11"
setuptools==78.1.1 ; python_version >= "3.11"
six==1.17.0 ; python_version >= "3.11"
smmap==5.0.3 ; python_version >= "3.11"
streamlit==1.55.0 ; python_version >= "3.11"
tabulate==0.10.0 ; python_version >= "3.11"
tenacity==9.1.4 ; python_version >= "3.11"
threadpoolctl==3.6.0 ; python_version >= "3.11"
toml==0.10.2 ; python_version >= "3.11"
tornado==6.5.5 ; python_version >= "3.11"
typing-extensions==4.15.0 ; python_version >= "3.11"
tzdata==2025.3 ; python_version >= "3.11"
urllib3==2.6.3 ; python_version >= "3.11"
watchdog==6.0.0 ; python_version >= "3.11" and platform_system != "Darwin"
Loading
Loading