Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 24 additions & 1 deletion magi_compiler/_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading