"""HF Inference Endpoint handler for SulphurAI/Sulphur-2-base (LTX-2.3 fine-tune). The pipeline is constructed against the low-level `ltx_pipelines.DistilledPipeline` API rather than `diffusers`, because Sulphur ships raw .safetensors and explicitly recommends the distilled checkpoint + distill LoRA workflow. Required environment: - LTX-2 (Lightricks/LTX-2) installed: `ltx_core`, `ltx_pipelines` importable. The official install path is `uv sync` from a git clone — see Dockerfile. - HF_TOKEN set as an Endpoint secret. The Gemma text encoder is gated, so the HF account behind the token must have accepted its license. - GPU with >= 48 GB VRAM (H100 80GB / A100 80GB recommended). Distilled bf16 plus Gemma 12B comfortably exceeds 40 GB at fp16. """ from __future__ import annotations import base64 import logging import os import tempfile import uuid from typing import Any, Dict, Iterator import torch from huggingface_hub import hf_hub_download, snapshot_download logger = logging.getLogger("sulphur2.handler") logging.basicConfig(level=logging.INFO) SULPHUR_REPO = "SulphurAI/Sulphur-2-base" SULPHUR_DISTILLED = "sulphur_distil_bf16.safetensors" SULPHUR_DISTILL_LORA = "distill_loras/ltx-2.3-22b-distilled-lora-1.1_fro90_ceil72_condsafe.safetensors" LTX_REPO = "Lightricks/LTX-2.3" LTX_SPATIAL_UPSCALER = "ltx-2.3-spatial-upscaler-x2-1.1.safetensors" GEMMA_REPO = "google/gemma-3-12b-it-qat-q4_0-unquantized" def _import_ltx(): """Import LTX-2 lazily so import errors surface at handler init, not module load.""" from ltx_core.loader import LoraPathStrengthAndSDOps from ltx_pipelines import DistilledPipeline from ltx_pipelines.utils.media_io import encode_video return DistilledPipeline, LoraPathStrengthAndSDOps, encode_video class EndpointHandler: def __init__(self, model_dir: str, **kwargs: Any) -> None: token = os.getenv("HF_TOKEN") or os.getenv("HUGGING_FACE_HUB_TOKEN") if not token: raise RuntimeError( "HF_TOKEN (or HUGGING_FACE_HUB_TOKEN) is required. Accept the Gemma " "license on huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized " "and add the token as an Endpoint secret." ) DistilledPipeline, LoraPathStrengthAndSDOps, encode_video = _import_ltx() self._encode_video = encode_video logger.info("Downloading Sulphur distilled checkpoint...") checkpoint_path = hf_hub_download( repo_id=SULPHUR_REPO, filename=SULPHUR_DISTILLED, token=token ) logger.info("Downloading Sulphur distill LoRA...") distill_lora_path = hf_hub_download( repo_id=SULPHUR_REPO, filename=SULPHUR_DISTILL_LORA, token=token ) logger.info("Downloading LTX-2.3 spatial upscaler...") spatial_upscaler_path = hf_hub_download( repo_id=LTX_REPO, filename=LTX_SPATIAL_UPSCALER, token=token ) logger.info("Downloading Gemma 3 text encoder snapshot...") gemma_root = snapshot_download( repo_id=GEMMA_REPO, token=token, allow_patterns=[ "*.json", "*.model", "*.safetensors", "tokenizer*", "special_tokens_map.json", ], ) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") logger.info("Initializing DistilledPipeline on %s...", device) loras = [LoraPathStrengthAndSDOps(path=distill_lora_path, strength=1.0)] self.pipeline = DistilledPipeline( distilled_checkpoint_path=checkpoint_path, gemma_root=gemma_root, spatial_upsampler_path=spatial_upscaler_path, loras=loras, device=device, ) logger.info("Pipeline ready.") @torch.inference_mode() def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]: prompt = data.get("inputs") or data.get("prompt") if not isinstance(prompt, str) or not prompt.strip(): return {"error": "Request must include a non-empty 'inputs' or 'prompt' string."} params = data.get("parameters") or {} width = int(params.get("width", 768)) height = int(params.get("height", 512)) num_frames = int(params.get("num_frames", 97)) fps = int(params.get("fps", 24)) seed = int(params.get("seed", 0)) enhance_prompt = bool(params.get("enhance_prompt", False)) if num_frames < 9 or num_frames > 257: return {"error": "num_frames must be between 9 and 257."} if width % 32 or height % 32: return {"error": "width and height must be multiples of 32."} logger.info( "Generating: prompt=%r seed=%d %dx%d frames=%d fps=%d", prompt[:80], seed, width, height, num_frames, fps, ) video_iter, audio = self.pipeline( prompt=prompt, seed=seed, height=height, width=width, num_frames=num_frames, frame_rate=float(fps), images=[], tiling_config=None, enhance_prompt=enhance_prompt, ) chunks = list(video_iter) if isinstance(video_iter, Iterator) else [video_iter] if not chunks: return {"error": "Pipeline produced no frames."} video_chunks_number = len(chunks) tmp_path = os.path.join(tempfile.gettempdir(), f"sulphur_{uuid.uuid4().hex}.mp4") try: self._encode_video( video=iter(chunks), fps=fps, audio=audio, output_path=tmp_path, video_chunks_number=video_chunks_number, ) with open(tmp_path, "rb") as f: payload = f.read() finally: if os.path.exists(tmp_path): try: os.remove(tmp_path) except OSError: pass return { "content_type": "video/mp4", "encoding": "base64", "video_base64": base64.b64encode(payload).decode("ascii"), "metadata": { "width": width, "height": height, "num_frames": num_frames, "fps": fps, "seed": seed, }, }