Coverage for hopwise/data/dataset/kg_path_dataset.py: 81%
639 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2025
2# @Author : Giacomo Medda
3# @Email : giacomo.medda@unica.it
5from dataclasses import dataclass
6from itertools import chain, zip_longest
8import numba
9import numpy as np
11from hopwise.data import Interaction
12from hopwise.data.dataset import KnowledgeBasedDataset, UserItemKnowledgeBasedDataset
13from hopwise.utils import PathLanguageModelingTokenType, PathSamplingStrategy, progress_bar, set_color
16@dataclass
17class CSRGraph:
18 """Container for CSR graph arrays used in parallel random walks.
20 This dataclass bundles together the CSR sparse matrix components
21 needed for efficient graph traversal in numba.
23 Attributes:
24 indptr: CSR row pointers (int64)
25 indices: CSR column indices (int64)
26 relations: Edge relation types (int64)
27 """
29 indptr: np.ndarray
30 indices: np.ndarray
31 relations: np.ndarray
33 @classmethod
34 def from_sparse_matrix(cls, csr_matrix):
35 """Create CSRGraph from scipy sparse matrix with relation data.
37 Args:
38 csr_matrix: Scipy CSR matrix with relations stored in data field
40 Returns:
41 CSRGraph instance
42 """
43 indptr = csr_matrix.indptr.astype(np.int64)
44 indices = csr_matrix.indices.astype(np.int64)
45 relations = csr_matrix.data.astype(np.int64)
47 return cls(indptr=indptr, indices=indices, relations=relations)
49 def unpack(self):
50 """Unpack arrays for passing to numba functions.
52 Returns:
53 tuple: (indptr, indices, relations)
54 """
55 return self.indptr, self.indices, self.relations
58def _run_batched_numba(
59 numba_func,
60 batched_arrays,
61 fixed_args_before=(),
62 fixed_args_after=(),
63 batch_size=50000,
64 desc="Processing",
65):
66 """Run a numba function in batches with progress tracking.
68 This helper enables progress monitoring for parallel numba operations by splitting
69 the work into batches and updating a progress bar between batch executions.
71 The progress bar shows total paths processed (jumping by batch_size each iteration)
72 rather than number of batches, giving better visibility into actual progress.
74 The numba function is called as:
75 numba_func(*fixed_args_before, *batched_slices, *fixed_args_after)
77 Args:
78 numba_func: The numba-compiled function to call
79 batched_arrays: List of arrays to slice per batch (must all have same length)
80 fixed_args_before: Tuple of args to pass before batched arrays
81 fixed_args_after: Tuple of args to pass after batched arrays
82 batch_size: Number of samples per batch
83 desc: Description for the progress bar
85 Returns:
86 Tuple of (paths, relations) concatenated from all batches
87 """
88 n_total = len(batched_arrays[0])
89 if n_total == 0:
90 return None, None
92 n_batches = (n_total + batch_size - 1) // batch_size
93 all_paths = []
94 all_rels = []
96 # Progress bar shows total paths, updated by batch_size each iteration
97 pbar = progress_bar(
98 total=n_total,
99 ncols=100,
100 desc=set_color(desc, "red", progress=True),
101 )
103 for batch_idx in range(n_batches):
104 batch_start = batch_idx * batch_size
105 batch_end = min(batch_start + batch_size, n_total)
106 actual_batch_size = batch_end - batch_start
108 # Slice all batched arrays
109 batch_slices = tuple(arr[batch_start:batch_end] for arr in batched_arrays)
111 # Call numba function
112 batch_paths, batch_rels = numba_func(*fixed_args_before, *batch_slices, *fixed_args_after)
114 all_paths.append(batch_paths)
115 all_rels.append(batch_rels)
117 # Update progress bar by actual batch size (jumps)
118 pbar.update(actual_batch_size)
120 pbar.close()
122 paths = np.concatenate(all_paths, axis=0)
123 path_rels = np.concatenate(all_rels, axis=0)
125 return paths, path_rels
128class KnowledgePathDataset(KnowledgeBasedDataset):
129 """:class:`KnowledgePathDataset` is based on :class:`~hopwise.data.dataset.KnowledgeBasedDataset`,
130 and provides an interface to prepare tokenized knowledge graph path for path language modeling.
132 Attributes:
133 path_hop_length (int): The same as ``config["path_hop_length"]``.
135 max_paths_per_user (int): The same as ``config["max_paths_per_user"]``.
137 temporal_causality (bool): The same as ``config["path_sample_args"]["temporal_causality"]``.
139 collaborative_path (bool): The same as ``config["path_sample_args"]["collaborative_path"]``.
141 strategy (str): The same as ``config["path_sample_args"]["strategy"]``.
143 path_token_separator (str): The same as ``config["path_sample_args"]["path_token_separator"]``.
145 restrict_by_phase (bool): The same as ``config["path_sample_args"]["restrict_by_phase"]``.
147 max_consecutive_invalid (int): The same as ``config["MAX_CONSECUTIVE_INVALID"]``.
149 tokenizer (PreTrainedTokenizerFast): Tokenizer to process the sample paths.
150 """
152 PATH_PADDING = -1
153 # Default batch size for progress tracking in numba-parallel operations
154 # Smaller batches = more frequent progress updates but slightly more overhead
155 PARALLEL_BATCH_SIZE = 50000
157 def __init__(self, config):
158 super().__init__(config)
159 self._path_dataset = None # path dataset is generated with generate_user_path_dataset
160 self._tokenized_dataset = None # tokenized path dataset is generated with tokenize_path_dataset
161 self._tokenizer = None
162 self.used_ids = None
164 # The starting index for auxiliary entity tokens.
165 # If None, it defaults to the number of items, meaning that all entities are considered.
166 # When KG includes user entities, this is set to `self.user_num + self.item_num`.
167 if self.config["tokenizer"].get("auxiliary_entity_start_id") is None:
168 # only entities that are not items are considered
169 self.config["tokenizer"]["auxiliary_entity_start_id"] = self.item_num
171 self._init_tokenizer()
173 def _get_field_from_config(self):
174 super()._get_field_from_config()
176 self.context_length = self.config["context_length"]
178 # Path sampling parameters
179 self.path_hop_length = self.config["path_hop_length"]
180 assert self.path_hop_length % 2 == 1, "Path hop length must be odd"
181 self.max_paths_per_user = self.config["MAX_PATHS_PER_USER"]
183 # path_hop_length = n_relations => (n_relations + user_starting_node) + n_relations + 2 (BOS, EOS)
184 self.token_sequence_length = (1 + self.path_hop_length) + self.path_hop_length + 2
186 path_sample_args = self.config["path_sample_args"]
187 self.temporal_causality = path_sample_args["temporal_causality"]
188 self.collaborative_path = path_sample_args["collaborative_path"]
189 try:
190 self.strategy = PathSamplingStrategy(path_sample_args["strategy"])
191 except ValueError:
192 raise ValueError(
193 f"Invalid path sampling strategy [{path_sample_args['strategy']}]. "
194 f"Valid strategies are: {[s.value for s in PathSamplingStrategy]}"
195 )
196 self.path_token_separator = path_sample_args["path_token_separator"]
197 self.restrict_by_phase = path_sample_args["restrict_by_phase"]
198 self.max_consecutive_invalid = path_sample_args["MAX_CONSECUTIVE_INVALID"]
200 # Tokenizer parameters
201 self.tokenizer_model = self.config["tokenizer"]["model"]
203 # Special tokens
204 if self.config["tokenizer"]["special_tokens"] is not None:
205 for token_name, token_value in self.config["tokenizer"]["special_tokens"].items():
206 setattr(self, token_name, token_value)
207 self.special_tokens = list(self.config["tokenizer"]["special_tokens"].values())
208 else:
209 self.special_tokens = []
211 self.logger.debug(set_color("tokenizer", "blue") + f": {self.tokenizer_model}")
213 @property
214 def path_dataset(self):
215 if self._path_dataset is None:
216 raise ValueError("Path dataset has not been generated yet, build the dataset first.")
218 return self._path_dataset
220 @property
221 def tokenizer(self):
222 return self._tokenizer
224 @property
225 def tokenized_dataset(self):
226 if self._tokenized_dataset is None:
227 raise ValueError("Tokenized path dataset has not been generated yet, build the dataset first.")
229 return self._tokenized_dataset
231 def __len__(self):
232 """Return the length of the tokenized dataset."""
233 return len(self.tokenized_dataset)
235 def __getitem__(self, idx):
236 """Return the item at index `idx` from the tokenized dataset."""
237 if self._tokenized_dataset is None:
238 # It avoids issues with hopwise flops calculation.
239 dummy_data = self.tokenize(["U1"])
240 return Interaction(dummy_data.data)
242 return self.tokenized_dataset[idx]
244 def _init_tokenizer(self):
245 """Initialize the HuggingFace tokenizer.
247 Args:
248 auxiliary_entity_start_id (int, optional):
250 """
251 from tokenizers import Tokenizer, pre_tokenizers
252 from tokenizers import models as token_models
253 from tokenizers import processors as token_processors
254 from tokenizers import trainers as token_trainers
255 from transformers import PreTrainedTokenizerFast
257 tokenizer_model_class = getattr(token_models, self.tokenizer_model)
259 tokenizer_object = Tokenizer(tokenizer_model_class(unk_token=self.unk_token))
261 # Pre-tokenizer definition based on :attr:`path_token_separator`
262 tokenizer_object.pre_tokenizer = pre_tokenizers.Split(self.path_token_separator, "removed")
264 auxiliary_entity_start_id = self.config["tokenizer"]["auxiliary_entity_start_id"]
265 entity_range = np.arange(auxiliary_entity_start_id, self.entity_num)
266 token_vocab = np.concatenate(
267 [
268 np.char.add(PathLanguageModelingTokenType.USER.token, np.arange(self.user_num).astype(str)),
269 np.char.add(PathLanguageModelingTokenType.ITEM.token, np.arange(self.item_num).astype(str)),
270 np.char.add(PathLanguageModelingTokenType.ENTITY.token, entity_range.astype(str)),
271 np.char.add(PathLanguageModelingTokenType.RELATION.token, np.arange(self.relation_num).astype(str)),
272 ]
273 )
275 tokenizer_trainer_class = getattr(token_trainers, self.tokenizer_model + "Trainer")
276 tokenizer_trainer = tokenizer_trainer_class(
277 vocab_size=len(token_vocab) + len(self.special_tokens), special_tokens=self.special_tokens
278 )
280 tokenizer_object.train_from_iterator(token_vocab, trainer=tokenizer_trainer)
282 tokenizer_object.post_processor = token_processors.TemplateProcessing(
283 single=f"{self.bos_token} $A {self.eos_token}",
284 special_tokens=[
285 (spec_token, tokenizer_object.token_to_id(spec_token))
286 for spec_token in [self.bos_token, self.eos_token]
287 ],
288 )
289 self._tokenizer = PreTrainedTokenizerFast(
290 tokenizer_object=tokenizer_object,
291 model_max_length=self.context_length,
292 eos_token=self.eos_token,
293 bos_token=self.bos_token,
294 pad_token=self.pad_token,
295 unk_token=self.unk_token,
296 mask_token=self.mask_token,
297 )
299 def _igraph_triple_to_tokenizer_triple(
300 self, vertex_metadata, igraph_head, igraph_relation, igraph_tail, token_vocab=None
301 ):
302 """Convert igraph ids to tokenizer ids."""
303 if token_vocab is None:
304 token_vocab = self.tokenizer.get_vocab()
306 ret = []
307 triple = [igraph_head, igraph_relation, igraph_tail]
308 for term, term_type in zip(triple, ["node", "relation", "node"]):
309 term_id = term
310 if term_type == "node":
311 if vertex_metadata[term_id]["type"] == self.uid_field:
312 prefix = PathLanguageModelingTokenType.USER.token
313 elif vertex_metadata[term_id]["type"] == self.iid_field:
314 term_id -= self.user_num
315 prefix = PathLanguageModelingTokenType.ITEM.token
316 elif vertex_metadata[term_id]["type"] == self.entity_field:
317 prefix = PathLanguageModelingTokenType.ENTITY.token
318 if self.config["tokenizer"]["auxiliary_entity_start_id"] == self.item_num:
319 # it means the KG does not have user nodes, but the igraph graph
320 # is a CKG with users so entity ids are shifted by user_num and
321 # we need to shift them back to match the tokenizer vocab
322 term_id -= self.user_num
323 else:
324 raise ValueError(
325 f"Unknown vertex type [{vertex_metadata[term_id]['type']}] "
326 "in igraph during tokenized_kg generation."
327 )
328 else:
329 prefix = PathLanguageModelingTokenType.RELATION.token
331 token_id = token_vocab[prefix + str(term_id)]
332 ret.append(token_id)
334 return ret
336 def get_tokenized_ckg(self):
337 """Return the tokenized collaborative knowledge graph.
339 We assume the any path is bidirectional except for user-item relations and :attr:`collaborative_path` is False.
341 Returns:
342 dict[dict[set]]: The tokenized collaborative knowledge graph.
343 """
344 graph = self._create_ckg_igraph(show_relation=True, directed=False)
345 vertex_metadata, edge_metadata = graph.to_dict_list()
346 token_vocab = self.tokenizer.get_vocab()
348 tokenized_kg = {}
349 for edge in edge_metadata:
350 head = edge["source"]
351 tail = edge["target"]
352 relation = edge["type"]
353 relation_id = self.field2token_id[self.relation_field][relation]
355 head_token, relation_token, tail_token = self._igraph_triple_to_tokenizer_triple(
356 vertex_metadata, head, relation_id, tail, token_vocab=token_vocab
357 )
359 # head is always the user in user-item relations. The check to add the reverse path is done later
360 if relation == self.ui_relation and vertex_metadata[head]["type"] != self.uid_field:
361 head_token, tail_token = tail_token, head_token
363 if head_token not in tokenized_kg:
364 tokenized_kg[head_token] = {}
365 if tail_token not in tokenized_kg:
366 tokenized_kg[tail_token] = {}
368 if relation_token not in tokenized_kg[head_token]:
369 tokenized_kg[head_token][relation_token] = set()
371 tokenized_kg[head_token][relation_token].add(tail_token)
373 if relation_token not in tokenized_kg[tail_token]:
374 tokenized_kg[tail_token][relation_token] = set()
376 tokenized_kg[tail_token][relation_token].add(head_token)
378 return tokenized_kg
380 def tokenize(self, data):
381 """Tokenize the input data using the tokenizer."""
382 return self.tokenizer(
383 data,
384 truncation=True,
385 padding=True,
386 max_length=self.context_length,
387 add_special_tokens=True,
388 return_token_type_ids=True,
389 )
391 def tokenize_path_dataset(self):
392 """Tokenize the path dataset."""
394 if self._tokenized_dataset is None:
395 tokenized_dataset = self.tokenize(self.path_dataset.split("\n"))
396 tokenized_dataset = Interaction(tokenized_dataset.data)
397 correct_path_mask = [
398 all(spec_token not in path[1:-1] for spec_token in self.tokenizer.all_special_ids)
399 for path in tokenized_dataset["input_ids"]
400 ]
401 tokenized_dataset = tokenized_dataset[correct_path_mask]
402 self._tokenized_dataset = tokenized_dataset
404 def build(self):
405 """Extends the build method to generate user path dataset and tokenize it."""
406 datasets = super().build()
407 datasets[0].generate_user_path_dataset()
408 datasets[0].tokenize_path_dataset()
410 return datasets
412 def get_tokenized_used_ids(self):
413 """Convert the used ids to tokenized ids.
415 Args:
416 used_ids: A numpy array of sets, where each set contains the item ids
417 that a user has interacted with.
418 tokenizer: The tokenizer to convert ids to tokenized ids.
419 Returns:
420 dict: A dictionary where keys are tokenized user ids and values are sets of tokenized item ids.
421 A numpy array of sets cannot be used as user tokens are not in the range [0, user_num].
422 """
423 user_token_type = PathLanguageModelingTokenType.USER.token
424 item_token_type = PathLanguageModelingTokenType.ITEM.token
426 used_ids = self.get_user_used_ids()
427 tokenized_used_ids = {}
428 for uid in range(used_ids.shape[0]):
429 uid_token = self.tokenizer.convert_tokens_to_ids(user_token_type + str(uid))
430 tokenized_used_ids[uid_token] = set(
431 [self.tokenizer.convert_tokens_to_ids(item_token_type + str(item)) for item in used_ids[uid]]
432 )
433 return tokenized_used_ids
435 def generate_user_path_dataset(self):
436 """Generate path dataset by sampling paths from the knowledge graph.
438 Paths represent walks in the graph that connect :attr:`hop_length` + 1 entities through
439 :attr:`hop_length` relations.
440 Each path connects two items. In the common scenario, the first item is a positive item
441 for the user and the second item is a recommendation candidate.
443 Refer to :meth:`generate_user_paths` for more details about path generation strategies.
444 """
445 if not isinstance(self.inter_feat, Interaction):
446 raise ValueError("The data should be prepared before generating the path dataset.")
448 if self._path_dataset is None:
449 generated_paths = self.generate_user_paths()
451 formatted_paths = [self._format_path(path) for path in generated_paths]
452 path_string = "\n".join(formatted_paths)
453 self._path_dataset = path_string
455 def generate_user_paths(self):
456 """Generate paths from the knowledge graph.
458 It currently supports three sampling strategies:
460 - weighted-rw: sampling-and-discarding approach through weighted random walk.
461 Paths are sampled from the knowledge graph and discarded if they are not valid, i.e.,
462 they do not end in a positive item
464 - constrained-rw: faithful random walk with constraints based on expected path output.
466 - simple-ui: per-interaction random-walk sampling. For every user-item interaction it samples up to
467 MAX_PATHS_PER_USER paths ending at other positive items, iteratively re-sampling uncovered interactions.
469 Returns:
470 list: List of paths with relations.
471 """
472 temporal_matrix = None
473 if self.temporal_causality:
474 if self.time_field in self.inter_feat:
475 temporal_matrix = self.inter_matrix(value_field=self.time_field).toarray()
476 else:
477 self.logger.warning(
478 "time_field has not been loaded or set,"
479 "thus temporal causality will not be used for path generation."
480 )
482 used_ids = self.get_user_used_ids()
484 csr_matrix = self._create_ckg_sparse_matrix(form="csr", show_relation=True)
485 csr_graph = CSRGraph.from_sparse_matrix(csr_matrix)
487 if self.strategy == PathSamplingStrategy.WEIGHTED_RW:
488 max_tries_per_iid = self.config["path_sample_args"]["MAX_RW_TRIES_PER_IID"]
489 paths_with_relations = self._generate_user_paths_weighted_random_walk(
490 csr_graph, used_ids, temporal_matrix=temporal_matrix, max_tries_per_iid=max_tries_per_iid
491 )
492 elif self.strategy == PathSamplingStrategy.CONSTRAINED_RW:
493 max_paths_per_hop = self.config["path_sample_args"]["MAX_RW_PATHS_PER_HOP"]
494 paths_with_relations = self._generate_user_paths_constrained_random_walk(
495 csr_graph, used_ids, temporal_matrix=temporal_matrix, paths_per_hop=max_paths_per_hop
496 )
497 elif self.strategy == PathSamplingStrategy.SIMPLE_UI:
498 paths_with_relations = self._generate_user_paths_all_simple_ui(
499 csr_graph, used_ids, temporal_matrix=temporal_matrix
500 )
501 else:
502 raise NotImplementedError(f"Path generation method [{self.strategy}] has not been implemented.")
504 return paths_with_relations
506 def _generate_user_paths_weighted_random_walk(
507 self, csr_graph, used_ids, temporal_matrix=None, max_tries_per_iid=50
508 ):
509 """Generate paths from the knowledge graph using weighted random walk with CSR + numba.
511 Uses batched parallel numba implementation for efficient path generation.
513 Args:
514 csr_graph: CSRGraph instance containing graph arrays
515 used_ids: Array of sets containing positive item ids per user
516 temporal_matrix: Optional temporal ordering matrix
517 max_tries_per_iid: Maximum attempts per starting item
518 """
519 indptr, indices, relations = csr_graph.unpack()
520 path_hop_length = self.path_hop_length - 2 # First and last hops handled separately
521 graph_min_iid = np.int64(self.user_num)
522 graph_max_iid = np.int64(self.item_num - 1 + self.user_num)
524 # Prepare all start nodes and user mappings for batch processing
525 all_start_nodes = []
526 all_user_ids = []
527 all_item_candidates = [] # For restrict_by_phase
529 self.logger.info(set_color("Preparing batch data for weighted-rw...", "blue"))
531 for u in range(1, self.user_num):
532 pos_iid = np.array(list(used_ids[u]), dtype=np.int64)
533 if len(pos_iid) == 0:
534 continue
536 if temporal_matrix is not None:
537 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])]
539 pos_iid_graph = pos_iid + self.user_num
541 if temporal_matrix is not None and len(pos_iid) <= 1:
542 continue
543 pos_iid_range = len(pos_iid) - (1 if temporal_matrix is not None else 0)
545 # Generate multiple start nodes per user for batch processing
546 n_samples_per_user = min(self.max_paths_per_user * max_tries_per_iid, pos_iid_range * max_tries_per_iid)
547 start_indices = np.random.randint(0, pos_iid_range, size=n_samples_per_user)
548 start_nodes = pos_iid_graph[start_indices]
550 all_start_nodes.extend(start_nodes)
551 all_user_ids.extend([u] * n_samples_per_user)
553 # Store candidate info for each start
554 for idx in start_indices:
555 if self.restrict_by_phase:
556 if temporal_matrix is not None:
557 candidates = pos_iid_graph[idx + 1 :]
558 else:
559 candidates = np.concatenate([pos_iid_graph[:idx], pos_iid_graph[idx + 1 :]])
560 all_item_candidates.append(set(candidates))
561 else:
562 all_item_candidates.append(None)
564 if len(all_start_nodes) == 0:
565 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
567 all_start_nodes = np.array(all_start_nodes, dtype=np.int64)
568 all_user_ids = np.array(all_user_ids, dtype=np.int64)
570 # Run random walks in batches with progress tracking
571 paths, path_rels = _run_batched_numba(
572 _csr_parallel_random_walks,
573 batched_arrays=[all_start_nodes],
574 fixed_args_before=(indptr, indices, relations),
575 fixed_args_after=(path_hop_length, graph_min_iid, self.collaborative_path),
576 batch_size=self.PARALLEL_BATCH_SIZE,
577 desc="KG Path Sampling (weighted-rw)",
578 )
580 # Filter valid paths (no -1 in the middle)
581 valid_mask = paths[:, -1] != -1
583 # Check path structure: first node must be item, intermediate nodes valid
584 first_node_valid = (paths[:, 0] >= graph_min_iid) & (paths[:, 0] <= graph_max_iid)
585 valid_mask &= first_node_valid
587 if not self.collaborative_path and path_hop_length > 1:
588 # Intermediate nodes must not be users
589 intermediate = paths[:, 1:-1]
590 intermediate_valid = (intermediate >= graph_min_iid).all(axis=1) | (intermediate == -1).any(axis=1)
591 valid_mask &= intermediate_valid
593 valid_indices = np.where(valid_mask)[0]
595 # Process valid paths and add final hop
596 all_final_paths = []
597 user_path_counts = {}
598 ui_rel_id = self.field2token_id[self.relation_field][self.ui_relation]
600 # Calculate target: max_paths_per_user * num_users for early stopping
601 target_total_paths = self.max_paths_per_user * self.user_num
603 pbar_validation = progress_bar(
604 valid_indices,
605 total=len(valid_indices),
606 ncols=100,
607 desc=set_color("Path Validation (weighted-rw)", "red", progress=True),
608 )
610 for idx in pbar_validation:
611 u = all_user_ids[idx]
613 # Check if user already has enough paths
614 if user_path_counts.get(u, 0) >= self.max_paths_per_user:
615 continue
617 path = paths[idx]
618 path_rel = path_rels[idx]
619 start_node = all_start_nodes[idx]
620 item_candidates = all_item_candidates[idx]
622 # Find valid last hop to an item
623 last_node = path[-1]
624 neighbor_start = indptr[last_node]
625 neighbor_end = indptr[last_node + 1]
627 if neighbor_end <= neighbor_start:
628 continue
630 neighbor_nodes = indices[neighbor_start:neighbor_end]
631 neighbor_rels = relations[neighbor_start:neighbor_end]
633 # Filter to item nodes only
634 item_mask = (neighbor_nodes >= graph_min_iid) & (neighbor_nodes <= graph_max_iid)
635 item_mask &= neighbor_nodes != start_node
637 if item_candidates is not None:
638 item_mask &= np.array([n in item_candidates for n in neighbor_nodes])
640 valid_neighbors = neighbor_nodes[item_mask]
641 valid_rels = neighbor_rels[item_mask]
643 if len(valid_neighbors) == 0:
644 continue
646 choice_idx = np.random.randint(len(valid_neighbors))
647 final_node = valid_neighbors[choice_idx]
648 final_rel = valid_rels[choice_idx]
650 # Build interleaved path: user, rel, item, rel, ..., item
651 path_with_rels = [u, ui_rel_id]
653 for i in range(len(path)):
654 path_with_rels.append(path[i])
655 if i < len(path_rel):
656 path_with_rels.append(path_rel[i])
658 path_with_rels.append(final_rel)
659 path_with_rels.append(final_node)
661 path_tuple = tuple(path_with_rels)
662 all_final_paths.append(path_tuple)
663 user_path_counts[u] = user_path_counts.get(u, 0) + 1
665 # Early stopping: if we've collected enough paths, stop iterating
666 if len(all_final_paths) >= target_total_paths:
667 break
669 # Deduplicate and convert to array. Sort for consistency.
670 unique_paths = sorted(list(set(all_final_paths)))
672 if len(unique_paths) == 0:
673 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
675 complete_path_length = self.path_hop_length * 2 + 1
676 paths_array = np.full((len(unique_paths), complete_path_length), fill_value=self.PATH_PADDING, dtype=np.int64)
677 for i, path in enumerate(unique_paths):
678 paths_array[i, : len(path)] = path
680 return paths_array
682 def _generate_user_paths_constrained_random_walk(self, csr_graph, used_ids, temporal_matrix=None, paths_per_hop=1):
683 """Generate paths from the knowledge graph using constrained random walks with CSR + numba.
685 Uses batched parallel numba implementation for efficient path generation.
686 The walk is constrained to follow entity types (items -> entities -> items).
688 Args:
689 csr_graph: CSRGraph instance containing graph arrays
690 used_ids: Array of sets containing positive item ids per user
691 temporal_matrix: Optional temporal ordering matrix
692 paths_per_hop: Number of paths to sample at each hop branching (used for oversampling)
693 """
694 indptr, indices, relations = csr_graph.unpack()
695 path_hop_length = self.path_hop_length - 1 # First hop (user-item) handled separately
696 graph_min_iid = np.int64(self.user_num)
697 graph_max_iid = np.int64(self.item_num - 1 + self.user_num)
699 # Prepare all start nodes and user mappings for batch processing
700 all_start_nodes = []
701 all_user_ids = []
702 all_start_node_indices = [] # Index within user's pos_iid array
704 # Build flattened positive item arrays for candidate lookup
705 pos_iid_flat = []
706 pos_iid_offsets = [0] # pos_iid_offsets[u] gives start index in pos_iid_flat for user u
708 self.logger.info(set_color("Preparing batch data for constrained-rw...", "blue"))
710 for u in range(self.user_num):
711 if u == 0:
712 pos_iid_offsets.append(0)
713 continue
715 pos_iid = np.array(list(used_ids[u]), dtype=np.int64)
716 if len(pos_iid) == 0:
717 pos_iid_offsets.append(pos_iid_offsets[-1])
718 continue
720 if temporal_matrix is not None:
721 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])]
723 pos_iid_graph = pos_iid + self.user_num
724 pos_iid_flat.extend(pos_iid_graph)
725 pos_iid_offsets.append(len(pos_iid_flat))
727 if temporal_matrix is not None and len(pos_iid) <= 1:
728 continue
729 pos_iid_range = len(pos_iid) - (1 if temporal_matrix is not None else 0)
731 # Generate multiple start nodes per user for batch processing
732 # Oversample to account for invalid paths
733 n_samples_per_user = self.max_paths_per_user * self.max_consecutive_invalid * paths_per_hop
734 start_indices = np.random.randint(0, pos_iid_range, size=n_samples_per_user)
735 start_nodes = pos_iid_graph[start_indices]
737 all_start_nodes.extend(start_nodes)
738 all_user_ids.extend([u] * n_samples_per_user)
739 all_start_node_indices.extend(start_indices)
741 if len(all_start_nodes) == 0:
742 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
744 all_start_nodes = np.array(all_start_nodes, dtype=np.int64)
745 all_user_ids = np.array(all_user_ids, dtype=np.int64)
746 all_start_node_indices = np.array(all_start_node_indices, dtype=np.int64)
747 pos_iid_flat = np.array(pos_iid_flat, dtype=np.int64)
748 pos_iid_offsets = np.array(pos_iid_offsets, dtype=np.int64)
750 # Run constrained random walks in batches with progress tracking
751 paths, path_rels = _run_batched_numba(
752 numba_func=_csr_constrained_random_walks,
753 batched_arrays=[all_start_nodes],
754 fixed_args_before=(indptr, indices, relations),
755 fixed_args_after=(
756 path_hop_length,
757 graph_min_iid,
758 graph_max_iid,
759 self.collaborative_path,
760 ),
761 batch_size=self.PARALLEL_BATCH_SIZE,
762 desc="KG Path Sampling (constrained-rw)",
763 )
765 # Filter valid paths (completed walks that end at items)
766 valid_mask = paths[:, -1] != -1
767 valid_mask &= (paths[:, -1] >= graph_min_iid) & (paths[:, -1] <= graph_max_iid)
768 # Ensure end node is different from start node
769 valid_mask &= paths[:, -1] != all_start_nodes
771 # Apply restrict_by_phase filtering if needed
772 if self.restrict_by_phase:
773 phase_filter_indices = np.where(valid_mask)[0]
774 pbar_phase = progress_bar(
775 phase_filter_indices,
776 total=len(phase_filter_indices),
777 ncols=100,
778 desc=set_color("Phase Filtering (constrained-rw)", "red", progress=True),
779 )
780 for idx in pbar_phase:
781 u = all_user_ids[idx]
782 start_idx = all_start_node_indices[idx]
783 end_node = paths[idx, -1]
785 # Get user's positive items
786 user_pos_start = pos_iid_offsets[u]
787 user_pos_end = pos_iid_offsets[u + 1]
788 user_pos_items = pos_iid_flat[user_pos_start:user_pos_end]
790 # Check if end_node is a valid candidate
791 if temporal_matrix is not None:
792 # Only items after start_idx are valid
793 valid_candidates = user_pos_items[start_idx + 1 :]
794 else:
795 # All items except the start item are valid
796 valid_candidates = (
797 np.concatenate([user_pos_items[:start_idx], user_pos_items[start_idx + 1 :]])
798 if len(user_pos_items) > 1
799 else np.array([], dtype=np.int64)
800 )
802 if end_node not in valid_candidates:
803 valid_mask[idx] = False
805 valid_indices = np.where(valid_mask)[0]
807 # Build final paths with user and relations
808 all_final_paths = []
809 user_path_counts = {}
810 ui_rel_id = self.field2token_id[self.relation_field][self.ui_relation]
812 # Calculate target: max_paths_per_user * num_users for early stopping
813 target_total_paths = self.max_paths_per_user * self.user_num
815 pbar_validation = progress_bar(
816 valid_indices,
817 total=len(valid_indices),
818 ncols=100,
819 desc=set_color("Path Validation (constrained-rw)", "red", progress=True),
820 )
822 for idx in pbar_validation:
823 u = all_user_ids[idx]
825 # Check if user already has enough paths
826 if user_path_counts.get(u, 0) >= self.max_paths_per_user:
827 continue
829 path = paths[idx]
830 path_rel = path_rels[idx]
832 # Build interleaved path: user, rel, item, rel, ..., item
833 path_with_rels = [u, ui_rel_id]
835 for i in range(len(path)):
836 if path[i] == -1:
837 break
838 path_with_rels.append(path[i])
839 if i < len(path_rel) and path_rel[i] != -1:
840 path_with_rels.append(path_rel[i])
842 path_tuple = tuple(path_with_rels)
843 all_final_paths.append(path_tuple)
844 user_path_counts[u] = user_path_counts.get(u, 0) + 1
846 # Early stopping: if we've collected enough paths, stop iterating
847 if len(all_final_paths) >= target_total_paths:
848 break
850 # Deduplicate and convert to array. Sort for consistency.
851 unique_paths = sorted(list(set(all_final_paths)))
853 if len(unique_paths) == 0:
854 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
856 complete_path_length = self.path_hop_length * 2 + 1
857 paths_array = np.full((len(unique_paths), complete_path_length), fill_value=self.PATH_PADDING, dtype=np.int64)
858 for i, path in enumerate(unique_paths):
859 paths_array[i, : len(path)] = path
861 return paths_array
863 def _generate_user_paths_all_simple_ui(self, csr_graph, used_ids, temporal_matrix=None):
864 """Generate paths from users to their positive items using parallel random walks.
866 This method uses parallel random walks with iterative re-sampling to find paths
867 connecting positive items. It keeps retrying until all paths are found or
868 no progress is made for max_consecutive_invalid consecutive attempts.
870 Strategy:
871 1. Run parallel random walks from all positive items needing paths
872 2. Check coverage: count pairs that still need paths
873 3. Re-sample: if missing pairs count unchanged for max_consecutive_invalid attempts, stop
874 4. Otherwise, continue until all pairs are satisfied
876 Args:
877 csr_graph: CSRGraph instance containing graph arrays
878 used_ids: Array of sets containing positive item ids per user
879 temporal_matrix: Optional temporal ordering matrix
880 """
881 indptr, indices, relations = csr_graph.unpack()
882 path_hop_length = self.path_hop_length - 1 # First hop (user-item) handled separately
883 graph_min_iid = np.int64(self.user_num)
884 graph_max_iid = np.int64(self.item_num - 1 + self.user_num)
886 # Build user data structures
887 user_pos_items = {} # user -> list of positive item graph IDs
888 user_pos_set = {} # user -> set of positive item graph IDs (for fast lookup)
890 self.logger.info(set_color("Preparing batch data for simple-ui...", "blue"))
892 for u in range(1, self.user_num):
893 pos_iid = np.array(list(used_ids[u]), dtype=np.int64)
894 if len(pos_iid) == 0:
895 continue
897 if temporal_matrix is not None:
898 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])]
900 pos_iid_graph = pos_iid + self.user_num
901 user_pos_items[u] = pos_iid_graph
902 user_pos_set[u] = set(pos_iid_graph)
904 if len(user_pos_items) == 0:
905 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
907 # Paths per (user, positive_item) pair
908 paths_per_pair = self.max_paths_per_user
910 # Info and warning messages for simple-ui behavior
911 self.logger.info(
912 set_color(
913 f"simple-ui: max_paths_per_user ({paths_per_pair}) limits paths per (user, positive_item) pair",
914 "blue",
915 )
916 )
918 # Warn if paths_per_pair seems high relative to average user actions
919 if paths_per_pair > 0.1 * self.avg_actions_of_users:
920 self.logger.warning(
921 set_color(
922 f"max_paths_per_user ({paths_per_pair}) > 10% of avg_actions_of_users in training set "
923 f"({self.avg_actions_of_users:.2f}). This may result in excessive path generation.",
924 "yellow",
925 )
926 )
928 # simple-ui forces restrict_by_phase by design (paths must connect positive items in order)
929 if not self.restrict_by_phase:
930 self.logger.info(set_color("simple-ui forces restrict_by_phase=True by design", "blue"))
932 # Track paths per (user, start_item) pair
933 # Key: (user_id, start_item_graph_id), Value: set of path tuples
934 user_item_paths = {}
935 ui_rel_id = self.field2token_id[self.relation_field][self.ui_relation]
937 # Initialize tracking for all (user, start_item) pairs
938 total_pairs = sum(len(pos_items) for pos_items in user_pos_items.values())
939 for u, pos_items in user_pos_items.items():
940 for start_item in pos_items:
941 user_item_paths[(u, start_item)] = set()
943 # Target: total paths we want to find
944 target_total_paths = total_pairs * paths_per_pair
946 # Tracking for early stopping
947 consecutive_no_progress = 0
948 prev_missing_pairs_count = total_pairs
949 attempt = 0
951 # Create a single progress bar showing total paths found
952 pbar = progress_bar(
953 total=target_total_paths,
954 ncols=100,
955 desc=set_color("KG Path Sampling (simple-ui)", "red", progress=True),
956 )
957 current_total_paths = 0
959 # Iterative sampling until all paths found or no progress for max_consecutive_invalid attempts
960 while True:
961 attempt += 1
963 # Find pairs that still need more paths
964 pairs_needing_paths = [
965 (u, start_item) for (u, start_item), paths in user_item_paths.items() if len(paths) < paths_per_pair
966 ]
968 missing_pairs_count = len(pairs_needing_paths)
970 if missing_pairs_count == 0:
971 self.logger.info(set_color(f"All pairs satisfied after {attempt} attempts", "green"))
972 break
974 # Check for progress: if missing pairs count unchanged, increment counter
975 if missing_pairs_count == prev_missing_pairs_count:
976 consecutive_no_progress += 1
977 if consecutive_no_progress >= self.max_consecutive_invalid:
978 self.logger.info(
979 set_color(
980 f"No progress for {self.max_consecutive_invalid} consecutive attempts, "
981 f"stopping with {missing_pairs_count} pairs still missing paths",
982 "yellow",
983 )
984 )
985 break
986 else:
987 # Progress was made, reset counter
988 consecutive_no_progress = 0
989 prev_missing_pairs_count = missing_pairs_count
991 # Prepare batch data for this iteration
992 all_start_nodes = []
993 all_user_ids = []
995 # Oversample more aggressively to find paths faster
996 samples_per_pair = max(4, paths_per_pair * 4)
998 for u, start_item in pairs_needing_paths:
999 needed = paths_per_pair - len(user_item_paths[(u, start_item)])
1000 n_samples = needed * samples_per_pair
1001 all_start_nodes.extend([start_item] * n_samples)
1002 all_user_ids.extend([u] * n_samples)
1004 if len(all_start_nodes) == 0:
1005 break
1007 all_start_nodes = np.array(all_start_nodes, dtype=np.int64)
1008 all_user_ids = np.array(all_user_ids, dtype=np.int64)
1010 # Run parallel random walks (without individual progress bar - we have the outer one)
1011 paths, path_rels = _csr_constrained_random_walks(
1012 indptr,
1013 indices,
1014 relations,
1015 all_start_nodes,
1016 path_hop_length,
1017 graph_min_iid,
1018 graph_max_iid,
1019 self.collaborative_path,
1020 )
1022 # Filter and validate paths
1023 valid_mask = paths[:, -1] != -1
1024 valid_mask &= (paths[:, -1] >= graph_min_iid) & (paths[:, -1] <= graph_max_iid)
1025 valid_mask &= paths[:, -1] != all_start_nodes # End different from start
1027 # Single merged validation loop: check end node is positive item AND respects temporal order
1028 # simple-ui always enforces restrict_by_phase behavior
1029 for idx in np.where(valid_mask)[0]:
1030 u = all_user_ids[idx]
1031 start_node = all_start_nodes[idx]
1032 end_node = paths[idx, -1]
1034 # End node must be in user's positive items
1035 if end_node not in user_pos_set[u]:
1036 valid_mask[idx] = False
1037 continue
1039 # Enforce restrict_by_phase: end must be valid relative to start
1040 pos_items = user_pos_items[u]
1041 start_idx = np.where(pos_items == start_node)[0]
1042 if len(start_idx) == 0:
1043 valid_mask[idx] = False
1044 continue
1045 start_idx = start_idx[0]
1047 if temporal_matrix is not None:
1048 # End must come after start in temporal order
1049 valid_candidates = set(pos_items[start_idx + 1 :])
1050 else:
1051 # Any other positive item is valid
1052 valid_candidates = set(pos_items) - {start_node}
1054 if end_node not in valid_candidates:
1055 valid_mask[idx] = False
1057 valid_indices = np.where(valid_mask)[0]
1059 # Process valid paths
1060 new_paths_found = 0
1061 for idx in valid_indices:
1062 u = all_user_ids[idx]
1063 start_node = all_start_nodes[idx]
1064 key = (u, start_node)
1066 # Check if this pair already has enough paths
1067 if len(user_item_paths[key]) >= self.max_paths_per_user:
1068 continue
1070 path = paths[idx]
1071 path_rel = path_rels[idx]
1073 # Build interleaved path: user, ui_rel, item, rel, ..., item
1074 path_with_rels = [u, ui_rel_id]
1076 for i in range(len(path)):
1077 if path[i] == -1:
1078 break
1079 path_with_rels.append(path[i])
1080 if i < len(path_rel) and path_rel[i] != -1:
1081 path_with_rels.append(path_rel[i])
1083 path_tuple = tuple(path_with_rels)
1084 if path_tuple not in user_item_paths[key]:
1085 user_item_paths[key].add(path_tuple)
1086 new_paths_found += 1
1088 # Update progress bar with new paths found
1089 if new_paths_found > 0:
1090 current_total_paths += new_paths_found
1091 pbar.update(new_paths_found)
1092 # Update description with attempt info
1093 pbar.set_description(
1094 set_color(f"KG Path Sampling (simple-ui, attempt {attempt})", "red", progress=True)
1095 )
1097 pbar.close()
1099 # Collect all paths - each pair is already limited to max_paths_per_user during sampling
1100 all_final_paths = [path_tuple for paths_set in user_item_paths.values() for path_tuple in paths_set]
1101 # Sort for consistency.
1102 all_final_paths = sorted(all_final_paths)
1104 if len(all_final_paths) == 0:
1105 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1)
1107 complete_path_length = self.path_hop_length * 2 + 1
1108 paths_array = np.full(
1109 (len(all_final_paths), complete_path_length), fill_value=self.PATH_PADDING, dtype=np.int64
1110 )
1111 for i, path in enumerate(all_final_paths):
1112 paths_array[i, : len(path)] = path
1114 return paths_array
1116 @staticmethod
1117 def _check_kg_path(path, user_num, item_num, check_last_node=False, collaborative_path=False):
1118 """Check if the path is valid. The first node must be an item node and it assumes the user node is omitted.
1120 Args:
1121 path (list): The path to be checked.
1123 check_last_node (bool, optional): Whether to check the last node in the path.
1124 Defaults to ``False``.
1125 """
1126 path = np.array(path, dtype=int)
1127 graph_min_iid = 1 + user_num
1128 graph_max_iid = item_num - 1 + user_num
1130 user_check = path[0] < graph_min_iid
1131 pos_iid_check = graph_min_iid <= path[1] <= graph_max_iid
1132 valid_path = (path[2:-1] >= graph_min_iid).all() or collaborative_path
1133 check_rec_iid = not check_last_node or graph_min_iid <= path[-1] <= graph_max_iid
1135 return user_check and pos_iid_check and valid_path and check_rec_iid
1137 def _format_path(self, path):
1138 """Format the path to a string according to :class:`~hopwise.utils.enum_type.PathLanguageModelingTokenType`.
1140 Args:
1141 path (list): The path to be formatted.
1142 """
1143 path = path[path != self.PATH_PADDING] # remove padding for shorter paths
1144 print
1145 path_nodes = path[::2]
1146 path_relations = path[1::2]
1148 remapped_path_nodes = []
1149 graph_min_iid = self.user_num
1150 graph_max_iid = self.item_num - 1 + self.user_num
1151 for node in path_nodes:
1152 if graph_min_iid <= node <= graph_max_iid:
1153 remapped_path_nodes.append(PathLanguageModelingTokenType.ITEM.token + str(node - self.user_num))
1154 elif node < graph_min_iid:
1155 remapped_path_nodes.append(PathLanguageModelingTokenType.USER.token + str(node))
1156 else:
1157 remapped_path_nodes.append(PathLanguageModelingTokenType.ENTITY.token + str(node - self.user_num))
1159 relation_mapped_list = [PathLanguageModelingTokenType.RELATION.token + str(r) for r in path_relations]
1161 interleaved_entities_relations = zip_longest(remapped_path_nodes, relation_mapped_list)
1162 path_string = self.path_token_separator.join(list(chain(*interleaved_entities_relations))[:-1])
1164 return path_string
1166 def __str__(self):
1167 info = [
1168 super().__str__(),
1169 f"The number of hops used for path sampling: {self.path_hop_length}",
1170 f"Maximum number of paths sampled per user: {self.max_paths_per_user}",
1171 f"The path sampling strategy: {self.strategy}",
1172 f"The tokenizer model: {self.tokenizer_model}",
1173 ]
1174 return "\n".join(info)
1177class UserItemKnowledgePathDataset(KnowledgePathDataset, UserItemKnowledgeBasedDataset):
1178 """:class:`UserItemKnowledgePathDataset` is based on :class:`~hopwise.data.dataset.KnowledgePathDataset`,
1179 and :class:`~hopwise.data.dataset.UserItemKnowledgeBasedDataset` to be used with user-side KG too.
1180 It provides an interface to prepare tokenized knowledge graph path for path language modeling.
1182 Attributes:
1183 path_hop_length (int): The same as ``config["path_hop_length"]``.
1185 max_paths_per_user (int): The same as ``config["max_paths_per_user"]``.
1187 temporal_causality (bool): The same as ``config["path_sample_args"]["temporal_causality"]``.
1189 collaborative_path (bool): The same as ``config["path_sample_args"]["collaborative_path"]``.
1191 strategy (str): The same as ``config["path_sample_args"]["strategy"]``.
1193 reasoning_template (str): The same as ``config["path_sample_args"]["reasoning_template"]``.
1195 restrict_by_phase (bool): The same as ``config["path_sample_args"]["restrict_by_phase"]``.
1197 max_consecutive_invalid (int): The same as ``config["MAX_CONSECUTIVE_INVALID"]``.
1199 tokenizer (PreTrainedTokenizerFast): Tokenizer to process the sample paths.
1200 """
1202 def __init__(self, config):
1203 self._path_dataset = None
1204 self._tokenized_dataset = None
1205 self._tokenizer = None
1206 UserItemKnowledgeBasedDataset.__init__(self, config)
1208 config["tokenizer"]["auxiliary_entity_start_id"] = self.user_num + self.item_num
1209 KnowledgePathDataset.__init__(self, config)
1210 KnowledgePathDataset._get_field_from_config(self)
1213# ============================================================================
1214# Numba-accelerated helper functions for CSR-based random walks
1215# ============================================================================
1216@numba.njit(parallel=True)
1217def _csr_parallel_random_walks(indptr, indices, relations, start_nodes, num_steps, graph_min_iid, collaborative_path):
1218 """Parallel random walks on CSR graph.
1220 Note: Set np.random.seed() before calling this function for reproducibility.
1222 Args:
1223 indptr: CSR row pointers
1224 indices: CSR column indices
1225 relations: Edge relation types
1226 start_nodes: Array of starting nodes
1227 num_steps: Number of steps per walk
1228 graph_min_iid: Minimum item ID in graph (users have id < graph_min_iid)
1229 collaborative_path: If True, allow user nodes as intermediate nodes
1231 Returns:
1232 paths: (n_walks, num_steps + 1) node paths
1233 path_relations: (n_walks, num_steps) relation paths
1234 """
1235 n_walks = len(start_nodes)
1236 paths = np.full((n_walks, num_steps + 1), -1, dtype=np.int64)
1237 path_relations = np.full((n_walks, num_steps), -1, dtype=np.int64)
1239 for i in numba.prange(n_walks):
1240 node = start_nodes[i]
1241 paths[i, 0] = node
1243 for step in range(num_steps):
1244 start_idx = indptr[node]
1245 end_idx = indptr[node + 1]
1246 n_neighbors = end_idx - start_idx
1248 if n_neighbors == 0:
1249 break
1251 # Build list of valid neighbors
1252 valid_count = 0
1253 valid_indices = np.empty(n_neighbors, dtype=np.int64)
1255 for j in range(n_neighbors):
1256 neighbor = indices[start_idx + j]
1258 # If not collaborative_path, skip user nodes (id < graph_min_iid)
1259 if not collaborative_path and neighbor < graph_min_iid:
1260 continue
1262 valid_indices[valid_count] = j
1263 valid_count += 1
1265 if valid_count == 0:
1266 break
1268 # Uniform random selection from valid neighbors
1269 selected = np.random.randint(valid_count)
1270 selected_idx = valid_indices[selected]
1272 neighbor_idx = start_idx + selected_idx
1273 node = indices[neighbor_idx]
1274 paths[i, step + 1] = node
1275 path_relations[i, step] = relations[neighbor_idx]
1277 return paths, path_relations
1280@numba.njit(parallel=True)
1281def _csr_constrained_random_walks(
1282 indptr,
1283 indices,
1284 relations,
1285 start_nodes,
1286 num_steps,
1287 graph_min_iid,
1288 graph_max_iid,
1289 collaborative_path,
1290):
1291 """Parallel constrained random walks on CSR graph.
1293 Constraints:
1294 - Start from item
1295 - End at item (different from start)
1296 - If collaborative_path=True: user nodes CAN be intermediate nodes
1297 - If collaborative_path=False: user nodes are NOT allowed as intermediate nodes
1299 Args:
1300 indptr: CSR row pointers
1301 indices: CSR column indices
1302 relations: Edge relation types
1303 start_nodes: Array of starting nodes (items)
1304 num_steps: Number of steps per walk
1305 graph_min_iid: Minimum item ID in graph (users have id < graph_min_iid)
1306 graph_max_iid: Maximum item ID in graph
1307 collaborative_path: If True, allow user nodes as intermediate nodes
1309 Returns:
1310 paths: (n_walks, num_steps + 1) node paths
1311 path_relations: (n_walks, num_steps) relation paths
1312 """
1313 n_walks = len(start_nodes)
1314 paths = np.full((n_walks, num_steps + 1), -1, dtype=np.int64)
1315 path_relations = np.full((n_walks, num_steps), -1, dtype=np.int64)
1317 for i in numba.prange(n_walks):
1318 start_node = start_nodes[i]
1319 node = start_node
1320 paths[i, 0] = node
1322 # Track visited nodes to avoid cycles (max path length is small)
1323 visited = np.zeros(num_steps + 1, dtype=np.int64)
1324 visited[0] = node
1325 n_visited = 1
1327 for step in range(num_steps):
1328 start_idx = indptr[node]
1329 end_idx = indptr[node + 1]
1330 n_neighbors = end_idx - start_idx
1332 if n_neighbors == 0:
1333 break
1335 is_last_step = step == num_steps - 1
1337 # Build list of valid neighbors based on constraints
1338 valid_count = 0
1339 valid_indices = np.empty(n_neighbors, dtype=np.int64)
1341 for j in range(n_neighbors):
1342 neighbor = indices[start_idx + j]
1344 # Check if already visited
1345 is_visited = False
1346 for v in range(n_visited):
1347 if visited[v] == neighbor:
1348 is_visited = True
1349 break
1350 if is_visited:
1351 continue
1353 if is_last_step:
1354 # Last step: must end at item (not the start node)
1355 if neighbor >= graph_min_iid and neighbor <= graph_max_iid:
1356 if neighbor != start_node:
1357 valid_indices[valid_count] = j
1358 valid_count += 1
1359 # Intermediate step
1360 elif collaborative_path:
1361 # Allow any node (users, items, entities)
1362 valid_indices[valid_count] = j
1363 valid_count += 1
1364 # Only allow items and entities (no users)
1365 elif neighbor >= graph_min_iid:
1366 valid_indices[valid_count] = j
1367 valid_count += 1
1369 if valid_count == 0:
1370 break
1372 # Uniform random selection from valid neighbors
1373 selected = np.random.randint(valid_count)
1374 selected_idx = valid_indices[selected]
1376 neighbor_idx = start_idx + selected_idx
1377 node = indices[neighbor_idx]
1378 paths[i, step + 1] = node
1379 path_relations[i, step] = relations[neighbor_idx]
1381 visited[n_visited] = node
1382 n_visited += 1
1384 return paths, path_relations