Coverage for hopwise/cli.py: 0%

219 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1import importlib 

2import os 

3import sys 

4import traceback 

5from datetime import datetime 

6 

7import click 

8from rich import box 

9from rich.console import Console 

10from rich.panel import Panel 

11from rich.traceback import install 

12from setproctitle import setproctitle 

13 

14from hopwise.quick_start import objective_function, run 

15from hopwise.trainer import HyperTuning 

16from hopwise.utils import list_to_latex 

17 

18 

19class hopwiseClickCommand(click.Command): 

20 def parse_args(self, ctx, args): 

21 """Override to filter out HopWise parameters before Click's validation""" 

22 

23 click_args = [] 

24 hopwise_args = [] 

25 

26 for arg in args: 

27 if arg.startswith("--") and "=" in arg: 

28 hopwise_args.append(arg) 

29 else: 

30 click_args.append(arg) 

31 

32 result = super().parse_args(ctx, click_args) 

33 ctx.args.extend(hopwise_args) 

34 

35 return result 

36 

37 

38console = Console() 

39debug_message = """[dim] 

40 Use --debug for full traceback or --rich-traceback for enhanced formatting. Please, be careful 

41 that it should be placed after hopwise command and before any subcommand, e.g.: 

42 hopwise train --debug [model] [dataset] [--config-files config1 config2 ...] 

43 hopwise train --rich-traceback [model] [dataset] [--config-files config1 config2 ...] 

44 [/dim]""" 

45 

46 

47@click.group() 

48@click.version_option() 

49@click.option("--debug", is_flag=True, help="Enable debug mode with full tracebacks") 

50@click.option("--rich-traceback", is_flag=True, help="Use Rich's enhanced traceback formatting") 

51@click.pass_context 

52def cli(ctx, debug, rich_traceback): 

53 """ 

54 🔮 HopWise - Advanced Knowledge Graph-Enhanced Recommendation System 

55 

56 HopWise extends RecBole with knowledge graphs, path-based reasoning, 

57 and language modeling for explainable recommendations. 

58 """ 

59 ctx.ensure_object(dict) 

60 ctx.obj["debug"] = debug 

61 ctx.obj["rich_traceback"] = rich_traceback 

62 

63 # Only install Rich traceback if specifically requested 

64 if rich_traceback: 

65 install(console=console, show_locals=True) 

66 else: 

67 # Install minimal Rich traceback without the enhanced features 

68 install(console=console, show_locals=False, suppress=[click]) 

69 

70 

71@cli.command( 

72 cls=hopwiseClickCommand, 

73 context_settings=dict( 

74 allow_interspersed_args=True, 

75 ), 

76) 

77@click.option("--model", "-m", default="BPR", help="Model name to train") 

78@click.option("--dataset", "-d", default="ml-100k", help="Dataset name") 

79@click.option("--config-files", help="Space-separated config files") 

80@click.option("--checkpoint", help="Checkpoint (.pth) file path") 

81@click.option("--nproc", default=1, help="Number of processes") 

82@click.option("--ip", default="localhost", help="Master node IP") 

83@click.option("--port", default="5678", help="Master node port") 

84@click.option("--world-size", default=-1, help="Total number of jobs") 

85@click.option("--group-offset", default=0, help="Global rank offset") 

86@click.option("--proc-title", default=None, help="Processor Title, shown in top, nvidia utils, etc.") 

87@click.pass_context 

88def train(ctx, model, dataset, config_files, nproc, checkpoint, ip, port, world_size, group_offset, proc_title): 

89 """Train or evaluate a single model.""" 

90 

91 if proc_title is None: 

92 proc_title = f"[hopwise] {model} {dataset} training" 

93 setproctitle(proc_title) 

94 

95 config_file_list = config_files.strip().split(" ") if config_files else None 

96 

97 try: 

98 run( 

99 model, 

100 dataset, 

101 "train", 

102 checkpoint, 

103 config_file_list=config_file_list, 

104 nproc=nproc, 

105 world_size=world_size, 

106 ip=ip, 

107 port=port, 

108 group_offset=group_offset, 

109 ) 

110 except KeyboardInterrupt: 

111 console.print("\n[bold yellow]⚠️ Training interrupted by user[/bold yellow]") 

112 sys.exit(130) 

113 

114 except Exception as e: 

115 if ctx.obj.get("debug", False): 

116 if ctx.obj.get("rich_traceback", False): 

117 # Rich will handle this automatically with enhanced formatting 

118 raise 

119 else: 

120 console.print(f"[bold red]✗ Training failed:[/bold red] {type(e).__name__}: {str(e)}") 

121 console.print("\n[dim]Traceback:[/dim]") 

122 traceback.print_exc() 

123 sys.exit(1) 

124 else: 

125 console.print(f"[bold red]✗ Training failed:[/bold red] {type(e).__name__}: {str(e)}") 

126 console.print(debug_message) 

127 sys.exit(1) 

128 

129 

130@cli.command( 

131 cls=hopwiseClickCommand, 

132 context_settings=dict( 

133 allow_interspersed_args=True, 

134 ), 

135) 

136@click.option("--model", "-m", default="BPR", help="Model name to train") 

137@click.option("--dataset", "-d", default="ml-100k", help="Dataset name") 

138@click.option("--config-files", help="Space-separated config files") 

139@click.option("--checkpoint", help="Checkpoint (.pth) file path") 

140@click.option("--nproc", default=1, help="Number of processes") 

141@click.option("--ip", default="localhost", help="Master node IP") 

142@click.option("--port", default="5678", help="Master node port") 

143@click.option("--world-size", default=-1, help="Total number of jobs") 

144@click.option("--group-offset", default=0, help="Global rank offset") 

145@click.option("--proc-title", default=None, help="Processor Title, shown in top, nvidia utils, etc.") 

146@click.pass_context 

147def evaluate(ctx, model, dataset, config_files, nproc, checkpoint, ip, port, world_size, group_offset, proc_title): 

148 """Train or evaluate a single model.""" 

149 

150 if proc_title is None: 

151 proc_title = f"[hopwise] {model} {dataset} evaluation" 

152 setproctitle(proc_title) 

153 

154 config_file_list = config_files.strip().split(" ") if config_files else None 

155 

156 try: 

157 run( 

158 model, 

159 dataset, 

160 "evaluate", 

161 checkpoint, 

162 config_file_list=config_file_list, 

163 nproc=nproc, 

164 world_size=world_size, 

165 ip=ip, 

166 port=port, 

167 group_offset=group_offset, 

168 ) 

169 except KeyboardInterrupt: 

170 console.print("\n[bold yellow]⚠️ Evaluation interrupted by user[/bold yellow]") 

171 sys.exit(130) 

172 

173 except Exception as e: 

174 if ctx.obj.get("debug", False): 

175 if ctx.obj.get("rich_traceback", False): 

176 # Rich will handle this automatically with enhanced formatting 

177 raise 

178 else: 

179 console.print(f"[bold red]✗ Evaluation failed:[/bold red] {type(e).__name__}: {str(e)}") 

180 console.print("\n[dim]Traceback:[/dim]") 

181 traceback.print_exc() 

182 sys.exit(1) 

183 else: 

184 console.print(f"[bold red]✗ Evaluation failed:[/bold red] {type(e).__name__}: {str(e)}") 

185 console.print(debug_message) 

186 sys.exit(1) 

187 

188 

189@cli.command( 

190 cls=hopwiseClickCommand, 

191 context_settings=dict( 

192 allow_interspersed_args=True, 

193 ), 

194) 

195@click.option("--models", "-m", required=True, help="Comma-separated model names") 

196@click.option("--dataset", "-d", default="ml-100k", help="Dataset name") 

197@click.option("--config-files", help="Space-separated config files") 

198@click.option("--valid-latex", default="./latex/valid.tex", help="Valid results LaTeX file") 

199@click.option("--test-latex", default="./latex/test.tex", help="Test results LaTeX file") 

200@click.option("--nproc", default=1, help="Number of processes") 

201@click.option("--ip", default="localhost", help="Master node IP") 

202@click.option("--port", default="5678", help="Master node port") 

203@click.option("--world-size", default=-1, help="Total number of jobs") 

204@click.option("--group-offset", default=0, help="Global rank offset") 

205@click.option("--proc-title", default=None, help="Processor Title, shown in top, nvidia utils, etc.") 

206def benchmark( 

207 models, dataset, config_files, valid_latex, test_latex, nproc, ip, port, world_size, group_offset, proc_title 

208): 

209 """ 

210 Run scientific benchmark experiments across multiple models. 

211 

212 Trains multiple models on the same dataset and generates comparative results 

213 in LaTeX table format for scientific publications. Ideal for reproducing 

214 paper results or conducting systematic model comparisons. 

215 

216 Example: 

217 hopwise benchmark --models "BPR,LightGCN,KGAT" --dataset ml-100k --show-progress 

218 """ 

219 

220 if proc_title is None: 

221 proc_title = f"[hopwise - benchmark] {models} {dataset}" 

222 setproctitle(proc_title) 

223 

224 model_list = [m.strip() for m in models.split(",")] 

225 config_file_list = config_files.strip().split(" ") if config_files else None 

226 

227 os.makedirs(os.path.dirname(valid_latex), exist_ok=True) 

228 os.makedirs(os.path.dirname(test_latex), exist_ok=True) 

229 

230 console.print( 

231 Panel( 

232 f"[bold blue]Model Benchmark[/bold blue]\n" 

233 f"Models: [green]{', '.join(model_list)}[/green]\n" 

234 f"Dataset: [green]{dataset}[/green]\n" 

235 f"Total runs: [yellow]{len(model_list)}[/yellow]", 

236 box=box.ROUNDED, 

237 ) 

238 ) 

239 

240 valid_result_list = [] 

241 test_result_list = [] 

242 

243 for idx, model in enumerate(model_list): 

244 console.print(f"\n[bold blue]📊 Training {model} ({idx + 1}/{len(model_list)})[/bold blue]") 

245 

246 try: 

247 result = run( 

248 model, 

249 dataset, 

250 config_file_list=config_file_list, 

251 nproc=nproc, 

252 world_size=world_size, 

253 ip=ip, 

254 port=port, 

255 group_offset=group_offset, 

256 ) 

257 

258 valid_res_dict = {"Model": model} 

259 test_res_dict = {"Model": model} 

260 valid_res_dict.update(result["best_valid_result"]) 

261 test_res_dict.update(result["test_result"]) 

262 

263 valid_result_list.append(valid_res_dict) 

264 test_result_list.append(test_res_dict) 

265 

266 console.print(f"[green]✅ {model} completed successfully[/green]") 

267 

268 except Exception as e: 

269 console.print(f"[red]❌ {model} failed: {str(e)}[/red]") 

270 

271 successful = len(valid_result_list) 

272 failed = len(model_list) - successful 

273 console.print(f"\n[bold green]📋 Summary:[/bold green] {successful} successful, {failed} failed") 

274 

275 # Generate LaTeX tables 

276 try: 

277 bigger_flag = result["valid_score_bigger"] 

278 subset_columns = list(result["best_valid_result"].keys()) 

279 

280 df_valid, tex_valid = list_to_latex(valid_result_list, bigger_flag, subset_columns) 

281 df_test, tex_test = list_to_latex(test_result_list, bigger_flag, subset_columns) 

282 

283 with open(valid_latex, "w") as f: 

284 f.write(tex_valid) 

285 with open(test_latex, "w") as f: 

286 f.write(tex_test) 

287 

288 console.print(f"[bold green]✓[/bold green] Results saved to {valid_latex} and {test_latex}") 

289 

290 except Exception as e: 

291 console.print(f"[bold red]✗[/bold red] Failed to generate LaTeX: {str(e)}") 

292 

293 

294@cli.command( 

295 cls=hopwiseClickCommand, 

296 context_settings=dict( 

297 allow_interspersed_args=True, 

298 ), 

299) 

300@click.argument("params-file", type=click.Path(exists=True, dir_okay=False, readable=True)) 

301@click.option("--config-files", help="Fixed config files") 

302@click.option("--output-path", default="saved/hyper", help="Output directory") 

303@click.option("--display-file", help="Visualization file") 

304@click.option("--max-evals", default=10, help="Maximum evaluations") 

305@click.option("--tool", type=click.Choice(["hyperopt", "ray", "optuna"]), default="optuna", help="Tuning tool") 

306@click.option("--study-name", help="Study name for tuning") 

307@click.option("--algo", help="Algorithm for the tuner") 

308@click.option("--resume", is_flag=True, help="Resume from checkpoint") 

309@click.option("--proc-title", default=None, help="Processor Title, shown in top, nvidia utils, etc.") 

310@click.pass_context 

311def tune( 

312 ctx, params_file, config_files, output_path, display_file, max_evals, tool, study_name, algo, resume, proc_title 

313): 

314 """Run hyperparameter tuning.""" 

315 

316 if proc_title is None: 

317 proc_title = f"[hopwise - hyper] {study_name}" 

318 setproctitle(proc_title) 

319 

320 if not study_name: 

321 study_name = f"hyper_{datetime.now().strftime('%d_%m_%Y_%H_%M_%S')}" 

322 

323 config_file_list = config_files.strip().split(" ") if config_files else None 

324 

325 console.print( 

326 Panel( 

327 f"[bold blue]Hyperparameter Tuning[/bold blue]\n" 

328 f"Tool: [green]{tool}[/green]\n" 

329 f"Max Evaluations: [yellow]{max_evals}[/yellow]\n" 

330 f"Study: [cyan]{study_name}[/cyan]", 

331 box=box.ROUNDED, 

332 ) 

333 ) 

334 

335 try: 

336 ht = HyperTuning( 

337 objective_function, 

338 tuner=tool, 

339 algo=algo, 

340 early_stop=10, 

341 max_evals=max_evals, 

342 params_file=params_file, 

343 fixed_config_file_list=config_file_list, 

344 display_file=display_file, 

345 output_path=output_path, 

346 study_name=study_name, 

347 resume=resume, 

348 ) 

349 

350 console.print("[bold green]🚀[/bold green] Starting hyperparameter tuning...") 

351 ht.run() 

352 ht.export_result(output_path=output_path) 

353 

354 console.print( 

355 Panel( 

356 f"[bold green]✓ Tuning Completed![/bold green]\n" 

357 f"Best params: [cyan]{ht.best_params}[/cyan]\n" 

358 f"Best result: [yellow]{ht.params2result[ht.params2str(ht.best_params)]}[/yellow]", 

359 title="Results", 

360 box=box.ROUNDED, 

361 ) 

362 ) 

363 

364 except Exception as e: 

365 if ctx.obj.get("debug", False): 

366 if ctx.obj.get("rich_traceback", False): 

367 raise 

368 else: 

369 console.print(f"[bold red]✗ Tuning failed:[/bold red] {type(e).__name__}: {str(e)}") 

370 console.print("\n[dim]Traceback:[/dim]") 

371 traceback.print_exc() 

372 sys.exit(1) 

373 else: 

374 console.print(f"[bold red]✗[/bold red] Tuning failed: {str(e)}") 

375 console.print(debug_message) 

376 sys.exit(1) 

377 

378 

379@cli.command() 

380@click.option("--verbose", is_flag=True, help="Show detailed model list (docstrings)") 

381@click.option( 

382 "--type", 

383 "model_types", 

384 type=click.Choice( 

385 ["all", "Context", "Exlib", "General", "KG-aware", "KG-embed", "PathLM", "Sequential"], case_sensitive=False 

386 ), 

387 default=["all"], 

388 multiple=True, 

389 help="Filter by model type. Default is 'all' which shows all models.", 

390) 

391def models(verbose, model_types): 

392 """List available models.""" 

393 model_types_map = { 

394 "Context": "context_aware_recommender", 

395 "Exlib": "exlib_recommender", 

396 "General": "general_recommender", 

397 "KG-aware": "knowledge_aware_recommender", 

398 "KG-embed": "knowledge_graph_embedding_recommender", 

399 "PathLM": "path_language_modeling_recommender", 

400 "Sequential": "sequential_recommender", 

401 "all": "all", 

402 } 

403 model_types = [model_types_map[t] for t in model_types] 

404 

405 console.print("[bold blue]Available Models:[/bold blue]") 

406 models_dir = os.path.join(os.path.dirname(__file__), "model") 

407 models_info = [] 

408 for dir_content in os.scandir(models_dir): 

409 if dir_content.is_dir(): 

410 if "all" in model_types or dir_content.name.lower() in model_types: 

411 for filename in os.listdir(dir_content.path): 

412 if filename.endswith(".py") and not filename.startswith("_"): 

413 model_type = dir_content.name.lower() 

414 model_lcase = filename[:-3] 

415 model_module = importlib.import_module(f"hopwise.model.{model_type}.{model_lcase}") 

416 model_name = [name for name in dir(model_module) if name.lower() == model_lcase][0] 

417 model_module = getattr(model_module, model_name) 

418 models_info.append( 

419 { 

420 "name": model_name, 

421 "type": model_type, 

422 "doc": getattr(model_module, "__doc__", "No documentation available"), 

423 } 

424 ) 

425 for info in models_info: 

426 console.print(f"- [blue]{info['type']}[/blue] [green]{info['name']}[/green]") 

427 if verbose: 

428 console.print(f" [yellow]{info['doc']}[/yellow]") 

429 

430 

431if __name__ == "__main__": 

432 cli()