File size: 2,028 Bytes
08646a9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 | # SPDX-FileCopyrightText: 2026 devtaji
# SPDX-License-Identifier: Apache-2.0
"""Minimal example: load recICL and rank a toy catalog for one user.
It embeds nothing. A real catalog needs precomputed item embeddings from the encoder named in config.json
("item_embeddings": gte-Qwen2-1.5B-instruct, D=1536). Here the catalog is 500 random unit vectors, so the ranking
itself carries no meaning; the calls are exactly the ones you would make with real embeddings.
pip install -r requirements.txt
python example.py
"""
import os
import numpy as np
import recicl as ir
HERE = os.path.dirname(os.path.abspath(__file__))
# 1. model: config.json + model.safetensors
model = ir.load_model(HERE)
cfg = ir.load_config(HERE)
n_params = sum(p.numel() for p in model.parameters())
print(f"loaded {cfg['name']}: {n_params:,} parameters, D={model.input_dim}, "
f"device={next(model.parameters()).device}")
# 2. catalog: precomputed item embeddings [N, D] (row i = item index i). Toy stand-in here.
rng = np.random.default_rng(0)
N = 500
item_embeddings = rng.standard_normal((N, model.input_dim)).astype(np.float32)
catalog = ir.ItemCatalog(item_embeddings, dim=model.input_dim)
# 3. context pool: other users' item sequences over the same catalog (item indices, oldest first)
pool = [rng.integers(0, N, size=int(rng.integers(3, 20))).tolist() for _ in range(300)]
pool += [[12, 7, 311, 45, 88], [7, 311, 45, 90], [311, 45, 88, 402]] # a few users who share this user's items
retriever = ir.ContextRetriever(pool, model.config)
# 4. query: the user's recent items (oldest first) + up to 8 retrieved context sequences
history = [3, 12, 7, 311]
context = retriever.retrieve(history)
print(f"history: {history}")
print(f"context: {len(context)} sequences, e.g. {context[:3]}")
# 5. rank the whole catalog
items, scores = ir.recommend(model, catalog, history, context, k=10)
print("top-10 items (index, score):")
for rank, (i, s) in enumerate(zip(items, scores), 1):
print(f" {rank:2d}. item {i:3d} score {s:8.3f}")
|