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
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
32 changes: 20 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 _calculate_grid_size(
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,23 @@ def run_optimization(
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"
"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:
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"
)
logger.info(
f"Starting optimization: {total_trials} trials, "
Expand All @@ -508,7 +511,7 @@ def run_optimization(
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 +613,14 @@ def run_optimization(
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
):
"""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()

# 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