import json import os import torch ARTIFACT_REPO = "multimodalart/flux-3-action-so101-aoti" INDUCTOR_CONFIGS = {"max_autotune": True, "coordinate_descent_tuning": True, "triton.cudagraphs": False} def block_groups(dit): groups = {"single": list(dit.single_blocks), "mode_txt": list(dit.txt_mode_blocks)} for name, blocks in dit.content_mode_blocks.items(): groups[f"mode_{name}"] = list(blocks) return groups class Fallback: def __init__(self, compiled, eager, state): self.compiled = compiled self.eager = eager self.state = state def __call__(self, *args, **kwargs): if self.state["ok"]: try: return self.compiled(*args, **kwargs) except Exception as e: self.state["ok"] = False self.state["error"] = repr(e)[:300] return self.eager(*args, **kwargs) def load(dit, repo_id=ARTIFACT_REPO, token=None): from huggingface_hub import hf_hub_download from spaces.zero.torch.aoti import LazyAOTIModel, aoti_patch config = json.load(open(hf_hub_download(repo_id, "config.json", token=token))) if config.get("torch") != torch.__version__: raise RuntimeError(f"artifacts built for torch {config.get('torch')}, runtime is {torch.__version__}") state = {"ok": True, "error": None, "groups": []} groups = block_groups(dit) for name in config["groups"]: path = hf_hub_download(repo_id, "package.pt2", subfolder=name, token=token) model = LazyAOTIModel(path) for block in groups[name]: eager = block.forward aoti_patch(block, model) block.forward = Fallback(block.forward, eager, state) state["groups"].append(name) return state def enabled(): return os.environ.get("FLUX3_AOTI", "1") != "0"