InstructAV2AV / app.py
suimu's picture
init
e0177dc
Raw
History Blame Contribute Delete
3.5 kB
"""Hugging Face Space entry point for InstructAV2AV."""
from __future__ import annotations
import logging
import os
import tempfile
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import gradio as gr
try:
import spaces
except ImportError:
class _SpacesFallback:
@staticmethod
def GPU(*decorator_args: Any, **_decorator_kwargs: Any):
if decorator_args and callable(decorator_args[0]):
return decorator_args[0]
return lambda function: function
spaces = _SpacesFallback()
from scripts.demo import CSS, DEFAULT_CONFIG, DemoRuntime, build_demo
from space_model_store import HubModelStore
logging.basicConfig(
level=logging.INFO,
format="[%(asctime)s] %(levelname)s: %(message)s",
)
def default_hf_home() -> Path:
persistent_storage = Path("/data")
if persistent_storage.is_dir() and os.access(persistent_storage, os.W_OK):
return persistent_storage / ".huggingface"
return Path.home() / ".cache" / "huggingface"
HF_HOME = Path(
os.getenv("HF_HOME", default_hf_home())
).expanduser()
MODEL_HOME = Path(
os.getenv("INSTRUCTAV2AV_MODEL_HOME", HF_HOME / "instructav2av")
).expanduser()
OUTPUT_DIR = Path(
os.getenv(
"INSTRUCTAV2AV_OUTPUT_DIR",
Path(tempfile.gettempdir()) / "instructav2av_outputs",
)
).expanduser()
GPU_DURATION = int(os.getenv("INSTRUCTAV2AV_ZEROGPU_DURATION", "300"))
model_store = HubModelStore(MODEL_HOME, hub_cache_dir=HF_HOME / "hub")
if os.getenv("INSTRUCTAV2AV_EAGER_DOWNLOAD", "1").lower() not in {
"0",
"false",
"no",
}:
try:
model_store.preload_default()
logging.info("Shared weights and the General checkpoint are cached.")
except Exception:
logging.exception(
"Startup model download failed; the first edit request will retry it."
)
runtime_args = SimpleNamespace(
config_file=str(DEFAULT_CONFIG),
model_dir=str(model_store.ckpt_dir / "InstructAV2AV"),
ckpt_dir=str(model_store.ckpt_dir),
output_dir=str(OUTPUT_DIR),
device=int(os.getenv("INSTRUCTAV2AV_CUDA_DEVICE", "0")),
no_cpu_offload=os.getenv("INSTRUCTAV2AV_CPU_OFFLOAD", "0").lower()
in {"0", "false", "no"},
)
runtime = DemoRuntime(
runtime_args,
checkpoint_resolver=model_store.resolve_checkpoint,
)
@spaces.GPU(duration=GPU_DURATION)
def zero_gpu_task(
operation: str,
video_value: Any,
instruction: str,
model_key: str,
seed: float,
sample_steps: float,
video_guidance_scale: float,
audio_guidance_scale: float,
progress: gr.Progress = gr.Progress(),
) -> str:
if operation == "warmup":
return runtime.warmup("general", progress)
if operation != "generate":
raise gr.Error(f"Unsupported ZeroGPU operation: {operation}")
return runtime.generate(
video_value,
instruction,
model_key,
seed,
sample_steps,
video_guidance_scale,
audio_guidance_scale,
progress,
)
demo = build_demo(
runtime,
zero_gpu_fn=zero_gpu_task,
)
demo.queue(max_size=8, default_concurrency_limit=1)
if __name__ == "__main__":
demo.launch(
server_name="0.0.0.0",
server_port=int(os.getenv("PORT", "7860")),
allowed_paths=[str(runtime.output_dir)],
show_error=True,
theme=gr.themes.Default(),
css=CSS,
max_file_size=os.getenv("GRADIO_MAX_FILE_SIZE", "500mb"),
)