"""WaveDiT - 3D Brain MRI Generator (Hugging Face ZeroGPU Gradio Space). Single-page Gradio app: * pick an age (and optionally model / seed / steps / CFG / sampler / Morpheus), * generate a full-resolution synthetic 3D brain MRI on a ZeroGPU GPU, * explore it in-browser with an MRIcroGL-like Niivue viewer (triplane + 3D render), * sweep an "aging time-lapse", apply presets, switch colormap/gamma live, * download the float32 NIfTI and read an always-on generation/VRAM badge. The wavedit/ package is VENDORED into this Space (do not pip-install it: its pyproject pins torch==2.6.0, which is unsupported on ZeroGPU >= 2.8). The only source patch is a lazy `import wandb` in BOTH - wavedit/models/wavelet_flow_matching.py (used at sampling-import time), and - wavedit/training/trainer.py (pulled in transitively because wavedit/training/__init__.py imports Trainer, and app.py imports load_model_weights from wavedit.training.checkpoint). Research artifact only - synthetic images, not a medical device, not for clinical use. """ from __future__ import annotations import os # CRITICAL: select the pure-PyTorch neighborhood-attention fallback BEFORE importing # wavedit (NATTEN is unavailable on Spaces; the torch backend is numerically equivalent). os.environ.setdefault("WAVEDIT_NA_BACKEND", "torch") import base64 import gzip import html as html_lib import random import time import traceback from pathlib import Path import gradio as gr import nibabel as nib import numpy as np import torch try: import spaces # True no-op off ZeroGPU; gates the real GPU on @spaces.GPU. _HAS_SPACES = True except Exception: # pragma: no cover - running fully off-platform _HAS_SPACES = False class _SpacesShim: @staticmethod def GPU(*dargs, **dkwargs): def _wrap(fn): return fn # Support both @spaces.GPU and @spaces.GPU(duration=...) if dargs and callable(dargs[0]) and not dkwargs: return dargs[0] return _wrap spaces = _SpacesShim() # type: ignore from huggingface_hub import hf_hub_download from wavedit import Config from wavedit.models import build_model # IMPORT load_model_weights DIRECTLY from .training.checkpoint. NOTE: this still runs # wavedit/training/__init__.py (which imports trainer.py); the vendored trainer.py has # the lazy-wandb patch, so this import succeeds even when wandb is not installed. from wavedit.training.checkpoint import load_model_weights from wavedit.generation.generator import center_crop_bounds # --------------------------------------------------------------------------- # # Constants # --------------------------------------------------------------------------- # HF_REPO = "danesed/WaveDiT" HF_REVISION = "main" CHECKPOINTS = { "Base (fast)": "WaveDiT-Base.pth", "FinePatch (detailed)": "WaveDiT-FinePatch.pth", "Wide (largest)": "WaveDiT-Wide.pth", } DEFAULT_MODEL = "Base (fast)" # Standard MNI-like target grid; FULL_SIZE is read per-checkpoint at build time so a # future Deep/Wide variant with a different image_size still crops correctly. DEFAULT_FULL_SIZE = (224, 224, 224) CROP_SIZE = (182, 218, 182) # Fallbacks if a checkpoint somehow lacks an age range (the real ones carry 6..95). AGE_MIN_FALLBACK, AGE_MAX_FALLBACK = 6, 95 AGE_DEFAULT = 72 SEED_DEFAULT = 42 SEED_MAX = 2_147_483_647 STEPS_DEFAULT = 10 STEPS_MIN, STEPS_MAX = 1, 200 # Aging time-lapse runs many frames; keep each frame within one short GPU window and # bound total session work. Per-frame steps are clamped to SWEEP_STEPS_MAX and the # frame count is bounded so a worst-case FinePatch sweep never overruns the per-call # @spaces.GPU budget (each frame is its OWN GPU call -- see gpu_sweep_frame). SWEEP_FRAMES_MIN, SWEEP_FRAMES_MAX = 3, 8 SWEEP_STEPS_MAX = 50 CFG_MIN, CFG_MAX, CFG_DEFAULT = 1.0, 8.0, 1.0 CFG_RESCALE = 0.7 # fixed; only active when cfg_scale != 1.0 MORPHEUS_DEFAULT = 1.0 SAMPLERS = ["Heun", "Euler"] SAMPLER_MAP = {"Heun": "heun", "Euler": "euler"} COLORMAPS = ["gray", "bone", "viridis", "plasma", "inferno", "magma", "hot", "cubehelix"] DELIGHT_COLORMAPS = ["viridis", "plasma", "bone", "magma"] SPACE_VERSION = "0.1.0" DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") HERE = Path(__file__).resolve().parent VIEWER_TEMPLATE_PATH = HERE / "viewer_template.html" HERO_NII_PATH = HERE / "assets" / "wavedit_hero_finepatch_age72_seed42.nii.gz" DOWNLOAD_DIR = Path(os.environ.get("WAVEDIT_OUT_DIR", "/tmp/wavedit_outputs")) DOWNLOAD_DIR.mkdir(parents=True, exist_ok=True) # --------------------------------------------------------------------------- # # Links / copy # --------------------------------------------------------------------------- # LINK_PAPER = "https://arxiv.org/abs/2606.08670" LINK_CODE = "https://github.com/sisinflab/WaveDiT" LINK_PROJECT = "https://danesed.github.io/wavedit-page/" LINK_MODEL = "https://huggingface.co/danesed/WaveDiT" CITATION = """@misc{danese2026waveditdistributionawarewaveletflow, title = {WaveDiT: Distribution-Aware Wavelet Flow Matching for Efficient 3D Brain MRI Synthesis}, author = {Danilo Danese and Angela Lombardi and Giuseppe Fasano and Matteo Attimonelli and Tommaso Di Noia}, year = {2026}, eprint = {2606.08670}, archivePrefix = {arXiv}, primaryClass = {cs.CV}, url = {https://arxiv.org/abs/2606.08670} }""" # --------------------------------------------------------------------------- # # Viewer template # --------------------------------------------------------------------------- # _INLINE_VIEWER_TEMPLATE = """