CoolFace
Modelpublic

rdsi2026/Wan2.2-I2V-A14B-Diffusers

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes3downloads
handler.py121 linesDownload Raw Back to root
1"""
2handler.py
3----------
4Handler personalizado para un Inference Endpoint dedicado de Hugging Face.
5
6Este fichero se sube AL REPO DEL MODELO (no al Space), en la raiz, junto
7a los pesos. Cuando creas el Endpoint desde la web de HF, el sistema
8detecta este handler.py y lo usa en lugar de intentar mapear el modelo
9a una tarea estandar.
10
11Requiere que el repo tenga tambien un requirements.txt con, como minimo:
12    diffusers>=0.31.0
13    transformers
14    accelerate
15    imageio[ffmpeg]
16    torch
17
18Nota: la clase real del pipeline de Wan (WanImageToVideoPipeline o el
19nombre que tenga en tu version de diffusers) puede variar segun cuando
20la instales. Comprueba el nombre exacto en la documentacion de diffusers
21antes de desplegar, porque el soporte de Wan es reciente y puede cambiar.
22"""
23
24import base64
25import io
26import tempfile
27
28import ftfy
29import torch
30from PIL import Image
31from diffusers import WanImageToVideoPipeline
32from diffusers.utils import export_to_video
33
34# Bug conocido de transformers: el tokenizer de CLIP importa ftfy de forma
35# local dentro de __init__, pero luego lo referencia como nombre global en
36# otros metodos (p.ej. bpe()), lo que provoca NameError: name 'ftfy' is not
37# defined aunque el paquete este instalado. Se corrige inyectando el modulo
38# como atributo global antes de que se cargue el pipeline.
39import transformers.models.clip.tokenization_clip as _clip_tokenization_module
40_clip_tokenization_module.ftfy = ftfy
41
42
43class EndpointHandler:
44    """Handler que Hugging Face Inference Endpoints instancia una vez al arrancar."""
45
46    def __init__(self, path: str = ""):
47        """
48        Carga el pipeline en GPU al arrancar el contenedor.
49
50        path es la ruta local donde el Endpoint descarga los pesos del repo.
51        """
52        self.pipe = WanImageToVideoPipeline.from_pretrained(
53            path,
54            torch_dtype=torch.bfloat16,
55        )
56        self.pipe.to("cuda")
57
58        # El VAE y el image_encoder suelen necesitar float32 por estabilidad
59        # numerica en estos pipelines de imagen-a-video; forzarlos aqui evita
60        # el mismatch de dtype con el resto del pipeline en bfloat16.
61        if getattr(self.pipe, "vae", None) is not None:
62            self.pipe.vae = self.pipe.vae.to(torch.float32)
63        if getattr(self.pipe, "image_encoder", None) is not None:
64            self.pipe.image_encoder = self.pipe.image_encoder.to(torch.float32)
65
66    def __call__(self, data: dict) -> dict:
67        """
68        Procesa una peticion de generacion de video.
69
70        data llega con la forma {"inputs": {...}} tal y como lo envia
71        el cliente REST. Devuelve un dict serializable en JSON con el
72        video codificado en base64.
73        """
74        inputs = data.get("inputs", data)
75
76        imagen_b64 = inputs["image"]
77        prompt = inputs.get("prompt", "")
78        negative_prompt = inputs.get("negative_prompt", "")
79        num_frames = inputs.get("num_frames", 33)
80        num_inference_steps = inputs.get("num_inference_steps", 20)
81        guidance_scale = inputs.get("guidance_scale", 5.0)
82        seed = inputs.get("seed", 42)
83        height = inputs.get("height", 480)
84        width = inputs.get("width", 832)
85        fps = inputs.get("fps", 16)
86
87        imagen = decodificar_imagen(imagen_b64)
88        generador = torch.Generator(device="cuda").manual_seed(seed)
89
90        salida = self.pipe(
91            image=imagen,
92            prompt=prompt,
93            negative_prompt=negative_prompt,
94            num_frames=num_frames,
95            num_inference_steps=num_inference_steps,
96            guidance_scale=guidance_scale,
97            height=height,
98            width=width,
99            generator=generador,
100        )
101
102        frames = salida.frames[0]
103        video_b64 = exportar_video_base64(frames, fps)
104
105        return {"video_base64": video_b64}
106
107
108def decodificar_imagen(imagen_b64: str) -> Image.Image:
109    """Convierte una imagen codificada en base64 a un objeto PIL en modo RGB."""
110    imagen_bytes = base64.b64decode(imagen_b64)
111    return Image.open(io.BytesIO(imagen_bytes)).convert("RGB")
112
113
114def exportar_video_base64(frames, fps: int) -> str:
115    """Exporta una lista de frames a mp4 temporal y devuelve el contenido en base64."""
116    with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp:
117        export_to_video(frames, tmp.name, fps=fps)
118        tmp.seek(0)
119        video_bytes = tmp.read()
120
121    return base64.b64encode(video_bytes).decode("utf-8")