diff --git a/benchmark/benchmark_ane_inference.py b/benchmark/benchmark_ane_inference.py index cdf9e9d50..f48623ffe 100644 --- a/benchmark/benchmark_ane_inference.py +++ b/benchmark/benchmark_ane_inference.py @@ -203,6 +203,7 @@ def save_plots( ax1.set_xlabel("env_num") ax1.set_ylabel("Mean time (ms)") ax1.set_xscale("log", base=2) + ax1.set_yscale("log") ax1.set_xticks(env_nums) ax1.set_xticklabels([str(v) for v in env_nums]) ax1.grid(True, alpha=0.3) @@ -229,6 +230,7 @@ def save_plots( ax2.set_xlabel("env_num") ax2.set_ylabel("Envs per second") ax2.set_xscale("log", base=2) + ax2.set_yscale("log") ax2.set_xticks(env_nums) ax2.set_xticklabels([str(v) for v in env_nums]) ax2.grid(True, alpha=0.3) diff --git a/benchmark/benchmark_backends.py b/benchmark/benchmark_backends.py index 591cee712..b4fa3496c 100644 --- a/benchmark/benchmark_backends.py +++ b/benchmark/benchmark_backends.py @@ -188,7 +188,7 @@ def ew(): def main(): parser = argparse.ArgumentParser(description="Benchmark NumPy/Torch/MLX compute.") - parser.add_argument("--sizes", type=str, default=",".join(str(2**k) for k in range(5, 15))) + parser.add_argument("--sizes", type=str, default=",".join(str(2**k) for k in range(5, 12))) parser.add_argument("--warmup", type=int, default=1) parser.add_argument("--repeat", type=int, default=2) parser.add_argument("--dtypes", type=str, default="float16,float32") diff --git a/benchmark/benchmark_mlp_inference.py b/benchmark/benchmark_mlp_inference.py index 421c52452..4e760675e 100644 --- a/benchmark/benchmark_mlp_inference.py +++ b/benchmark/benchmark_mlp_inference.py @@ -1017,6 +1017,7 @@ def save_plots( ax1.set_xlabel("env_num") ax1.set_ylabel("Mean time (ms)") ax1.set_xscale("log", base=2) + ax1.set_yscale("log") ax1.set_xticks(env_nums) ax1.set_xticklabels([str(v) for v in env_nums]) ax1.grid(True, alpha=0.3) @@ -1043,6 +1044,7 @@ def save_plots( ax2.set_xlabel("env_num") ax2.set_ylabel("Envs per second") ax2.set_xscale("log", base=2) + ax2.set_yscale("log") ax2.set_xticks(env_nums) ax2.set_xticklabels([str(v) for v in env_nums]) ax2.grid(True, alpha=0.3) diff --git a/scripts/completions/unilab.bash b/scripts/completions/unilab.bash index cc4895634..b324e051d 100644 --- a/scripts/completions/unilab.bash +++ b/scripts/completions/unilab.bash @@ -15,7 +15,15 @@ _unilab_uv_complete() { return 0 fi + if [[ ${#candidates[@]} -eq 0 ]]; then + if [[ $COMP_CWORD -le 2 || ${COMP_WORDS[2]} != "train" && ${COMP_WORDS[2]} != "eval" ]]; then + compopt -o default -o bashdefault 2>/dev/null || true + return 1 + fi + return 0 + fi + COMPREPLY=("${candidates[@]}") } -complete -F _unilab_uv_complete uv +complete -o default -o bashdefault -F _unilab_uv_complete uv diff --git a/scripts/completions/unilab.zsh b/scripts/completions/unilab.zsh index 3552e4986..3e0f01977 100644 --- a/scripts/completions/unilab.zsh +++ b/scripts/completions/unilab.zsh @@ -12,7 +12,12 @@ _unilab_uv_complete() { || uv run --no-sync python -m unilab.tools.completion --cword "$((CURRENT - 1))" -- "${words[@]}" 2>/dev/null )" || return 0 - [[ -n "$output" ]] || return 0 + if [[ -z "$output" ]]; then + if (( CURRENT <= 3 )) || [[ ${words[3]} != "train" && ${words[3]} != "eval" ]]; then + _files + fi + return 0 + fi local -a candidates candidates=("${(@f)output}") diff --git a/src/unilab/tools/completion.py b/src/unilab/tools/completion.py index 5ea0e313f..479bdbae1 100644 --- a/src/unilab/tools/completion.py +++ b/src/unilab/tools/completion.py @@ -17,6 +17,9 @@ "train": "unilab.cli:train_main", "eval": "unilab.cli:eval_main", } +RUN_PATH_ROOTS = ("benchmark", "scripts") +RUN_PATH_SUFFIXES = (".py", ".sh") +RUN_PATH_IGNORED_PARTS = ("__pycache__", "outputs") SCRIPT_ASSIGNMENT_PATTERN = re.compile(r'^([A-Za-z0-9_.-]+)\s*=\s*"([^"]+)"\s*(?:#.*)?$') COMPLETION_BLOCK_START = "# >>> unilab completion >>>" COMPLETION_BLOCK_END = "# <<< unilab completion <<<" @@ -37,6 +40,7 @@ class CompletionMetadata: flags: dict[str, tuple[str, ...]] choices: dict[str, dict[str, tuple[str, ...]]] tasks: tuple[TaskCompletionEntry, ...] + run_paths: tuple[str, ...] = () def _find_project_root(start: Path) -> Path | None: @@ -75,6 +79,13 @@ def _training_commands(scripts: Mapping[str, str]) -> tuple[str, ...]: ) +def _project_commands(scripts: Mapping[str, str]) -> tuple[str, ...]: + training_commands = _training_commands(scripts) + commands = [*training_commands] + commands.extend(command for command in sorted(scripts) if command not in training_commands) + return tuple(commands) + + def _parser_for_command(command: str) -> argparse.ArgumentParser: return cli._train_eval_parser(mode=command) @@ -162,15 +173,49 @@ def _task_entries(root: Path) -> tuple[TaskCompletionEntry, ...]: ) +def _is_run_path(path: Path) -> bool: + return path.is_file() and path.suffix in RUN_PATH_SUFFIXES and path.name != "__init__.py" + + +def _has_run_path_descendant(path: Path) -> bool: + return any( + child.name != "__init__.py" and child.suffix in RUN_PATH_SUFFIXES + for child in path.rglob("*") + ) + + +def _run_path_entries(root: Path) -> tuple[str, ...]: + entries: set[str] = set() + for path_root in RUN_PATH_ROOTS: + base = root / path_root + if not base.is_dir(): + continue + entries.add(f"{path_root}/") + for path in base.rglob("*"): + if path.name.startswith("."): + continue + relative_parts = path.relative_to(root).parts + if any(part in RUN_PATH_IGNORED_PARTS for part in relative_parts): + continue + relative_path = path.relative_to(root).as_posix() + if path.is_dir(): + if _has_run_path_descendant(path): + entries.add(f"{relative_path}/") + elif _is_run_path(path): + entries.add(relative_path) + return tuple(sorted(entries)) + + def build_metadata(root: Path | None = None) -> CompletionMetadata: selected_root = root or _find_project_root(Path.cwd()) or cli.repo_root() scripts = _read_project_scripts(selected_root / "pyproject.toml") - commands = _training_commands(scripts) + training_commands = _training_commands(scripts) return CompletionMetadata( - commands=commands, - flags={command: _parser_flags(command) for command in commands}, - choices={command: _parser_choices(command) for command in commands}, + commands=_project_commands(scripts), + flags={command: _parser_flags(command) for command in training_commands}, + choices={command: _parser_choices(command) for command in training_commands}, tasks=_task_entries(selected_root), + run_paths=_run_path_entries(selected_root), ) @@ -190,6 +235,25 @@ def _matching(candidates: Sequence[str], prefix: str) -> list[str]: return [candidate for candidate in candidates if candidate.startswith(prefix)] +def _dedupe(candidates: Sequence[str]) -> list[str]: + return list(dict.fromkeys(candidates)) + + +def _path_choices(candidates: Sequence[str], prefix: str) -> list[str]: + choices: set[str] = set() + for candidate in candidates: + if not candidate.startswith(prefix): + continue + remainder = candidate[len(prefix) :] + if remainder == "": + continue + if "/" in remainder: + choices.add(f"{prefix}{remainder.split('/', 1)[0]}/") + else: + choices.add(candidate) + return sorted(choices) + + def _option_state(words: Sequence[str], cword: int) -> tuple[dict[str, str], set[str]]: values: dict[str, str] = {} used: set[str] = set() @@ -253,7 +317,12 @@ def complete_words( current = _current_word(words, cword) if cword <= 2: - return _matching(selected_metadata.commands, current) + return _dedupe( + [ + *_matching(selected_metadata.commands, current), + *_path_choices(selected_metadata.run_paths, current), + ] + ) command = words[2] if command not in selected_metadata.commands: diff --git a/tests/test_completion.py b/tests/test_completion.py new file mode 100644 index 000000000..8ecc1bb88 --- /dev/null +++ b/tests/test_completion.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from pathlib import Path + +from unilab.tools.completion import build_metadata, complete_words + + +def _write_completion_fixture(root: Path) -> None: + (root / "pyproject.toml").write_text( + "\n".join( + [ + "[project.scripts]", + 'train = "unilab.cli:train_main"', + 'eval = "unilab.cli:eval_main"', + 'demo = "unilab.cli:demo_main"', + ] + ), + encoding="utf-8", + ) + (root / "benchmark" / "core").mkdir(parents=True) + (root / "benchmark" / "benchmark_sim.py").write_text("", encoding="utf-8") + (root / "benchmark" / "core" / "runner.py").write_text("", encoding="utf-8") + (root / "scripts").mkdir() + (root / "scripts" / "play_viser.py").write_text("", encoding="utf-8") + + +def test_uv_run_command_position_includes_project_scripts_and_run_paths(tmp_path: Path) -> None: + _write_completion_fixture(tmp_path) + metadata = build_metadata(tmp_path) + + assert "demo" in complete_words(["uv", "run", "d"], 2, metadata) + assert complete_words(["uv", "run", "b"], 2, metadata) == ["benchmark/"] + + +def test_uv_run_path_completion_is_hierarchical(tmp_path: Path) -> None: + _write_completion_fixture(tmp_path) + metadata = build_metadata(tmp_path) + + benchmark_choices = complete_words(["uv", "run", "benchmark/"], 2, metadata) + assert "benchmark/benchmark_sim.py" in benchmark_choices + assert "benchmark/core/" in benchmark_choices + assert "benchmark/core/runner.py" not in benchmark_choices + + assert complete_words(["uv", "run", "benchmark/core/"], 2, metadata) == [ + "benchmark/core/runner.py" + ] + + +def test_uv_run_unknown_command_arguments_defer_to_shell_completion(tmp_path: Path) -> None: + _write_completion_fixture(tmp_path) + metadata = build_metadata(tmp_path) + + assert complete_words(["uv", "run", "benchmark/benchmark_sim.py", ""], 3, metadata) == []