Spaces:
Running on Zero
Running on Zero
Download aoti_blocks.py from multimodalart/flux-3-action-so101-sim: direct link, hf CLI and curl.
- Browser
- Download file 1.84 kB
-
https://huggingface.co/spaces/multimodalart/flux-3-action-so101-sim/resolve/main/aoti_blocks.py
- Command line
-
hf download hf://spaces/multimodalart/flux-3-action-so101-sim/aoti_blocks.py
-
curl -L -o aoti_blocks.py https://huggingface.co/spaces/multimodalart/flux-3-action-so101-sim/resolve/main/aoti_blocks.py
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" | |