sasrec-pytorch / utils.py
MongrelIntruder's picture
Upload utils.py with huggingface_hub
5006e18 verified
Raw
History Blame Contribute Delete
6.39 kB
import copy
import random
from collections import defaultdict
from multiprocessing import Process, Queue
import numpy as np
from tqdm.notebook import tqdm
def build_index(dataset_name):
ui_mat = np.loadtxt(f"data/{dataset_name}.txt", dtype=np.int32)
n_users = ui_mat[:, 0].max()
n_items = ui_mat[:, 1].max()
u2i_index = [[] for _ in range(n_users + 1)]
i2u_index = [[] for _ in range(n_items + 1)]
for ui_pair in ui_mat:
u2i_index[ui_pair[0]].append(ui_pair[1])
i2u_index[ui_pair[1]].append(ui_pair[0])
return u2i_index, i2u_index
# sampler for batch generation
def random_neq(l, r, s):
t = np.random.randint(l, r)
while t in s:
t = np.random.randint(l, r)
return t
def sample_function(
user_train, usernum, itemnum, batch_size, maxlen, result_queue, SEED
):
def sample(uid):
# uid = np.random.randint(1, usernum + 1)
while len(user_train[uid]) <= 1:
uid = np.random.randint(1, usernum + 1)
seq = np.zeros([maxlen], dtype=np.int32)
pos = np.zeros([maxlen], dtype=np.int32)
neg = np.zeros([maxlen], dtype=np.int32)
nxt = user_train[uid][-1]
idx = maxlen - 1
ts = set(user_train[uid])
for i in reversed(user_train[uid][:-1]):
seq[idx] = i
pos[idx] = nxt
neg[idx] = random_neq(1, itemnum + 1, ts) # Don't need "if nxt != 0"
nxt = i
idx -= 1
if idx == -1:
break
return (uid, seq, pos, neg)
np.random.seed(SEED)
uids = np.arange(1, usernum + 1, dtype=np.int32)
counter = 0
while True:
if counter % usernum == 0:
np.random.shuffle(uids)
one_batch = []
for i in range(batch_size):
one_batch.append(sample(uids[counter % usernum]))
counter += 1
result_queue.put(zip(*one_batch))
class WarpSampler(object):
def __init__(self, User, usernum, itemnum, batch_size=64, maxlen=10, n_workers=1):
self.result_queue = Queue(maxsize=n_workers * 10)
self.processors = []
for i in range(n_workers):
self.processors.append(
Process(
target=sample_function,
args=(
User,
usernum,
itemnum,
batch_size,
maxlen,
self.result_queue,
np.random.randint(2e9),
),
)
)
self.processors[-1].daemon = True
self.processors[-1].start()
def next_batch(self):
return self.result_queue.get()
def close(self):
for p in self.processors:
p.terminate()
p.join()
def data_partition(fname):
usernum = 0
itemnum = 0
User = defaultdict(list)
user_train = {}
user_valid = {}
user_test = {}
# assuming user/item index starts from 1
with open(f"data/{fname}.txt", "r") as f:
for line in f:
u, i = line.rstrip().split(" ")
u = int(u)
i = int(i)
usernum = max(u, usernum)
itemnum = max(i, itemnum)
User[u].append(i)
f.close()
for user in User:
nfeedback = len(User[user])
if nfeedback < 4:
user_train[user] = User[user]
user_valid[user] = []
user_test[user] = []
else:
user_train[user] = User[user][:-2]
user_valid[user] = []
user_valid[user].append(User[user][-2])
user_test[user] = []
user_test[user].append(User[user][-1])
return [user_train, user_valid, user_test, usernum, itemnum]
def evaluate(model, dataset, args):
[train, valid, test, usernum, itemnum] = copy.deepcopy(dataset)
NDCG = 0.0
HT = 0.0
valid_user = 0.0
if usernum > 10000:
users = random.sample(range(1, usernum + 1), 10000)
else:
users = range(1, usernum + 1)
for u in tqdm(users, desc="evaluating (test)", leave=False):
if len(train[u]) < 1 or len(test[u]) < 1:
continue
seq = np.zeros([args.maxlen], dtype=np.int32)
idx = args.maxlen - 1
seq[idx] = valid[u][0]
idx -= 1
for i in reversed(train[u]):
seq[idx] = i
idx -= 1
if idx == -1:
break
rated = set(train[u])
rated.add(0)
item_idx = [test[u][0]]
for _ in range(100):
t = np.random.randint(1, itemnum + 1)
while t in rated:
t = np.random.randint(1, itemnum + 1)
item_idx.append(t)
predictions = -model.predict(*[np.array(l) for l in [[u], [seq], item_idx]])
predictions = predictions[0]
rank = predictions.argsort().argsort()[0].item()
valid_user += 1
if rank < 10:
NDCG += 1 / np.log2(rank + 2)
HT += 1
return NDCG / valid_user, HT / valid_user
def evaluate_valid(model, dataset, args):
[train, valid, test, usernum, itemnum] = copy.deepcopy(dataset)
NDCG = 0.0
valid_user = 0.0
HT = 0.0
if usernum > 10000:
users = random.sample(range(1, usernum + 1), 10000)
else:
users = range(1, usernum + 1)
for u in tqdm(users, desc="evaluating (valid)", leave=False):
if len(train[u]) < 1 or len(valid[u]) < 1:
continue
seq = np.zeros([args.maxlen], dtype=np.int32)
idx = args.maxlen - 1
for i in reversed(train[u]):
seq[idx] = i
idx -= 1
if idx == -1:
break
rated = set(train[u])
rated.add(0)
item_idx = [valid[u][0]]
for _ in range(100):
t = np.random.randint(1, itemnum + 1)
while t in rated:
t = np.random.randint(1, itemnum + 1)
item_idx.append(t)
predictions = -model.predict(*[np.array(l) for l in [[u], [seq], item_idx]])
predictions = predictions[0]
rank = predictions.argsort().argsort()[0].item()
valid_user += 1
if rank < 10:
NDCG += 1 / np.log2(rank + 2)
HT += 1
return NDCG / valid_user, HT / valid_user