diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 4f2d8cd54a..5b32103dc8 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -1697,9 +1697,20 @@ def _kv_unified_from_args( return enabled -def _flash_attn_enabled_from_args(args: Optional[Iterable[str]], default: bool = True) -> bool: - """Resolve llama.cpp's last-wins flash-attention CLI setting.""" +def _flash_attn_enabled_from_args( + args: Optional[Iterable[str]], + default: bool = True, + env: Optional[Mapping[str, str]] = None, +) -> bool: + """Resolve llama.cpp's environment and last-wins flash-attention settings.""" enabled = default + # llama.cpp applies LLAMA_ARG_FLASH_ATTN before parsing argv (arg.cpp set_env), + # so the CLI still wins. --flash-attn has no args_neg, so no LLAMA_ARG_NO_ twin. + value = (os.environ if env is None else env).get("LLAMA_ARG_FLASH_ATTN") + if value in _LLAMA_ARG_FALSE_VALUES: + enabled = False + elif value in _LLAMA_ARG_TRUE_OR_AUTO_VALUES: + enabled = True values = [str(arg) for arg in args] if args else [] for i, raw in enumerate(values): if _flag_name(raw) not in {"-fa", "--flash-attn"}: @@ -9054,7 +9065,8 @@ class LlamaCppBackend: int(self._DEFAULT_N_UBATCH if _effective_ubatch is None else _effective_ubatch), ) self._flash_attn_enabled = ( - _flash_attn_enabled_from_args(_last_spawn_cmd) and self._architecture != "grok" + _flash_attn_enabled_from_args(_last_spawn_cmd, env = env) + and self._architecture != "grok" ) self._effective_cache_types = _effective_main_cache_types( _last_spawn_cmd, diff --git a/studio/backend/tests/test_mtp_vram_budget.py b/studio/backend/tests/test_mtp_vram_budget.py index 77ca76325f..3742018e5e 100644 --- a/studio/backend/tests/test_mtp_vram_budget.py +++ b/studio/backend/tests/test_mtp_vram_budget.py @@ -817,7 +817,26 @@ class TestExtraArgsMtpDetection: ], ) def test_flash_attn_last_value_wins(self, args, expected): - assert _flash_attn_enabled_from_args(args) is expected + assert _flash_attn_enabled_from_args(args, env = {}) is expected + + @pytest.mark.parametrize( + "value,expected", + [ + ("off", False), + ("disabled", False), + ("false", False), + ("0", False), + ("on", True), + ("auto", True), + ("garbage", True), # llama.cpp refuses to start, so the default is moot + ], + ) + def test_flash_attn_env_applies(self, value, expected): + env = {"LLAMA_ARG_FLASH_ATTN": value} + assert _flash_attn_enabled_from_args([], env = env) is expected + # llama.cpp parses the environment first, so an explicit flag still wins. + assert _flash_attn_enabled_from_args(["-fa", "on"], env = env) is True + assert _flash_attn_enabled_from_args(["-fa", "off"], env = env) is False def test_effective_main_cache_types_follow_env_then_cli(self): env = {