redkits's picture
Upload server.py with huggingface_hub
06447aa verified
Raw
History Blame Contribute Delete
4.09 kB
import os
import sys
import subprocess
import base64
import tempfile
import traceback
# Install latest diffusers + dependencies
subprocess.check_call([sys.executable, "-m", "pip", "install", "-q", "--upgrade",
"diffusers", "flask", "accelerate", "sentencepiece", "protobuf", "imageio[ffmpeg]", "transformers", "huggingface_hub"])
import torch
from flask import Flask, request, jsonify
print(f"[video] torch version: {torch.__version__}", flush=True)
print(f"[video] CUDA available: {torch.cuda.is_available()}", flush=True)
if torch.cuda.is_available():
print(f"[video] GPU: {torch.cuda.get_device_name(0)}", flush=True)
print(f"[video] VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB", flush=True)
app = Flask(__name__)
pipe = None
def load_model():
global pipe
from diffusers import AutoencoderKLWan, WanPipeline
print("[video] Loading Wan2.2-TI2V-5B pipeline...", flush=True)
vae = AutoencoderKLWan.from_pretrained(
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
subfolder="vae",
torch_dtype=torch.bfloat16,
)
pipe = WanPipeline.from_pretrained(
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
vae=vae,
torch_dtype=torch.bfloat16,
)
pipe.enable_model_cpu_offload()
print("[video] Model loaded with CPU offload!", flush=True)
@app.route("/health", methods=["GET"])
def health():
if pipe is not None:
return jsonify({"status": "healthy"}), 200
return jsonify({"status": "loading"}), 503
@app.route("/", methods=["POST"])
def generate():
try:
from diffusers.utils import export_to_video
data = request.get_json(force=True)
# Support both HF format (inputs) and backend format (prompt)
prompt = data.get("inputs", "") or data.get("prompt", "")
params = data.get("parameters", {})
num_frames = int(params.get("num_frames", 25))
height = int(params.get("height", 480))
width = int(params.get("width", 832))
steps = int(params.get("num_inference_steps", 15))
fps = int(params.get("fps", 24))
guidance = float(params.get("guidance_scale", 5.0))
negative = params.get("negative_prompt", "low quality, blurry, distorted")
print(f"[video] Generating {num_frames} frames: {prompt[:100]}", flush=True)
print(f"[video] Params: {width}x{height}, steps={steps}, fps={fps}", flush=True)
frames = pipe(
prompt=prompt,
negative_prompt=negative,
num_frames=num_frames,
height=height,
width=width,
num_inference_steps=steps,
guidance_scale=guidance,
).frames[0]
# Write to temp file (export_to_video needs a file path, not BytesIO)
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp:
tmp_path = tmp.name
export_to_video(frames, tmp_path, fps=fps)
with open(tmp_path, "rb") as f:
video_bytes = f.read()
os.unlink(tmp_path)
video_b64 = base64.b64encode(video_bytes).decode("utf-8")
duration = len(frames) / fps
print(f"[video] Done: {len(frames)} frames, {duration:.1f}s, {len(video_bytes)} bytes", flush=True)
return jsonify({
"video": video_b64,
"format": "mp4",
"frames": len(frames),
"duration": duration,
"success": True,
"video_base64": video_b64,
"duration_seconds": duration,
"model": "wan2.2-ti2v-5b",
})
except Exception as e:
tb = traceback.format_exc()
print(f"[video] ERROR: {e}\n{tb}", flush=True)
return jsonify({"error": str(e), "traceback": tb, "success": False}), 500
# Also add /api/generate-video route for backend compatibility
@app.route("/api/generate-video", methods=["POST"])
def api_generate_video():
"""Handle the backend's /api/generate-video format by forwarding to generate()"""
return generate()
if __name__ == "__main__":
load_model()
app.run(host="0.0.0.0", port=8000)