Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions benchmark/benchmark_ane_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion benchmark/benchmark_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
2 changes: 2 additions & 0 deletions benchmark/benchmark_mlp_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand Down
10 changes: 9 additions & 1 deletion scripts/completions/unilab.bash
Original file line number Diff line number Diff line change
Expand Up @@ -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
7 changes: 6 additions & 1 deletion scripts/completions/unilab.zsh
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down
79 changes: 74 additions & 5 deletions src/unilab/tools/completion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 <<<"
Expand All @@ -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:
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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),
)


Expand All @@ -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()
Expand Down Expand Up @@ -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:
Expand Down
53 changes: 53 additions & 0 deletions tests/test_completion.py
Original file line number Diff line number Diff line change
@@ -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) == []
Loading