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
2 changes: 1 addition & 1 deletion src/lenskit/flexmf/_explicit.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@ def create_model(self) -> FlexMFModel:

@override
def train_batch(self, batch: FlexMFTrainingBatch) -> float:
if self.config.reg_method == "L2":
if self.explicit_norm:
result = self.model(batch.users, batch.items, return_norm=True)
pred = result[0, :]
norm = torch.mean(result[1, :]) * self.config.regularization
Expand Down
54 changes: 37 additions & 17 deletions src/lenskit/flexmf/_implicit.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from __future__ import annotations

import math
from collections.abc import Callable
from dataclasses import dataclass
from typing import Literal, TypeAlias

Expand Down Expand Up @@ -121,6 +122,20 @@ def create_trainer(self, data, options):


class FlexMFImplicitTrainer(FlexMFTrainerBase[FlexMFImplicitScorer, FlexMFImplicitConfig]):
_loss: Callable[[torch.Tensor, torch.Tensor, float, torch.Tensor | None], torch.Tensor]

def __init__(self, component, data, options):
super().__init__(component, data, options)
match self.config.loss:
case "logistic":
self._loss = _loss_logistic
case "pairwise":
self._loss = _loss_pairwise
case "warp":
self._loss = _loss_warp
case other: # pragma: nocover
raise ValueError(f"unknown loss {other}")

def prepare_data(self, data: Dataset) -> FlexMFTrainingData:
"""
Set up the training data and context for the scorer.
Expand Down Expand Up @@ -164,7 +179,7 @@ def create_model(self) -> FlexMFModel:
)

def score(self, users, items) -> tuple[torch.Tensor, torch.Tensor]:
if self.config.reg_method == "L2":
if self.explicit_norm:
result = self.model(users, items, return_norm=True)
scores = result[0, ...]

Expand All @@ -183,23 +198,9 @@ def train_batch(self, batch: FlexMFTrainingBatch) -> float:

neg_items, neg_pred, neg_norm, weights = self.scored_negatives(batch, users, pos_pred)

match self.config.loss:
case "logistic":
pos_lp = -F.logsigmoid(pos_pred) * self.config.positive_weight
neg_lp = -F.logsigmoid(-neg_pred)
tot_lp = pos_lp.sum() + neg_lp.sum()
tot_n = pos_lp.nelement() + neg_lp.nelement()
loss = tot_lp / tot_n
loss = self._loss(pos_pred, neg_pred, self.config.positive_weight, weights)

case "pairwise":
lp = -F.logsigmoid(pos_pred - neg_pred)
loss = lp.mean()
case "warp":
assert weights is not None
lp = -F.logsigmoid(pos_pred - neg_pred) * weights
loss = lp.mean()

if self.config.reg_method == "L2":
if self.explicit_norm:
loss = loss + self.config.regularization * 0.5 * (pos_norm.mean() + neg_norm.mean())

loss.backward()
Expand Down Expand Up @@ -289,3 +290,22 @@ def scored_negatives(self, batch, users, pos_scores):
+ 1 / (120 * ranks**4)
)
return neg_items, neg_scores.reshape(-1, 1), neg_norms, weights.to(pos_scores.device)


def _loss_logistic(pos_pred, neg_pred, pos_weight, weights):
pos_lp = -F.logsigmoid(pos_pred) * pos_weight
neg_lp = -F.logsigmoid(-neg_pred)
tot_lp = pos_lp.sum() + neg_lp.sum()
tot_n = pos_lp.nelement() + neg_lp.nelement()
return tot_lp / tot_n


def _loss_pairwise(pos_pred, neg_pred, pos_weight, weights):
lp = -F.logsigmoid(pos_pred - neg_pred)
return lp.mean()


def _loss_warp(pos_pred, neg_pred, pos_weight, weights):
assert weights is not None
lp = -F.logsigmoid(pos_pred - neg_pred) * weights
return lp.mean()
6 changes: 6 additions & 0 deletions src/lenskit/flexmf/_training.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,11 @@ class FlexMFTrainerBase(ModelTrainer, Generic[Comp, Cfg]):
Torch device for training.
"""

explicit_norm: bool
"""
Use explicit (L2) norm for normalizing data.
"""

rng: np.random.Generator
"""
NumPy generator for random number generation.
Expand All @@ -82,6 +87,7 @@ def __init__(self, component: Comp, data: Dataset, options: TrainingOptions):
ensure_parallel_init()

self.component = component
self.explicit_norm = component.config.reg_method == "L2"

self.log = _log.bind(scorer=self.__class__.__name__, size=self.config.embedding_size)

Expand Down