andrasnest's picture
Wan Loop Generator
7d03019
Raw
History Blame Contribute Delete
11.3 kB
from __future__ import annotations
import os
import json
import atexit
from pathlib import Path
import time
from datetime import datetime, timezone
# ZeroGPU v2 serializes global virtual CUDA tensors before the replica starts.
# Keep the service-managed NVMe location, matching the reference Space.
os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1")
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
# PyTorch's default CUDA availability check invokes the CUDA device-count API,
# which poisons forked children. NVML discovery is explicitly fork-safe.
os.environ.setdefault("PYTORCH_NVML_BASED_CUDA_CHECK", "1")
import spaces
import torch
_ZERO_GPU = os.environ.get("SPACES_ZERO_GPU", "").strip().lower() in {"1", "true", "yes", "on"}
def _is_zerogpu_parent() -> bool:
if not _ZERO_GPU:
return False
try:
from spaces.zero import wrappers as zero_wrappers
return not zero_wrappers.forked
except Exception:
return True
def _assert_cuda_fork_clean(event: str) -> None:
"""Reproduce the CUDA bad-fork check used by the ZeroGPU worker."""
read_fd, write_fd = os.pipe()
pid = os.fork()
if pid == 0:
os.close(read_fd)
try:
verdict = b"1" if torch.cuda._is_in_bad_fork() else b"0"
os.write(write_fd, verdict)
finally:
os.close(write_fd)
os._exit(0)
os.close(write_fd)
try:
verdict = os.read(read_fd, 1)
finally:
os.close(read_fd)
_, status = os.waitpid(pid, 0)
if status != 0 or verdict != b"0":
raise RuntimeError("ZeroGPU parent CUDA state poisons the worker fork")
print(f'[WAN_SERVICE] {{"bad_fork": false, "event": "{event}"}}', flush=True)
def _install_parent_cuda_firewall() -> None:
"""Keep discovery fork-safe in the parent and restore real CUDA in workers."""
original_is_available = torch.cuda.is_available
original_device_count = torch.cuda.device_count
original_is_bf16_supported = torch.cuda.is_bf16_supported
original_get_allocator_backend = torch.cuda.get_allocator_backend
original_lazy_init = torch.cuda._lazy_init
original_c_device_count = torch._C._cuda_getDeviceCount
original_c_init = torch._C._cuda_init
def dispatch(original, parent_value):
def wrapped(*args, **kwargs):
if _is_zerogpu_parent():
return parent_value
return original(*args, **kwargs)
return wrapped
def guarded_lazy_init(*args, **kwargs):
if _is_zerogpu_parent():
raise RuntimeError("CUDA context initialization attempted in ZeroGPU parent")
return original_lazy_init(*args, **kwargs)
def guarded_c_init(*args, **kwargs):
if _is_zerogpu_parent():
raise RuntimeError("Direct CUDA initialization attempted in ZeroGPU parent")
return original_c_init(*args, **kwargs)
torch.cuda.is_available = dispatch(original_is_available, False)
torch.cuda.device_count = dispatch(original_device_count, 0)
torch.cuda.is_bf16_supported = dispatch(original_is_bf16_supported, True)
torch.cuda.get_allocator_backend = dispatch(original_get_allocator_backend, "zerogpu")
torch.cuda._lazy_init = guarded_lazy_init
torch.cuda.init = guarded_lazy_init
torch._C._cuda_getDeviceCount = dispatch(original_c_device_count, 0)
torch._C._cuda_init = guarded_c_init
if _ZERO_GPU:
_assert_cuda_fork_clean("parent.cuda_import_clean")
_install_parent_cuda_firewall()
import gradio as gr
if _ZERO_GPU:
_assert_cuda_fork_clean("parent.cuda_gradio_clean")
from core_anim.space_config import load_space_config
from core_anim.space_postprocess import encode_video_with_preview
from core_anim.space_runtime import LoopGeneratorService
if _ZERO_GPU:
_assert_cuda_fork_clean("parent.cuda_runtime_imports_clean")
CONFIG = load_space_config()
_SERVICE = None
if _ZERO_GPU:
_assert_cuda_fork_clean("parent.cuda_config_clean")
def _shutdown_global_service() -> None:
global _SERVICE
service, _SERVICE = _SERVICE, None
if service is not None:
service.close()
def get_service() -> LoopGeneratorService:
global _SERVICE
if _SERVICE is None:
_SERVICE = LoopGeneratorService(CONFIG)
return _SERVICE
# ZeroGPU optimizes CUDA placement performed during module initialization.
# Tests and manifest validation can explicitly skip the multi-GB model setup.
if os.environ.get("SPACE_SKIP_MODEL_LOAD") != "1":
_SERVICE = LoopGeneratorService(CONFIG)
if _ZERO_GPU:
if torch.cuda.is_initialized():
raise RuntimeError("ZeroGPU parent initialized CUDA before worker fork")
print('[WAN_SERVICE] {"cuda_initialized": false, "event": "parent.cuda_clean"}', flush=True)
_assert_cuda_fork_clean("parent.cuda_fork_clean")
atexit.register(_shutdown_global_service)
def _diag(event: str, **fields) -> None:
payload = {
"ts": datetime.now(timezone.utc).isoformat(timespec="milliseconds"),
"event": event,
**fields,
}
print(f"[WAN_JOB] {json.dumps(payload, sort_keys=True)}", flush=True)
def validate_request(image, prompt: str):
if image is None:
raise gr.Error("Upload an image.")
if not (prompt or "").strip():
raise gr.Error("Enter a prompt.")
return image, prompt
def finish_video(frame_bundle: str, output_format: str) -> tuple[str, str]:
return encode_video_with_preview(
frame_bundle,
fps=CONFIG.fps,
output_format=output_format,
crossfade=CONFIG.use_crossfade,
)
@spaces.GPU(duration=CONFIG.gpu_duration, size=CONFIG.gpu_size)
def generate_loop(image, prompt: str, output_format: str, progress=gr.Progress(track_tqdm=False)):
image, prompt = validate_request(image, prompt)
if os.environ.get("SPACE_SKIP_MODEL_LOAD") == "1":
raise RuntimeError("Model runtime disabled by SPACE_SKIP_MODEL_LOAD.")
service = None
frame_bundle = None
started = time.perf_counter()
outcome = "failed"
_diag(
"job.accepted",
input_width=getattr(image, "width", None),
input_height=getattr(image, "height", None),
)
try:
progress(0.0, desc="Requesting GPU…")
yield gr.skip(), gr.skip(), "Requesting GPU…"
progress(0.03, desc="Initializing runtime…")
yield gr.skip(), gr.skip(), "Initializing runtime…"
init_started = time.perf_counter()
service = get_service()
_diag("runtime.init.done", elapsed_s=round(time.perf_counter() - init_started, 3))
for stage_result in service.generate_iter(
image,
prompt,
progress_callback=progress,
):
if stage_result is None: # Compatibility and cancellation checkpoint.
yield gr.skip(), gr.skip(), "Processing…"
elif isinstance(stage_result, dict) and "stage" in stage_result:
yield gr.skip(), gr.skip(), stage_result["stage"]
elif isinstance(stage_result, dict) and "bundle" in stage_result:
frame_bundle = stage_result["bundle"]
else: # Compatibility with simple test doubles.
frame_bundle = stage_result
# Release per-job diffusion/VAE objects before CPU-only video encoding;
# reusable process-global components remain prepared for the next job.
service.cleanup_job()
format_name = "MKV" if output_format == "mkv" else "MP4"
encoding_label = f"Encoding {format_name} and preview…" if output_format == "mkv" else "Encoding MP4…"
progress(0.97, desc=encoding_label)
yield gr.skip(), gr.skip(), encoding_label
encode_started = time.perf_counter()
_diag("video.encode.start", output_format=output_format)
preview, download = finish_video(frame_bundle, output_format)
_diag(
"video.encode.done",
output_format=output_format,
elapsed_s=round(time.perf_counter() - encode_started, 3),
)
frame_bundle = None # encode_video_with_preview removes its source bundle.
progress(1.0, desc="Complete")
outcome = "complete"
yield preview, download, "Complete"
except GeneratorExit:
outcome = "cancelled"
_diag("job.cancelled", elapsed_s=round(time.perf_counter() - started, 3))
raise
except BaseException as exc:
_diag(
"job.failed",
elapsed_s=round(time.perf_counter() - started, 3),
error_type=type(exc).__name__,
error=str(exc)[:500],
)
raise
finally:
_diag("job.cleanup.start", outcome=outcome)
if service is not None:
service.cleanup_job()
if frame_bundle is not None:
Path(frame_bundle).unlink(missing_ok=True)
_diag(
"job.cleanup.done",
outcome=outcome,
elapsed_s=round(time.perf_counter() - started, 3),
)
with gr.Blocks(title="Wan Loop Generator", delete_cache=(86400, 86400)) as demo:
gr.Markdown(
"# Wan Loop Generator\n"
"This experimental Space is a test for generating one-second video loops from a single image."
)
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### Image Upload")
image_input = gr.Image(
type="pil",
label="Image",
sources=["upload"],
height=420,
)
prompt_input = gr.Textbox(label="Prompt", placeholder="Describe the motion of the loop…", lines=4)
format_input = gr.Dropdown(
choices=[
("MKV (High Quality)", "mkv"),
("MP4 (Compressed)", "mp4"),
],
value="mkv",
label="Output format",
interactive=True,
)
with gr.Row():
generate_button = gr.Button("Generate loop", variant="primary")
cancel_button = gr.Button("Cancel", variant="stop")
with gr.Column(scale=1):
gr.Markdown("### Progress / Output")
status_output = gr.Textbox(
label="Status",
value="Ready",
interactive=False,
)
preview_output = gr.Video(
label="Preview",
format="mp4",
height=360,
autoplay=True,
loop=True,
)
download_output = gr.File(
label="Download generated loop",
)
generation_event = generate_button.click(
generate_loop,
inputs=[image_input, prompt_input, format_input],
outputs=[preview_output, download_output, status_output],
concurrency_limit=CONFIG.concurrency_limit,
api_name="generate_loop",
show_progress="full",
)
cancel_button.click(
fn=None,
cancels=[generation_event],
queue=False,
api_visibility="private",
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=CONFIG.concurrency_limit, max_size=8).launch(ssr_mode=False)