Skip to content
Closed
100 changes: 91 additions & 9 deletions python/flydsl/utils/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from typing import Any, Callable, Dict, Generic, Optional, TypeVar

T = TypeVar("T")
NumberT = TypeVar("NumberT", int, float)


class EnvOption(Generic[T]):
Expand Down Expand Up @@ -77,33 +78,85 @@ def parse_value(self, raw: str) -> bool:
return raw.lower() in ("1", "true", "yes", "on")


class OptInt(EnvOption[int]):
"""Integer environment option with optional min/max validation."""
class _OptNumber(EnvOption[NumberT]):
"""Numeric environment option with optional min/max validation."""

def __init__(
self,
default: int = 0,
default: NumberT,
parser: Callable[[str], NumberT],
env_var: Optional[str] = None,
description: str = "",
min_value: Optional[int] = None,
max_value: Optional[int] = None,
min_value: Optional[NumberT] = None,
max_value: Optional[NumberT] = None,
empty_is_default: bool = False,
):
validator = None
if min_value is not None or max_value is not None:

def validator(v: int) -> bool:
def validator(v: NumberT) -> bool:
if min_value is not None and v < min_value:
return False
if max_value is not None and v > max_value:
return False
return True

super().__init__(default, env_var, description, validator)
self.parser = parser
self.min_value = min_value
self.max_value = max_value
self.empty_is_default = empty_is_default

def parse_value(self, raw: str) -> NumberT:
if self.empty_is_default and not raw.strip():
return self.default
return self.parser(raw)


class OptInt(_OptNumber[int]):
"""Integer environment option with optional min/max validation."""

def __init__(
self,
default: int = 0,
env_var: Optional[str] = None,
description: str = "",
min_value: Optional[int] = None,
max_value: Optional[int] = None,
empty_is_default: bool = False,
):
super().__init__(
default,
int,
env_var,
description,
min_value,
max_value,
empty_is_default,
)

def parse_value(self, raw: str) -> int:
return int(raw)

class OptFloat(_OptNumber[float]):
"""Floating-point environment option with optional min/max validation."""

def __init__(
self,
default: float = 0.0,
env_var: Optional[str] = None,
description: str = "",
min_value: Optional[float] = None,
max_value: Optional[float] = None,
empty_is_default: bool = False,
):
super().__init__(
default,
float,
env_var,
description,
min_value,
max_value,
empty_is_default,
)


class OptStr(EnvOption[str]):
Expand Down Expand Up @@ -224,6 +277,33 @@ class AutotuneEnvManager(EnvManager):
config_dir = OptStr("", description="Directory for offline config artifacts; empty disables artifacts")


class AotEnvManager(EnvManager):
"""AOT job options (``FLYDSL_AOT_*`` environment variables)."""

env_prefix = "AOT"

workers = OptInt(
0,
description=(
"Maximum concurrent worker processes; unset, empty, or non-positive values use the CPU and "
"available-memory based automatic limit"
),
empty_is_default=True,
)
Comment on lines +285 to +292
mem_per_worker_gb = OptFloat(
2.0,
description="Assumed GiB per worker for the automatic memory cap; non-positive disables the cap",
)
timeout = OptFloat(
1200.0,
description="Per-job wall-clock timeout in seconds; non-positive disables the timeout",
)
max_retries = OptInt(
2,
description="Retries after an abnormal worker exit or timeout; negative values clamp to zero",
)


class CompileEnvManager(EnvManager):
"""Compile-time options (``FLYDSL_COMPILE_*`` environment variables)."""

Expand Down Expand Up @@ -296,16 +376,18 @@ class RuntimeEnvManager(EnvManager):
enable_cache = OptBool(True, description="Enable kernel caching")
run_only = OptBool(
False,
description=("Skip JIT compilation; only load AOT cache. " "Raise RuntimeError on cache miss."),
description=("Skip JIT compilation; only load AOT cache. Raise RuntimeError on cache miss."),
)


aot = AotEnvManager()
autotune = AutotuneEnvManager()
compile = CompileEnvManager()
debug = DebugEnvManager()
runtime = RuntimeEnvManager()

__all__ = [
"aot",
"autotune",
"compile",
"debug",
Expand Down
Loading
Loading