import argparse import gc import importlib.util import io import json import os import sys import threading from pathlib import Path # Keep Gradio's local health check away from cluster/VS Code proxy settings. os.environ["NO_PROXY"] = "localhost,127.0.0.1,0.0.0.0" os.environ["no_proxy"] = os.environ["NO_PROXY"] import cv2 import gradio as gr import numpy as np import spaces import torch from diffusers import AutoencoderKL, UNetSpatioTemporalConditionModel from huggingface_hub import HfApi, hf_hub_download, snapshot_download from PIL import Image ROOT = Path(__file__).resolve().parent sys.path.insert(0, str(ROOT)) def load_module(name, path): spec = importlib.util.spec_from_file_location(name, path) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) return module def load_wan_module(): path = ROOT / "GeoNeXt-Wan" / "model_inference.py" return load_module("geonext_wan_model_inference", path) wan = load_wan_module() inference_lock = threading.Lock() pipe = None checkpoint = None model_args = None svd_pipe = None svd_helpers = None stats_lock = threading.Lock() STATS_REPO_ID = "happy0612/GeoNeXt-demo-stats" STATS_FILENAME = "stats.json" def format_prediction_count(count): return f"**{count:,} successful prediction{'s' if count != 1 else ''}**" def read_prediction_count(): token = os.environ.get("HF_TOKEN") if not token: return 0 try: path = hf_hub_download( repo_id=STATS_REPO_ID, filename=STATS_FILENAME, repo_type="dataset", token=token, force_download=True, ) with open(path, encoding="utf-8") as handle: return int(json.load(handle).get("successful_predictions", 0)) except Exception as error: print("[Stats] Could not read prediction count:", type(error).__name__) return 0 def increment_prediction_count(): token = os.environ.get("HF_TOKEN") if not token: return None with stats_lock: try: count = read_prediction_count() + 1 payload = json.dumps( {"successful_predictions": count}, indent=2 ).encode("utf-8") HfApi(token=token).upload_file( path_or_fileobj=io.BytesIO(payload), path_in_repo=STATS_FILENAME, repo_id=STATS_REPO_ID, repo_type="dataset", commit_message="Update successful prediction count", ) return count except Exception as error: print("[Stats] Could not update prediction count:", type(error).__name__) return None def get_prediction_count_text(): return format_prediction_count(read_prediction_count()) def unload_cuda(): gc.collect() torch.cuda.empty_cache() def unload_wan_pipeline(): global pipe if pipe is not None: pipe.load_models_to_device([]) del pipe pipe = None unload_cuda() def unload_svd_pipeline(): global svd_pipe if svd_pipe is not None: svd_pipe.to("cpu") del svd_pipe svd_pipe = None unload_cuda() def load_wan_pipeline(): global pipe, checkpoint, model_args if pipe is None: checkpoint = hf_hub_download( repo_id="happy0612/GeoNeXt", filename="GeoNeXt-Wan/geonext_wan.safetensors", ) base_model = snapshot_download(repo_id="Wan-AI/Wan2.1-T2V-1.3B") model_args = argparse.Namespace( wan_model_dir=base_model, override_vae_path="", vae_backend="wan", norm_type="trunc_disparity", rgb_condition_mode="concat", rgb_condition_scale=1.0, target_modalities="depth,normal", num_condition_frames=1, temporal_rope_scale=8, ) pipe = wan.build_pipe_for_depth_normal(model_args) state_dict = wan.load_state_dict(checkpoint) load_result = pipe.dit.load_state_dict(state_dict, strict=False) if load_result.unexpected_keys: print("Unexpected checkpoint keys:", len(load_result.unexpected_keys)) if load_result.missing_keys: print("Missing checkpoint keys:", len(load_result.missing_keys)) return pipe def load_svd_pipeline(): """Load SVD lazily on CPU so both backends can share one GPU Space.""" global svd_pipe, svd_helpers if svd_pipe is not None: return svd_pipe, svd_helpers token = os.environ.get("HF_TOKEN") if not token: raise gr.Error( "HF_TOKEN is not available in the Space runtime. Add it under " "Settings > Variables and secrets as a Secret, then restart the Space." ) try: auth = HfApi(token=token).whoami() HfApi(token=token).model_info( "stabilityai/stable-video-diffusion-img2vid-xt-1-1" ) print("[SVD Auth] authenticated user:", auth.get("name", "unknown")) print("[SVD Auth] gated SVD repository access: OK") except Exception as error: print("[SVD Auth] gated repository check failed:", type(error).__name__) raise gr.Error( "HF_TOKEN exists but cannot read the gated SVD repository. Create a " "fine-grained token with 'Read access to contents of all public gated " "repos you can access', replace the Space secret, and restart." ) from error svd_root = ROOT / "GeoNeXt-SVD" svd = load_module("geonext_svd_inference", svd_root / "inference.py") pipeline_module = load_module("geonext_svd_pipeline", svd_root / "pipeline.py") checkpoint_root = snapshot_download( repo_id="happy0612/GeoNeXt", allow_patterns="GeoNeXt-SVD/**", ) checkpoint = str(Path(checkpoint_root) / "GeoNeXt-SVD") dtype = torch.float16 vae = AutoencoderKL.from_pretrained( "stabilityai/sd-vae-ft-mse", torch_dtype=dtype, ) unet = UNetSpatioTemporalConditionModel.from_pretrained( checkpoint, subfolder="unet", torch_dtype=dtype, low_cpu_mem_usage=False, ) svd_pipe = pipeline_module.GeoNeXtPipeline.from_pretrained( svd.SVD_BASE_MODEL, unet=unet, vae=vae, variant="fp16", torch_dtype=dtype, low_cpu_mem_usage=False, token=token, ) svd_pipe.set_progress_bar_config(disable=True) svd_helpers = svd return svd_pipe, svd_helpers @torch.inference_mode() def predict_wan(image, steps): if image is None: raise gr.Error("Please upload an image first.") image = Image.fromarray(np.asarray(image, dtype=np.uint8), mode="RGB") with inference_lock: unload_svd_pipeline() local_pipe = load_wan_pipeline() original_width, original_height = image.size prepared, (content_width, content_height) = wan._prepare_image_for_inference( image, processing_res=768, pipe=local_pipe, processing_res_side="long", ) width, height = prepared.size latents = local_pipe( prompt="", negative_prompt="", input_video=[prepared], seed=0, rand_device="cuda", cfg_scale=1.0, num_inference_steps=int(steps), num_frames=9, height=height, width=width, tiled=False, tile_size=(30, 52), tile_stride=(15, 26), zero_noise=False, direct_clean_output=False, output_type="latent", ) if latents.shape[2] < 3: raise RuntimeError("GeoNeXt-Wan returned fewer than three latent frames.") local_pipe.load_models_to_device(["vae"]) outputs = [] for index, name in ((1, "depth"), (2, "normal")): decoded = local_pipe.vae.decode( latents[:, :, index:index + 1], device=local_pipe.device, tiled=False, tile_size=(30, 52), tile_stride=(15, 26), ) visual = wan._frame_tensor_to_vis( decoded[0, :, 0], name, norm_type="trunc_disparity" ) visual = visual.crop((0, 0, content_width, content_height)) visual = visual.resize((original_width, original_height), Image.BILINEAR) outputs.append(np.asarray(visual)) # Keep the most recently used backend warm. ZeroGPU releases the # physical GPU after this call; the pipeline is destroyed only when # the user switches to SVD (or when the Space restarts). return outputs def predict_svd(image, steps): original = Image.fromarray(np.asarray(image, dtype=np.uint8), mode="RGB") with inference_lock: # A 24 GB GPU cannot retain both pipelines. Destroy Wan before loading # SVD, and destroy SVD after producing CPU outputs. unload_wan_pipeline() local_svd_pipe, svd = load_svd_pipeline() resized = svd._resize(original, 768, "long") width, height = resized.size width = max(64, round(width / 64) * 64) height = max(64, round(height / 64) * 64) resized = resized.resize((width, height), Image.Resampling.BICUBIC) local_svd_pipe.to("cuda") generator = torch.Generator(device="cuda").manual_seed(0) with torch.autocast("cuda", dtype=torch.float16): prediction = local_svd_pipe( resized, num_frames=3, width=width, height=height, min_guidance_scale=1.0, max_guidance_scale=1.2, noise_aug_strength=0.0, decode_chunk_size=8, generator=generator, motion_bucket_id=127, fps=7, num_inference_steps=int(steps), ) depth = prediction.geo_res[0].mean(dim=1).squeeze().float().cpu().numpy() normal = prediction.geo_res[1].squeeze().permute(1, 2, 0).float().cpu().numpy() del prediction # Keep SVD alive for consecutive requests. It is unloaded by # predict_wan() only when the user switches back to Wan. depth = np.asarray( Image.fromarray(depth, mode="F").resize( original.size, Image.Resampling.BILINEAR ) ) normal = np.stack( [ np.asarray( Image.fromarray(normal[..., channel], mode="F").resize( original.size, Image.Resampling.BILINEAR ) ) for channel in range(3) ], axis=-1, ) from utils.visualization import depth_to_vis, normal_to_vis depth_vis = depth_to_vis(np.clip(depth, 0.0, 1.0), reverse_color=True) normal_vis = normal_to_vis(np.clip(normal, -1.0, 1.0)) return np.asarray(depth_vis), np.asarray(normal_vis) @spaces.GPU(duration=180) def predict(image, steps, backend, output_view): if image is None: raise gr.Error("Please upload an image first.") if backend == "GeoNeXt-SVD": depth, normal = predict_svd(image, steps) else: depth, normal = predict_wan(image, steps) selected = normal if output_view == "Surface Normal" else depth count = increment_prediction_count() count_text = format_prediction_count(count) if count is not None else gr.skip() return selected, depth, normal, count_text def select_output(output_view, depth, normal): if depth is None or normal is None: return None return normal if output_view == "Surface Normal" else depth header = """
Predict monocular depth and surface normals with video generative priors.