From a0cf2494fea9c5aa2f4c68892d9cea3e5b659af2 Mon Sep 17 00:00:00 2001 From: cenzhiyao <2523403608@qq.com> Date: Thu, 1 Oct 2026 17:35:16 +0000 Subject: [PATCH] feat(aot): graceful fallback to JIT when AOT compile fails When MAGI_COMPILE_AOT=1, models using SimpleFSDP/DTensor may fail AOT compilation (e.g. "Attempted to read undefined local variable"). Instead of crashing, catch the exception and fall back to the JIT bytecode capture path. This allows AOT to be enabled globally while FSDP models auto-degrade to JIT + inductor cache hits. --- magi_compiler/_api.py | 25 ++++++++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/magi_compiler/_api.py b/magi_compiler/_api.py index 6f91ddf..2c25cb9 100644 --- a/magi_compiler/_api.py +++ b/magi_compiler/_api.py @@ -157,7 +157,30 @@ def _run_orchestration(state: MagiCompileState, args, kwargs): state._ensure_compiled() if state.compile_config.aot: - state.aot_compile(*args, **kwargs) + try: + state.aot_compile(*args, **kwargs) + except Exception as e: + # AOT compile can fail for models using SimpleFSDP/DTensor + # (e.g. "Attempted to read undefined local variable"). + # Fall back to the JIT path so the model still compiles. + magi_logger.warning( + "AOT compile failed (%s: %s), falling back to JIT path " + "for %s", + type(e).__name__, + e, + state.original_code_for_hook, + ) + state.compile_config = state.compile_config.model_copy( + update={"aot": False} + ) + # Reset compiled_entry: it was created with AOT-only + # guard_filter_fn, need a fresh one for JIT bytecode + # capture. + state.compiled_entry = None + torch._dynamo.reset() + state._ensure_compiled() + with state._jit_capture_compiled_bytecode(): + return state.compiled_entry(*args, **kwargs) else: with state._jit_capture_compiled_bytecode(): return state.compiled_entry(*args, **kwargs)