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

1# @Time : 2025 

2# @Author : Giacomo Medda 

3# @Email : giacomo.medda@unica.it 

4 

5from dataclasses import dataclass 

6from itertools import chain, zip_longest 

7 

8import numba 

9import numpy as np 

10 

11from hopwise.data import Interaction 

12from hopwise.data.dataset import KnowledgeBasedDataset, UserItemKnowledgeBasedDataset 

13from hopwise.utils import PathLanguageModelingTokenType, PathSamplingStrategy, progress_bar, set_color 

14 

15 

16@dataclass 

17class CSRGraph: 

18 """Container for CSR graph arrays used in parallel random walks. 

19 

20 This dataclass bundles together the CSR sparse matrix components 

21 needed for efficient graph traversal in numba. 

22 

23 Attributes: 

24 indptr: CSR row pointers (int64) 

25 indices: CSR column indices (int64) 

26 relations: Edge relation types (int64) 

27 """ 

28 

29 indptr: np.ndarray 

30 indices: np.ndarray 

31 relations: np.ndarray 

32 

33 @classmethod 

34 def from_sparse_matrix(cls, csr_matrix): 

35 """Create CSRGraph from scipy sparse matrix with relation data. 

36 

37 Args: 

38 csr_matrix: Scipy CSR matrix with relations stored in data field 

39 

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) 

46 

47 return cls(indptr=indptr, indices=indices, relations=relations) 

48 

49 def unpack(self): 

50 """Unpack arrays for passing to numba functions. 

51 

52 Returns: 

53 tuple: (indptr, indices, relations) 

54 """ 

55 return self.indptr, self.indices, self.relations 

56 

57 

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. 

67 

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. 

70 

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. 

73 

74 The numba function is called as: 

75 numba_func(*fixed_args_before, *batched_slices, *fixed_args_after) 

76 

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 

84 

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 

91 

92 n_batches = (n_total + batch_size - 1) // batch_size 

93 all_paths = [] 

94 all_rels = [] 

95 

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 ) 

102 

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 

107 

108 # Slice all batched arrays 

109 batch_slices = tuple(arr[batch_start:batch_end] for arr in batched_arrays) 

110 

111 # Call numba function 

112 batch_paths, batch_rels = numba_func(*fixed_args_before, *batch_slices, *fixed_args_after) 

113 

114 all_paths.append(batch_paths) 

115 all_rels.append(batch_rels) 

116 

117 # Update progress bar by actual batch size (jumps) 

118 pbar.update(actual_batch_size) 

119 

120 pbar.close() 

121 

122 paths = np.concatenate(all_paths, axis=0) 

123 path_rels = np.concatenate(all_rels, axis=0) 

124 

125 return paths, path_rels 

126 

127 

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. 

131 

132 Attributes: 

133 path_hop_length (int): The same as ``config["path_hop_length"]``. 

134 

135 max_paths_per_user (int): The same as ``config["max_paths_per_user"]``. 

136 

137 temporal_causality (bool): The same as ``config["path_sample_args"]["temporal_causality"]``. 

138 

139 collaborative_path (bool): The same as ``config["path_sample_args"]["collaborative_path"]``. 

140 

141 strategy (str): The same as ``config["path_sample_args"]["strategy"]``. 

142 

143 path_token_separator (str): The same as ``config["path_sample_args"]["path_token_separator"]``. 

144 

145 restrict_by_phase (bool): The same as ``config["path_sample_args"]["restrict_by_phase"]``. 

146 

147 max_consecutive_invalid (int): The same as ``config["MAX_CONSECUTIVE_INVALID"]``. 

148 

149 tokenizer (PreTrainedTokenizerFast): Tokenizer to process the sample paths. 

150 """ 

151 

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 

156 

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 

163 

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 

170 

171 self._init_tokenizer() 

172 

173 def _get_field_from_config(self): 

174 super()._get_field_from_config() 

175 

176 self.context_length = self.config["context_length"] 

177 

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"] 

182 

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 

185 

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"] 

199 

200 # Tokenizer parameters 

201 self.tokenizer_model = self.config["tokenizer"]["model"] 

202 

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 = [] 

210 

211 self.logger.debug(set_color("tokenizer", "blue") + f": {self.tokenizer_model}") 

212 

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.") 

217 

218 return self._path_dataset 

219 

220 @property 

221 def tokenizer(self): 

222 return self._tokenizer 

223 

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.") 

228 

229 return self._tokenized_dataset 

230 

231 def __len__(self): 

232 """Return the length of the tokenized dataset.""" 

233 return len(self.tokenized_dataset) 

234 

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) 

241 

242 return self.tokenized_dataset[idx] 

243 

244 def _init_tokenizer(self): 

245 """Initialize the HuggingFace tokenizer. 

246 

247 Args: 

248 auxiliary_entity_start_id (int, optional): 

249 

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 

256 

257 tokenizer_model_class = getattr(token_models, self.tokenizer_model) 

258 

259 tokenizer_object = Tokenizer(tokenizer_model_class(unk_token=self.unk_token)) 

260 

261 # Pre-tokenizer definition based on :attr:`path_token_separator` 

262 tokenizer_object.pre_tokenizer = pre_tokenizers.Split(self.path_token_separator, "removed") 

263 

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 ) 

274 

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 ) 

279 

280 tokenizer_object.train_from_iterator(token_vocab, trainer=tokenizer_trainer) 

281 

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 ) 

298 

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() 

305 

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 

330 

331 token_id = token_vocab[prefix + str(term_id)] 

332 ret.append(token_id) 

333 

334 return ret 

335 

336 def get_tokenized_ckg(self): 

337 """Return the tokenized collaborative knowledge graph. 

338 

339 We assume the any path is bidirectional except for user-item relations and :attr:`collaborative_path` is False. 

340 

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() 

347 

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] 

354 

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 ) 

358 

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 

362 

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] = {} 

367 

368 if relation_token not in tokenized_kg[head_token]: 

369 tokenized_kg[head_token][relation_token] = set() 

370 

371 tokenized_kg[head_token][relation_token].add(tail_token) 

372 

373 if relation_token not in tokenized_kg[tail_token]: 

374 tokenized_kg[tail_token][relation_token] = set() 

375 

376 tokenized_kg[tail_token][relation_token].add(head_token) 

377 

378 return tokenized_kg 

379 

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 ) 

390 

391 def tokenize_path_dataset(self): 

392 """Tokenize the path dataset.""" 

393 

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 

403 

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() 

409 

410 return datasets 

411 

412 def get_tokenized_used_ids(self): 

413 """Convert the used ids to tokenized ids. 

414 

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 

425 

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 

434 

435 def generate_user_path_dataset(self): 

436 """Generate path dataset by sampling paths from the knowledge graph. 

437 

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. 

442 

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.") 

447 

448 if self._path_dataset is None: 

449 generated_paths = self.generate_user_paths() 

450 

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 

454 

455 def generate_user_paths(self): 

456 """Generate paths from the knowledge graph. 

457 

458 It currently supports three sampling strategies: 

459 

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 

463 

464 - constrained-rw: faithful random walk with constraints based on expected path output. 

465 

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. 

468 

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 ) 

481 

482 used_ids = self.get_user_used_ids() 

483 

484 csr_matrix = self._create_ckg_sparse_matrix(form="csr", show_relation=True) 

485 csr_graph = CSRGraph.from_sparse_matrix(csr_matrix) 

486 

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.") 

503 

504 return paths_with_relations 

505 

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. 

510 

511 Uses batched parallel numba implementation for efficient path generation. 

512 

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) 

523 

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 

528 

529 self.logger.info(set_color("Preparing batch data for weighted-rw...", "blue")) 

530 

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 

535 

536 if temporal_matrix is not None: 

537 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])] 

538 

539 pos_iid_graph = pos_iid + self.user_num 

540 

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) 

544 

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] 

549 

550 all_start_nodes.extend(start_nodes) 

551 all_user_ids.extend([u] * n_samples_per_user) 

552 

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) 

563 

564 if len(all_start_nodes) == 0: 

565 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

566 

567 all_start_nodes = np.array(all_start_nodes, dtype=np.int64) 

568 all_user_ids = np.array(all_user_ids, dtype=np.int64) 

569 

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 ) 

579 

580 # Filter valid paths (no -1 in the middle) 

581 valid_mask = paths[:, -1] != -1 

582 

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 

586 

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 

592 

593 valid_indices = np.where(valid_mask)[0] 

594 

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] 

599 

600 # Calculate target: max_paths_per_user * num_users for early stopping 

601 target_total_paths = self.max_paths_per_user * self.user_num 

602 

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 ) 

609 

610 for idx in pbar_validation: 

611 u = all_user_ids[idx] 

612 

613 # Check if user already has enough paths 

614 if user_path_counts.get(u, 0) >= self.max_paths_per_user: 

615 continue 

616 

617 path = paths[idx] 

618 path_rel = path_rels[idx] 

619 start_node = all_start_nodes[idx] 

620 item_candidates = all_item_candidates[idx] 

621 

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] 

626 

627 if neighbor_end <= neighbor_start: 

628 continue 

629 

630 neighbor_nodes = indices[neighbor_start:neighbor_end] 

631 neighbor_rels = relations[neighbor_start:neighbor_end] 

632 

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 

636 

637 if item_candidates is not None: 

638 item_mask &= np.array([n in item_candidates for n in neighbor_nodes]) 

639 

640 valid_neighbors = neighbor_nodes[item_mask] 

641 valid_rels = neighbor_rels[item_mask] 

642 

643 if len(valid_neighbors) == 0: 

644 continue 

645 

646 choice_idx = np.random.randint(len(valid_neighbors)) 

647 final_node = valid_neighbors[choice_idx] 

648 final_rel = valid_rels[choice_idx] 

649 

650 # Build interleaved path: user, rel, item, rel, ..., item 

651 path_with_rels = [u, ui_rel_id] 

652 

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]) 

657 

658 path_with_rels.append(final_rel) 

659 path_with_rels.append(final_node) 

660 

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 

664 

665 # Early stopping: if we've collected enough paths, stop iterating 

666 if len(all_final_paths) >= target_total_paths: 

667 break 

668 

669 # Deduplicate and convert to array. Sort for consistency. 

670 unique_paths = sorted(list(set(all_final_paths))) 

671 

672 if len(unique_paths) == 0: 

673 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

674 

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 

679 

680 return paths_array 

681 

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. 

684 

685 Uses batched parallel numba implementation for efficient path generation. 

686 The walk is constrained to follow entity types (items -> entities -> items). 

687 

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) 

698 

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 

703 

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 

707 

708 self.logger.info(set_color("Preparing batch data for constrained-rw...", "blue")) 

709 

710 for u in range(self.user_num): 

711 if u == 0: 

712 pos_iid_offsets.append(0) 

713 continue 

714 

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 

719 

720 if temporal_matrix is not None: 

721 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])] 

722 

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)) 

726 

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) 

730 

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] 

736 

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) 

740 

741 if len(all_start_nodes) == 0: 

742 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

743 

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) 

749 

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 ) 

764 

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 

770 

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] 

784 

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] 

789 

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 ) 

801 

802 if end_node not in valid_candidates: 

803 valid_mask[idx] = False 

804 

805 valid_indices = np.where(valid_mask)[0] 

806 

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] 

811 

812 # Calculate target: max_paths_per_user * num_users for early stopping 

813 target_total_paths = self.max_paths_per_user * self.user_num 

814 

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 ) 

821 

822 for idx in pbar_validation: 

823 u = all_user_ids[idx] 

824 

825 # Check if user already has enough paths 

826 if user_path_counts.get(u, 0) >= self.max_paths_per_user: 

827 continue 

828 

829 path = paths[idx] 

830 path_rel = path_rels[idx] 

831 

832 # Build interleaved path: user, rel, item, rel, ..., item 

833 path_with_rels = [u, ui_rel_id] 

834 

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]) 

841 

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 

845 

846 # Early stopping: if we've collected enough paths, stop iterating 

847 if len(all_final_paths) >= target_total_paths: 

848 break 

849 

850 # Deduplicate and convert to array. Sort for consistency. 

851 unique_paths = sorted(list(set(all_final_paths))) 

852 

853 if len(unique_paths) == 0: 

854 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

855 

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 

860 

861 return paths_array 

862 

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. 

865 

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. 

869 

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 

875 

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) 

885 

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) 

889 

890 self.logger.info(set_color("Preparing batch data for simple-ui...", "blue")) 

891 

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 

896 

897 if temporal_matrix is not None: 

898 pos_iid = pos_iid[np.argsort(temporal_matrix[u, pos_iid])] 

899 

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) 

903 

904 if len(user_pos_items) == 0: 

905 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

906 

907 # Paths per (user, positive_item) pair 

908 paths_per_pair = self.max_paths_per_user 

909 

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 ) 

917 

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 ) 

927 

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")) 

931 

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] 

936 

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() 

942 

943 # Target: total paths we want to find 

944 target_total_paths = total_pairs * paths_per_pair 

945 

946 # Tracking for early stopping 

947 consecutive_no_progress = 0 

948 prev_missing_pairs_count = total_pairs 

949 attempt = 0 

950 

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 

958 

959 # Iterative sampling until all paths found or no progress for max_consecutive_invalid attempts 

960 while True: 

961 attempt += 1 

962 

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 ] 

967 

968 missing_pairs_count = len(pairs_needing_paths) 

969 

970 if missing_pairs_count == 0: 

971 self.logger.info(set_color(f"All pairs satisfied after {attempt} attempts", "green")) 

972 break 

973 

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 

990 

991 # Prepare batch data for this iteration 

992 all_start_nodes = [] 

993 all_user_ids = [] 

994 

995 # Oversample more aggressively to find paths faster 

996 samples_per_pair = max(4, paths_per_pair * 4) 

997 

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) 

1003 

1004 if len(all_start_nodes) == 0: 

1005 break 

1006 

1007 all_start_nodes = np.array(all_start_nodes, dtype=np.int64) 

1008 all_user_ids = np.array(all_user_ids, dtype=np.int64) 

1009 

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 ) 

1021 

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 

1026 

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] 

1033 

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 

1038 

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] 

1046 

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} 

1053 

1054 if end_node not in valid_candidates: 

1055 valid_mask[idx] = False 

1056 

1057 valid_indices = np.where(valid_mask)[0] 

1058 

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) 

1065 

1066 # Check if this pair already has enough paths 

1067 if len(user_item_paths[key]) >= self.max_paths_per_user: 

1068 continue 

1069 

1070 path = paths[idx] 

1071 path_rel = path_rels[idx] 

1072 

1073 # Build interleaved path: user, ui_rel, item, rel, ..., item 

1074 path_with_rels = [u, ui_rel_id] 

1075 

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]) 

1082 

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 

1087 

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 ) 

1096 

1097 pbar.close() 

1098 

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) 

1103 

1104 if len(all_final_paths) == 0: 

1105 return np.array([], dtype=np.int64).reshape(0, self.path_hop_length * 2 + 1) 

1106 

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 

1113 

1114 return paths_array 

1115 

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. 

1119 

1120 Args: 

1121 path (list): The path to be checked. 

1122 

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 

1129 

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 

1134 

1135 return user_check and pos_iid_check and valid_path and check_rec_iid 

1136 

1137 def _format_path(self, path): 

1138 """Format the path to a string according to :class:`~hopwise.utils.enum_type.PathLanguageModelingTokenType`. 

1139 

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] 

1147 

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)) 

1158 

1159 relation_mapped_list = [PathLanguageModelingTokenType.RELATION.token + str(r) for r in path_relations] 

1160 

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]) 

1163 

1164 return path_string 

1165 

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) 

1175 

1176 

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. 

1181 

1182 Attributes: 

1183 path_hop_length (int): The same as ``config["path_hop_length"]``. 

1184 

1185 max_paths_per_user (int): The same as ``config["max_paths_per_user"]``. 

1186 

1187 temporal_causality (bool): The same as ``config["path_sample_args"]["temporal_causality"]``. 

1188 

1189 collaborative_path (bool): The same as ``config["path_sample_args"]["collaborative_path"]``. 

1190 

1191 strategy (str): The same as ``config["path_sample_args"]["strategy"]``. 

1192 

1193 reasoning_template (str): The same as ``config["path_sample_args"]["reasoning_template"]``. 

1194 

1195 restrict_by_phase (bool): The same as ``config["path_sample_args"]["restrict_by_phase"]``. 

1196 

1197 max_consecutive_invalid (int): The same as ``config["MAX_CONSECUTIVE_INVALID"]``. 

1198 

1199 tokenizer (PreTrainedTokenizerFast): Tokenizer to process the sample paths. 

1200 """ 

1201 

1202 def __init__(self, config): 

1203 self._path_dataset = None 

1204 self._tokenized_dataset = None 

1205 self._tokenizer = None 

1206 UserItemKnowledgeBasedDataset.__init__(self, config) 

1207 

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) 

1211 

1212 

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. 

1219 

1220 Note: Set np.random.seed() before calling this function for reproducibility. 

1221 

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 

1230 

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) 

1238 

1239 for i in numba.prange(n_walks): 

1240 node = start_nodes[i] 

1241 paths[i, 0] = node 

1242 

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 

1247 

1248 if n_neighbors == 0: 

1249 break 

1250 

1251 # Build list of valid neighbors 

1252 valid_count = 0 

1253 valid_indices = np.empty(n_neighbors, dtype=np.int64) 

1254 

1255 for j in range(n_neighbors): 

1256 neighbor = indices[start_idx + j] 

1257 

1258 # If not collaborative_path, skip user nodes (id < graph_min_iid) 

1259 if not collaborative_path and neighbor < graph_min_iid: 

1260 continue 

1261 

1262 valid_indices[valid_count] = j 

1263 valid_count += 1 

1264 

1265 if valid_count == 0: 

1266 break 

1267 

1268 # Uniform random selection from valid neighbors 

1269 selected = np.random.randint(valid_count) 

1270 selected_idx = valid_indices[selected] 

1271 

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] 

1276 

1277 return paths, path_relations 

1278 

1279 

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. 

1292 

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 

1298 

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 

1308 

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) 

1316 

1317 for i in numba.prange(n_walks): 

1318 start_node = start_nodes[i] 

1319 node = start_node 

1320 paths[i, 0] = node 

1321 

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 

1326 

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 

1331 

1332 if n_neighbors == 0: 

1333 break 

1334 

1335 is_last_step = step == num_steps - 1 

1336 

1337 # Build list of valid neighbors based on constraints 

1338 valid_count = 0 

1339 valid_indices = np.empty(n_neighbors, dtype=np.int64) 

1340 

1341 for j in range(n_neighbors): 

1342 neighbor = indices[start_idx + j] 

1343 

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 

1352 

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 

1368 

1369 if valid_count == 0: 

1370 break 

1371 

1372 # Uniform random selection from valid neighbors 

1373 selected = np.random.randint(valid_count) 

1374 selected_idx = valid_indices[selected] 

1375 

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] 

1380 

1381 visited[n_visited] = node 

1382 n_visited += 1 

1383 

1384 return paths, path_relations