Coverage for hopwise/cli.py: 0%
219 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
1import importlib
2import os
3import sys
4import traceback
5from datetime import datetime
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
14from hopwise.quick_start import objective_function, run
15from hopwise.trainer import HyperTuning
16from hopwise.utils import list_to_latex
19class hopwiseClickCommand(click.Command):
20 def parse_args(self, ctx, args):
21 """Override to filter out HopWise parameters before Click's validation"""
23 click_args = []
24 hopwise_args = []
26 for arg in args:
27 if arg.startswith("--") and "=" in arg:
28 hopwise_args.append(arg)
29 else:
30 click_args.append(arg)
32 result = super().parse_args(ctx, click_args)
33 ctx.args.extend(hopwise_args)
35 return result
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]"""
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
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
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])
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."""
91 if proc_title is None:
92 proc_title = f"[hopwise] {model} {dataset} training"
93 setproctitle(proc_title)
95 config_file_list = config_files.strip().split(" ") if config_files else None
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)
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)
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."""
150 if proc_title is None:
151 proc_title = f"[hopwise] {model} {dataset} evaluation"
152 setproctitle(proc_title)
154 config_file_list = config_files.strip().split(" ") if config_files else None
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)
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)
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.
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.
216 Example:
217 hopwise benchmark --models "BPR,LightGCN,KGAT" --dataset ml-100k --show-progress
218 """
220 if proc_title is None:
221 proc_title = f"[hopwise - benchmark] {models} {dataset}"
222 setproctitle(proc_title)
224 model_list = [m.strip() for m in models.split(",")]
225 config_file_list = config_files.strip().split(" ") if config_files else None
227 os.makedirs(os.path.dirname(valid_latex), exist_ok=True)
228 os.makedirs(os.path.dirname(test_latex), exist_ok=True)
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 )
240 valid_result_list = []
241 test_result_list = []
243 for idx, model in enumerate(model_list):
244 console.print(f"\n[bold blue]📊 Training {model} ({idx + 1}/{len(model_list)})[/bold blue]")
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 )
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"])
263 valid_result_list.append(valid_res_dict)
264 test_result_list.append(test_res_dict)
266 console.print(f"[green]✅ {model} completed successfully[/green]")
268 except Exception as e:
269 console.print(f"[red]❌ {model} failed: {str(e)}[/red]")
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")
275 # Generate LaTeX tables
276 try:
277 bigger_flag = result["valid_score_bigger"]
278 subset_columns = list(result["best_valid_result"].keys())
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)
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)
288 console.print(f"[bold green]✓[/bold green] Results saved to {valid_latex} and {test_latex}")
290 except Exception as e:
291 console.print(f"[bold red]✗[/bold red] Failed to generate LaTeX: {str(e)}")
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."""
316 if proc_title is None:
317 proc_title = f"[hopwise - hyper] {study_name}"
318 setproctitle(proc_title)
320 if not study_name:
321 study_name = f"hyper_{datetime.now().strftime('%d_%m_%Y_%H_%M_%S')}"
323 config_file_list = config_files.strip().split(" ") if config_files else None
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 )
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 )
350 console.print("[bold green]🚀[/bold green] Starting hyperparameter tuning...")
351 ht.run()
352 ht.export_result(output_path=output_path)
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 )
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)
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]
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]")
431if __name__ == "__main__":
432 cli()