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
6 changes: 5 additions & 1 deletion scripts/completions/unilab.bash
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,14 @@ _unilab_uv_complete() {
return 0
fi

local script_dir repo_root
script_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
repo_root="$(cd -- "$script_dir/../.." && pwd -P)"

local candidates
if ! mapfile -t candidates < <(
uv run --no-sync unilab-complete --cword "$COMP_CWORD" -- "${COMP_WORDS[@]}" 2>/dev/null \
|| uv run --no-sync python -m unilab.tools.completion --cword "$COMP_CWORD" -- "${COMP_WORDS[@]}" 2>/dev/null
|| PYTHONPATH="$repo_root/src${PYTHONPATH:+:$PYTHONPATH}" uv run --no-sync python -m unilab.tools.completion --cword "$COMP_CWORD" -- "${COMP_WORDS[@]}" 2>/dev/null
); then
return 0
fi
Expand Down
6 changes: 5 additions & 1 deletion scripts/completions/unilab.zsh
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,14 @@ _unilab_uv_complete() {
return 0
fi

local script_path repo_root
script_path="${${(%):-%x}:A}"
repo_root="${script_path:h:h:h}"

local output
output="$(
uv run --no-sync unilab-complete --cword "$((CURRENT - 1))" -- "${words[@]}" 2>/dev/null \
|| uv run --no-sync python -m unilab.tools.completion --cword "$((CURRENT - 1))" -- "${words[@]}" 2>/dev/null
|| PYTHONPATH="$repo_root/src${PYTHONPATH:+:$PYTHONPATH}" uv run --no-sync python -m unilab.tools.completion --cword "$((CURRENT - 1))" -- "${words[@]}" 2>/dev/null
)" || return 0

if [[ -z "$output" ]]; then
Expand Down
6 changes: 2 additions & 4 deletions src/unilab/__init__.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,7 @@
"""UniLab package initialization."""

import os

from etils import epath
from pathlib import Path

__version__ = "0.0.0"

ROOT_PATH = epath.Path(__file__).parent
ROOT_PATH = Path(__file__).resolve().parent
168 changes: 168 additions & 0 deletions src/unilab/tools/completion.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,14 @@
RUN_PATH_SUFFIXES = (".py", ".sh")
RUN_PATH_IGNORED_PARTS = ("__pycache__", "outputs")
SCRIPT_ASSIGNMENT_PATTERN = re.compile(r'^([A-Za-z0-9_.-]+)\s*=\s*"([^"]+)"\s*(?:#.*)?$')
DEFAULT_ALGO_LOG_NAMES = {
"ppo": "rsl_rl_ppo",
"mlx_ppo": "mlx_rl_train",
"appo": "appo",
"sac": "fast_sac",
"td3": "fast_td3",
"flashsac": "flash_sac",
}
COMPLETION_BLOCK_START = "# >>> unilab completion >>>"
COMPLETION_BLOCK_END = "# <<< unilab completion <<<"
SUPPORTED_SHELLS = ("bash", "zsh")
Expand All @@ -41,6 +49,7 @@ class CompletionMetadata:
choices: dict[str, dict[str, tuple[str, ...]]]
tasks: tuple[TaskCompletionEntry, ...]
run_paths: tuple[str, ...] = ()
root: Path | None = None


def _find_project_root(start: Path) -> Path | None:
Expand Down Expand Up @@ -216,6 +225,7 @@ def build_metadata(root: Path | None = None) -> CompletionMetadata:
choices={command: _parser_choices(command) for command in training_commands},
tasks=_task_entries(selected_root),
run_paths=_run_path_entries(selected_root),
root=selected_root,
)


Expand Down Expand Up @@ -306,6 +316,147 @@ def _task_choices(
return _matching(tuple(sorted(tasks)), prefix)


def _profile_choices(
metadata: CompletionMetadata,
*,
prefix: str,
selected_algo: str | None,
selected_sim: str | None,
selected_task: str | None,
) -> list[str]:
profiles: set[str] = set()
for entry in metadata.tasks:
if selected_algo is not None and entry.algo != selected_algo:
continue
if selected_sim is not None and entry.sim != selected_sim:
continue
if selected_task is not None and entry.task != selected_task:
continue
owner_prefix = f"{entry.sim}_"
if not entry.owner.startswith(owner_prefix):
continue
profile = entry.owner[len(owner_prefix) :]
if profile:
profiles.add(profile)
return _matching(tuple(sorted(profiles)), prefix)


def _strip_yaml_scalar(value: str) -> str:
selected = value.split("#", 1)[0].strip()
if len(selected) >= 2 and selected[0] == selected[-1] and selected[0] in {'"', "'"}:
return selected[1:-1]
return selected


def _yaml_section_scalar(path: Path, section: str, key: str) -> str | None:
if not path.is_file():
return None
in_section = False
section_indent = 0
for raw_line in path.read_text(encoding="utf-8").splitlines():
stripped = raw_line.strip()
if not stripped or stripped.startswith("#"):
continue
indent = len(raw_line) - len(raw_line.lstrip())
if indent == 0 and stripped.endswith(":"):
in_section = stripped[:-1] == section
section_indent = indent
continue
if not in_section:
continue
if indent <= section_indent:
break
if stripped.startswith(f"{key}:"):
return _strip_yaml_scalar(stripped.split(":", 1)[1])
return None


def _owner_yaml_paths(
root: Path,
*,
selected_algo: str,
selected_task: str,
selected_sim: str,
selected_profile: str | None,
) -> tuple[Path, ...]:
try:
route = cli.build_route(selected_algo, selected_task, selected_sim, selected_profile)
except SystemExit:
return ()

route_path = root / "conf" / route.config_group / "task" / route.owner_task
if not route_path.is_file():
return ()

paths = [route_path]
if selected_profile is not None:
try:
base_route = cli.build_route(selected_algo, selected_task, selected_sim, None)
except SystemExit:
base_route = None
if base_route is not None:
base_path = root / "conf" / base_route.config_group / "task" / base_route.owner_task
if base_path not in paths:
paths.append(base_path)
return tuple(paths)


def _first_yaml_section_scalar(paths: Sequence[Path], section: str, key: str) -> str | None:
for path in paths:
selected = _yaml_section_scalar(path, section, key)
if selected not in (None, ""):
return selected
return None


def _load_run_choices(
metadata: CompletionMetadata,
*,
prefix: str,
selected_algo: str | None,
selected_task: str | None,
selected_sim: str | None,
selected_profile: str | None,
) -> list[str]:
candidates = ["-1"]
if metadata.root is None or not selected_algo or not selected_task or not selected_sim:
return _matching(candidates, prefix)

owner_paths = _owner_yaml_paths(
metadata.root,
selected_algo=selected_algo,
selected_task=selected_task,
selected_sim=selected_sim,
selected_profile=selected_profile,
)
if not owner_paths:
return _matching(candidates, prefix)

task_name = _first_yaml_section_scalar(owner_paths, "training", "task_name") or selected_task
log_root = _first_yaml_section_scalar(owner_paths, "training", "log_root")
if log_root is not None:
log_root_path = Path(log_root)
base_log_root = (
log_root_path if log_root_path.is_absolute() else metadata.root / log_root_path
)
else:
algo_log_name = _first_yaml_section_scalar(
owner_paths, "algo", "algo_log_name"
) or DEFAULT_ALGO_LOG_NAMES.get(selected_algo)
if algo_log_name is None:
return _matching(candidates, prefix)
base_log_root = metadata.root / "logs" / algo_log_name

task_log_root = base_log_root / task_name
if task_log_root.is_dir():
candidates.extend(
path.name
for path in sorted(task_log_root.iterdir())
if path.is_dir() and cli.RUN_ID_PATTERN.fullmatch(path.name) is not None
)
return _dedupe(_matching(candidates, prefix))


def complete_words(
words: Sequence[str],
cword: int,
Expand Down Expand Up @@ -341,6 +492,23 @@ def complete_words(
)
if previous in choices:
return _matching(choices[previous], current)
if command == "eval" and previous == "--load-run":
return _load_run_choices(
selected_metadata,
prefix=current,
selected_algo=option_values.get("--algo"),
selected_task=option_values.get("--task"),
selected_sim=option_values.get("--sim"),
selected_profile=option_values.get("--profile"),
)
if command in {"train", "eval"} and previous == "--profile":
return _profile_choices(
selected_metadata,
prefix=current,
selected_algo=option_values.get("--algo"),
selected_sim=option_values.get("--sim"),
selected_task=option_values.get("--task"),
)
if current == "" or current.startswith("-") or previous == command:
return _matching(_available_flags(selected_metadata, command, used_options), current)
return []
Expand Down
Loading
Loading