Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
72 changes: 42 additions & 30 deletions auto_tune_vllm/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,9 +121,9 @@
n_trials: Optional[int] = typer.Option(
None, "--trials", "-n", help="Number of trials (overrides config)"
),
max_concurrent: Optional[int] = typer.Option(
max_concurrent_trials: Optional[int] = typer.Option(
None,
"--max-concurrent",
"--max-concurrent-trials",
help="REQUIRED: Max concurrent trials to run simultaneously.",
),
verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose logging"),
Expand Down Expand Up @@ -241,26 +241,33 @@

# Run optimization
# Use CLI value or fall back to YAML config
final_max_concurrent = (
max_concurrent or study_config.optimization.max_concurrent
final_max_concurrent_trials = (
max_concurrent_trials or study_config.optimization.max_concurrent_trials
)
if final_max_concurrent is None:
if final_max_concurrent_trials is None:
console.print(
"[bold red]"
"❌ --max-concurrent is required "
"(or set optimization.max_concurrent in the config)"
"❌ --max-concurrent-trials is required "
"(or set optimization.max_concurrent_trials in the config)"
"[/bold red]"
)
console.print(
"YAML:\n optimization:\n max_concurrent: 2 # Match your GPU count"
"YAML:\n optimization:\n"
" max_concurrent_trials: 2"
)
raise typer.Exit(1)
if final_max_concurrent < 1:
console.print("[bold red]❌ --max-concurrent must be >= 1[/bold red]")
if final_max_concurrent_trials < 1:
console.print(
"[bold red]❌ --max-concurrent-trials must be >= 1[/bold red]"
)
raise typer.Exit(1)

run_optimization_sync(
execution_backend, study_config, n_trials, final_max_concurrent, create_db
execution_backend,
study_config,
n_trials,
final_max_concurrent_trials,
create_db,
)

except Exception as e:
Expand Down Expand Up @@ -358,7 +365,7 @@
backend,
config: StudyConfig,
n_trials: Optional[int],
max_concurrent: Optional[int],
max_concurrent_trials: Optional[int],
create_db: bool = False,
):
"""Synchronous optimization runner with progress display."""
Expand Down Expand Up @@ -388,7 +395,7 @@

try:
# Run optimization
controller.run_optimization(n_trials, max_concurrent)
controller.run_optimization(n_trials, max_concurrent_trials)

# Mark task as completed
progress.update(
Expand Down Expand Up @@ -436,7 +443,7 @@

if results.get("baseline_improvement") is not None:
improvement = results["baseline_improvement"]
if improvement > 0:

Check failure on line 446 in auto_tune_vllm/cli/main.py

View workflow job for this annotation

GitHub Actions / pyright

Operator ">" not supported for types "int | float | str | list[float] | list[dict[str, int | float | str | list[float] | None]] | dict[str, str | int | float] | None" and "Literal[0]"   Operator ">" not supported for types "str" and "Literal[0]"   Operator ">" not supported for types "list[float]" and "Literal[0]"   Operator ">" not supported for types "list[dict[str, int | float | str | list[float] | None]]" and "Literal[0]"   Operator ">" not supported for types "dict[str, str | int | float]" and "Literal[0]"   Operator ">" not supported for types "None" and "Literal[0]" (reportOperatorIssue)
improvement_text = f"+{improvement:.1f}%"
improvement_style = "green"
else:
Expand Down Expand Up @@ -818,13 +825,10 @@
"--total-trials",
help="Total number of trials to reach (overrides config)",
),
max_concurrent: Optional[int] = typer.Option(
max_concurrent_trials: Optional[int] = typer.Option(
None,
"--max-concurrent",
help=(
"REQUIRED: Max concurrent trials to run simultaneously"
"(should match your GPU count)"
),
"--max-concurrent-trials",
help="REQUIRED: Max concurrent trials to run simultaneously.",
),
verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose logging"),
start_ray_head: bool = typer.Option(
Expand Down Expand Up @@ -916,27 +920,35 @@

# Resume study
# Use CLI value or fall back to YAML config
final_max_concurrent = (
max_concurrent or study_config.optimization.max_concurrent
final_max_concurrent_trials = (
max_concurrent_trials or study_config.optimization.max_concurrent_trials
)
if n_trials or n_total_trials: # Only required if we will run new trials
if final_max_concurrent is None:
if final_max_concurrent_trials is None:
console.print(
"[bold red]"
"❌ --max-concurrent is required to resume with new trials"
"❌ --max-concurrent-trials is required to resume "
"with new trials"
"[/bold red]"
)
console.print("Set CLI flag or config: optimization.max_concurrent")
console.print(
"Set CLI flag or config: "
"optimization.max_concurrent_trials"
)
raise typer.Exit(1)
if final_max_concurrent < 1:
console.print("[bold red]❌ --max-concurrent must be >= 1[/bold red]")
if final_max_concurrent_trials < 1:
console.print(
"[bold red]"
"❌ --max-concurrent-trials must be >= 1"
"[/bold red]"
)
raise typer.Exit(1)
resume_study_sync(
execution_backend,
study_config,
n_trials,
n_total_trials,
final_max_concurrent,
final_max_concurrent_trials,
)

except Exception as e:
Expand All @@ -951,7 +963,7 @@
config: StudyConfig,
n_trials: Optional[int],
n_total_trials: Optional[int],
max_concurrent: Optional[int],
max_concurrent_trials: Optional[int],
):
"""Resume study execution."""
controller = StudyController.resume_from_config(backend, config)
Expand All @@ -978,7 +990,7 @@
if n_trials is not None:
# --trials specifies additional trials to run
console.print(f"Running {n_trials} additional trials...")
controller.run_optimization(n_trials, max_concurrent)
controller.run_optimization(n_trials, max_concurrent_trials)

# Display updated results after running additional trials
console.print(
Expand All @@ -1004,7 +1016,7 @@
f"Running {trials_to_run} more trials to reach total of "
f"{n_total_trials} trials..."
)
controller.run_optimization(trials_to_run, max_concurrent)
controller.run_optimization(trials_to_run, max_concurrent_trials)

# Display updated results after running additional trials
console.print(
Expand Down
4 changes: 2 additions & 2 deletions auto_tune_vllm/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,8 +69,8 @@ class OptimizationConfig:
objective: Union[str, List[str]] = None # Old format: "maximize", "minimize", list
sampler: str = "tpe" # "tpe", "random", "gp", "botorch", "nsga2", "grid"
n_trials: int = 100
n_startup_trials: int = 10 # Number of random startup trials
max_concurrent: Optional[int] = (
n_startup_trials: int = 10 # Number of random startup trials
max_concurrent_trials: Optional[int] = (
None # Maximum concurrent trials (required for resource management)
)

Expand Down
24 changes: 12 additions & 12 deletions auto_tune_vllm/core/study_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -460,15 +460,15 @@
def run_optimization(
self,
n_trials: int | None = None,
max_concurrent: int | None = None,
max_concurrent_trials: int | None = None,
poll_interval: float = 5.0,
) -> optuna.Study:
"""
Run optimization study.

Args:
n_trials: Number of trials to run (overrides config)
max_concurrent: Maximum concurrent trials (None = unlimited)
max_concurrent_trials: Maximum concurrent trials (None = unlimited)
poll_interval: How often to poll for completed trials (seconds)

Returns:
Expand All @@ -477,20 +477,20 @@
total_trials = n_trials or self.config.optimization.n_trials

# Require explicit, positive concurrency specification
if max_concurrent is None:
if max_concurrent_trials is None:
msg = (
"❌ --max-concurrent is required to prevent GPU memory conflicts!\n\n"
"❌ --max-concurrent-trials is required to prevent GPU memory conflicts!\n\n"

Check failure on line 482 in auto_tune_vllm/core/study_controller.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (E501)

auto_tune_vllm/core/study_controller.py:482:88: E501 Line too long (93 > 88)
"Add to your YAML config:\n"
" optimization:\n"
" max_concurrent: 2 # Match your GPU count\n\n"
"Or use CLI: --max-concurrent 2"
" max_concurrent_trials: 2 # Match your GPU count\n\n"
"Or use CLI: --max-concurrent-trials 2"
)
raise ValueError(msg)
if max_concurrent < 1:
raise ValueError("--max-concurrent must be >= 1")
if max_concurrent_trials < 1:
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
raise ValueError("--max-concurrent-trials must be >= 1")

max_concurrent_str = (
max_concurrent if max_concurrent != float("inf") else "unlimited"
max_concurrent_trials if max_concurrent_trials != float("inf") else "unlimited"

Check failure on line 493 in auto_tune_vllm/core/study_controller.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (E501)

auto_tune_vllm/core/study_controller.py:493:89: E501 Line too long (91 > 88)
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
logger.info(
f"Starting optimization: {total_trials} trials, "
Expand All @@ -508,7 +508,7 @@
remaining_trials=total_trials
- self.completed_trials
- len(self.active_trials),
max_concurrent=max_concurrent,
max_concurrent_trials=max_concurrent_trials,
)

# Poll for completed trials
Expand Down Expand Up @@ -610,9 +610,9 @@
self.backend.shutdown()
logger.info("Backend shutdown complete.")

def _submit_available_trials(self, remaining_trials: int, max_concurrent: float):
def _submit_available_trials(self, remaining_trials: int, max_concurrent_trials: float):

Check failure on line 613 in auto_tune_vllm/core/study_controller.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (E501)

auto_tune_vllm/core/study_controller.py:613:89: E501 Line too long (92 > 88)
"""Submit new trials up to limits."""
while remaining_trials > 0 and len(self.active_trials) < max_concurrent:
while remaining_trials > 0 and len(self.active_trials) < max_concurrent_trials:
trial = self.study.ask()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated

# Check if these exact parameters have already been tried and failed
Expand Down
2 changes: 1 addition & 1 deletion docs/quick_start.md
Original file line number Diff line number Diff line change
Expand Up @@ -223,7 +223,7 @@ You can now explore optimization history, parameter importance, and parallel coo
ray stop --force
ray start --head --dashboard-host=0.0.0.0
```
- **Important**: Add `--max-concurrent <count>` or set `max_concurrent: <count>` in your YAML config.
- **Important**: Add `--max-concurrent-trials <count>` or set `max_concurrent_trials: <count>` in your YAML config.

### Next Steps

Expand Down
Loading