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)