flux-3-action-so101-sim / aoti_blocks.py
multimodalart's picture
multimodalart HF Staff
Load AoTI-compiled DiT blocks with eager fallback
b91cde0 verified
Raw History Blame Contribute Delete
1.84 kB
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"