Skip to content
Merged
Show file tree
Hide file tree
Changes from 7 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
4 changes: 1 addition & 3 deletions auto_tune_vllm/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,9 +69,7 @@ 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: Optional[int] = (
None # Number of startup trials for samplers that support it
)
n_startup_trials: int = 10 # Number of random startup trials
max_concurrent: Optional[int] = (
None # Maximum concurrent trials (required for resource management)
)
Expand Down
34 changes: 30 additions & 4 deletions auto_tune_vllm/core/study_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -395,19 +395,45 @@ def _create_search_space(
def _create_sampler(config: StudyConfig) -> optuna.samplers.BaseSampler:
"""Create Optuna sampler from configuration."""
sampler_name = config.optimization.sampler.lower()
n_startup_trials = config.optimization.n_startup_trials
n_trials = config.optimization.n_trials

# Validate n_startup_trials < n_trials to ensure the sampler algorithm runs
# This applies to all samplers with startup trials (TPE, GP, BoTorch)
if n_startup_trials >= n_trials:
suggestion = max(1, n_trials // 10)
min_trials = n_startup_trials + 1
msg = (
f"n_startup_trials ({n_startup_trials}) must be less than "
f"n_trials ({n_trials}). Otherwise all trials would be random. "
f"Suggestion: Set n_startup_trials to {suggestion} "
f"or increase n_trials to at least {min_trials}."
)
raise ValueError(msg)

# Log sampler configuration
logger.info(
f"Creating {sampler_name.upper()} sampler "
f"(n_startup_trials={n_startup_trials}, n_trials={n_trials})"
)

if sampler_name == "tpe":
return TPESampler()
# TPESampler uses random sampling for first n_startup_trials
return TPESampler(n_startup_trials=n_startup_trials)
elif sampler_name == "random":
# RandomSampler is always random, no startup trials concept
return RandomSampler()
elif sampler_name == "gp":
return GPSampler()
# GPSampler uses random sampling for initial trials
return GPSampler(n_startup_trials=n_startup_trials)
elif sampler_name == "botorch":
return optuna.integration.BoTorchSampler()
# BoTorchSampler uses random sampling for initial trials
return optuna.integration.BoTorchSampler(n_startup_trials=n_startup_trials)
elif sampler_name == "nsga2":
# NSGA2 is a genetic algorithm, no startup trials concept
return NSGAIISampler()
elif sampler_name == "grid":
# Build search space for grid sampler
# GridSampler is deterministic, no startup trials concept
search_space = StudyController._create_search_space(config)
grid_size = StudyController._calculate_grid_size(search_space)
logger.info(
Expand Down
3 changes: 3 additions & 0 deletions auto_tune_vllm/core/trial.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,9 @@ def vllm_args(self) -> List[str]:
if isinstance(value, bool):
if value:
args.append(f"--{cli_param}")
else:
# When False, add --no- prefix for vLLM
args.append(f"--no-{cli_param}")
else:
args.extend([f"--{cli_param}", str(value)])

Expand Down
1 change: 0 additions & 1 deletion auto_tune_vllm/execution/trial_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -853,7 +853,6 @@ def _start_vllm_server(self, trial_config: TrialConfig) -> dict:
str(port),
"--host",
"0.0.0.0",
"--no-enable-prefix-caching",
]

# Add trial-specific parameters
Expand Down
Loading