diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 656b070d5..288b1e274 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -621,7 +621,7 @@ class BaseRunner(ABC): torch.get_device_module(mr.device).synchronize() mr.tp_group.barrier() with forward_context(ForwardContext(attn_backend=mr.attn_backend)): - with torch.inference_mode(), run_ctx or empty_context(): + with run_ctx or empty_context(): run_once() def _autotune_buffers(self) -> Tuple[Optional[Any], Optional[int]]: diff --git a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py index 09668a2e1..b32751716 100644 --- a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py +++ b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py @@ -190,7 +190,7 @@ def flashinfer_autotune_context(model_runner: ModelRunner, *, skip_logits: bool) maybe_skip_logits = autotune_dummy_run_mode() skip_ops = get_flashinfer_autotune_skip_ops(mr) - with torch.inference_mode(), autotune( + with autotune( True, cache=str(autotune_cache), skip_ops=skip_ops,