/
githubmirror
/
Agent4Rec
Обзор
Документация
Войти
/
githubmirror
/
Agent4Rec
Код
Запросы
0
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
recommenders/models/LightGCN.py
74 строки
3 KB
chenyuxin1999
init repo
12 окт 2023, 05:43
12 окт 2023, 05:43
9a08cc7
Код
Авторство
О чём код?
import numpy as np import time import torch import torch.nn as nn import torch.nn.functional as F from models.base.abstract_model import AbstractModel from models.base.abstract_RS import AbstractRS from tqdm import tqdm class LightGCN_RS(AbstractRS): def __init__(self, args, special_args) -> None: super().__init__(args, special_args) def train_one_epoch(self, epoch): running_loss, running_mf_loss, running_reg_loss, num_batches = 0, 0, 0, 0 pbar = tqdm(enumerate(self.data.train_loader), mininterval=2, total = len(self.data.train_loader)) for batch_i, batch in pbar: batch = [x.cuda(self.device) for x in batch] users, pos_items, users_pop, pos_items_pop, pos_weights = batch[0], batch[1], batch[2], batch[3], batch[4] if self.args.infonce == 0 or self.args.neg_sample != -1: neg_items = batch[5] neg_items_pop = batch[6] self.model.train() mf_loss, reg_loss = self.model(users, pos_items, neg_items) loss = mf_loss + reg_loss self.optimizer.zero_grad() loss.backward() self.optimizer.step() # print(self.model.embed_user.weight) # print("?") running_loss += loss.detach().item() running_reg_loss += reg_loss.detach().item() running_mf_loss += mf_loss.detach().item() num_batches += 1 return [running_loss/num_batches, running_mf_loss/num_batches, running_reg_loss/num_batches] class LightGCN(AbstractModel): def __init__(self, args, data) -> None: super().__init__(args, data) def forward(self, users, pos_items, neg_items): all_users, all_items = self.compute() users_emb = all_users[users] pos_emb = all_items[pos_items] neg_emb = all_items[neg_items] userEmb0 = self.embed_user(users) posEmb0 = self.embed_item(pos_items) negEmb0 = self.embed_item(neg_items) if(self.train_norm == True): users_emb = F.normalize(users_emb, dim = -1) pos_emb = F.normalize(pos_emb, dim = -1) neg_emb = F.normalize(neg_emb, dim = -1) pos_scores = torch.sum(torch.mul(users_emb, pos_emb), dim=1) # users, pos_items, neg_items have the same shape neg_scores = torch.sum(torch.mul(users_emb, neg_emb), dim=1) regularizer = 0.5 * torch.norm(userEmb0) ** 2 + 0.5 * torch.norm(posEmb0) ** 2 + 0.5 * torch.norm(negEmb0) ** 2 regularizer = regularizer / self.batch_size maxi = torch.log(torch.sigmoid(pos_scores - neg_scores) + 1e-10) mf_loss = torch.negative(torch.mean(maxi)) reg_loss = self.decay * regularizer return mf_loss, reg_loss