"""InfiniteTalk on ZeroGPU: portrait image + speech audio -> lip-synced talking-head video.

Audio-driven, so the spoken language does not matter (Russian, Kazakh, ... all work).
Uses the fp8-quantized T5 and a quantized DiT so the whole pipeline fits a 48 GB ZeroGPU slot
without ever materializing the 65 GB fp32 base checkpoint.
"""

import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

import spaces  # noqa: E402  must precede torch

import logging
import math
import tempfile
import time
from pathlib import Path
from types import SimpleNamespace

import gradio as gr
import librosa
import numpy as np
import pyloudnorm as pyln
import soundfile as sf
import torch
from einops import rearrange
from huggingface_hub import hf_hub_download, snapshot_download
from transformers import Wav2Vec2FeatureExtractor

import wan
from src.audio_analysis.wav2vec2 import Wav2Vec2Model
from wan.configs import WAN_CONFIGS
from wan.utils.multitalk_utils import save_video_ffmpeg

logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(levelname)s: %(message)s")
log = logging.getLogger("infinitetalk-space")

# --- configuration -------------------------------------------------------------------------
# DiT variant from MeiGen-AI/InfiniteTalk/quant_models. "single_int8_lora" ships with an
# acceleration LoRA fused in (8 steps, cfg 1.0/2.0, shift 2); "single_fp8" is the plain
# model (40 steps, cfg 5.0/4.0, shift 7). Switch via the Space's env settings.
DIT_VARIANT = os.environ.get("DIT_VARIANT", "single_int8_lora")
FUSED_LORA = DIT_VARIANT.endswith("_lora")
MAX_AUDIO_SEC = float(os.environ.get("MAX_AUDIO_SEC", "20"))
FRAME_NUM, MOTION_FRAME, FPS = 81, 9, 25

DEFAULTS = (
    dict(steps=8, text_cfg=1.0, audio_cfg=2.0, shift=2.0)
    if FUSED_LORA
    else dict(steps=40, text_cfg=5.0, audio_cfg=4.0, shift=7.0)
)

WEIGHTS = Path("weights")

# --- weights ---------------------------------------------------------------------------------
t0 = time.perf_counter()
base_dir = snapshot_download(
    "Wan-AI/Wan2.1-I2V-14B-480P",
    local_dir=WEIGHTS / "Wan2.1-I2V-14B-480P",
    allow_patterns=[
        "config.json",
        "Wan2.1_VAE.pth",
        "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth",
        "google/**",
        "xlm-roberta-large/**",
    ],
)
it_dir = snapshot_download(
    "MeiGen-AI/InfiniteTalk",
    local_dir=WEIGHTS / "InfiniteTalk",
    allow_patterns=[
        f"quant_models/infinitetalk_{DIT_VARIANT}.safetensors",
        f"quant_models/infinitetalk_{DIT_VARIANT}.json",
        "quant_models/t5_fp8.safetensors",
        "quant_models/t5_map_fp8.json",
    ],
)
wav2vec_dir = snapshot_download(
    "TencentGameMate/chinese-wav2vec2-base", local_dir=WEIGHTS / "chinese-wav2vec2-base"
)
hf_hub_download(
    "TencentGameMate/chinese-wav2vec2-base",
    "model.safetensors",
    revision="refs/pr/1",
    local_dir=WEIGHTS / "chinese-wav2vec2-base",
)
log.info("weights ready in %.0fs", time.perf_counter() - t0)

# --- models (module scope, moved to cuda eagerly so ZeroGPU can pack them) -----------------
t0 = time.perf_counter()
feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(wav2vec_dir, local_files_only=True)
audio_encoder = Wav2Vec2Model.from_pretrained(wav2vec_dir, local_files_only=True).eval()
audio_encoder.feature_extractor._freeze_parameters()

cfg = WAN_CONFIGS["infinitetalk-14B"]
pipe = wan.InfiniteTalkPipeline(
    config=cfg,
    checkpoint_dir=base_dir,
    quant_dir=os.path.join(it_dir, "quant_models", f"infinitetalk_{DIT_VARIANT}.safetensors"),
    device_id=0,
    rank=0,
    t5_cpu=False,
    quant="fp8",  # selects t5_fp8; the DiT file above carries its own quantization map
    infinitetalk_dir=None,
)
pipe.model.to("cuda")
pipe.text_encoder.model.to("cuda")
log.info("models built in %.0fs", time.perf_counter() - t0)

# Only the flags generate_infinitetalk() reads; TeaCache/APG stay off.
EXTRA = SimpleNamespace(
    use_teacache=False, teacache_thresh=0.2, size="infinitetalk-480",
    use_apg=False, apg_momentum=-0.75, apg_norm_threshold=55,
)

# --- helpers ---------------------------------------------------------------------------------
def prepare_audio(path: str, sr: int = 16000) -> np.ndarray:
    """Load any audio/video file as 16 kHz mono, loudness-normalized, trimmed to MAX_AUDIO_SEC."""
    audio, _ = librosa.load(path, sr=sr)
    audio = audio[: int(MAX_AUDIO_SEC * sr)]
    meter = pyln.Meter(sr)
    loudness = meter.integrated_loudness(audio)
    if abs(loudness) < 100:
        audio = pyln.normalize.loudness(audio, loudness, -23)
    return audio


def audio_embedding(audio: np.ndarray, sr: int = 16000) -> torch.Tensor:
    """wav2vec2 hidden states resampled to 25 fps, as the DiT expects (S, layers, 768)."""
    frames = int(len(audio) / sr * FPS)
    feats = np.squeeze(feature_extractor(audio, sampling_rate=sr).input_values)
    feats = torch.from_numpy(feats).float().unsqueeze(0)
    with torch.no_grad():
        out = audio_encoder(feats, seq_len=frames, output_hidden_states=True)
    emb = torch.stack(out.hidden_states[1:], dim=1).squeeze(0)
    return rearrange(emb, "b s d -> s b d").cpu()


def n_chunks(audio_sec: float) -> int:
    frames = int(audio_sec * FPS)
    return 1 + max(0, math.ceil((frames - FRAME_NUM) / (FRAME_NUM - MOTION_FRAME)))


def estimate_duration(image, audio, steps, *args, **kwargs) -> int:
    try:
        sec = min(librosa.get_duration(path=audio), MAX_AUDIO_SEC)
    except Exception:
        sec = MAX_AUDIO_SEC
    per_chunk = 20 + int(steps) * 4  # seconds; refined from measured runs
    return min(600, 45 + n_chunks(sec) * per_chunk)


# --- inference -------------------------------------------------------------------------------
@spaces.GPU(duration=estimate_duration)
def generate(
    image: str,
    audio: str,
    steps: int = DEFAULTS["steps"],
    audio_cfg: float = DEFAULTS["audio_cfg"],
    text_cfg: float = DEFAULTS["text_cfg"],
    prompt: str = "A person is talking to the camera.",
    seed: int = -1,
    progress=gr.Progress(track_tqdm=True),
) -> str:
    """Animate a portrait so it speaks the given audio (any language) with synced lips.

    Args:
        image: Path to a front-facing portrait (PNG/JPG).
        audio: Path to speech audio (wav/mp3/…); trimmed to MAX_AUDIO_SEC.
        steps: Diffusion steps per chunk.
        audio_cfg: Audio guidance scale (higher = tighter lip sync).
        text_cfg: Text guidance scale.
        prompt: Short scene description.
        seed: RNG seed, -1 for random.
    Returns:
        Path to the generated mp4 with the audio muxed in.
    """
    if image is None or audio is None:
        raise gr.Error("Please provide both a portrait image and an audio file.")

    t_start = time.perf_counter()
    work = Path(tempfile.mkdtemp(prefix="infinitetalk_"))
    speech = prepare_audio(audio)
    wav_path = work / "speech.wav"
    sf.write(wav_path, speech, 16000)
    emb_path = work / "speech.pt"
    torch.save(audio_embedding(speech), emb_path)

    if seed is None or int(seed) < 0:
        seed = int(torch.randint(0, 2**31 - 1, (1,)).item())

    input_data = {
        "prompt": prompt or "",
        "cond_video": image,
        "cond_audio": {"person1": str(emb_path)},
        "video_audio": str(wav_path),
    }
    video = pipe.generate_infinitetalk(
        input_data,
        size_buckget="infinitetalk-480",
        motion_frame=MOTION_FRAME,
        frame_num=FRAME_NUM,
        shift=DEFAULTS["shift"],
        sampling_steps=int(steps),
        text_guide_scale=float(text_cfg),
        audio_guide_scale=float(audio_cfg),
        seed=int(seed),
        offload_model=False,
        max_frames_num=1000,
        color_correction_strength=1.0,
        extra_args=EXTRA,
    )
    out_stem = str(work / "result")
    save_video_ffmpeg(video, out_stem, [str(wav_path)], high_quality_save=False)
    log.info(
        "generated %.1fs of video (%d chunks, %d steps) in %.0fs",
        len(speech) / 16000, n_chunks(len(speech) / 16000), int(steps), time.perf_counter() - t_start,
    )
    return out_stem + ".mp4"


# --- UI --------------------------------------------------------------------------------------
with gr.Blocks(title="InfiniteTalk — talking-head from image + audio") as demo:
    gr.Markdown(
        f"""
        # 🗣️ InfiniteTalk — image + audio → lip-synced video
        Upload a portrait and speech audio in **any language** (the model is audio-driven, so
        Russian, Kazakh, English… all work). Audio is trimmed to **{MAX_AUDIO_SEC:.0f} s**; output is 480p @ 25 fps.
        Model: [MeiGen-AI/InfiniteTalk](https://huggingface.co/MeiGen-AI/InfiniteTalk) · variant `{DIT_VARIANT}`.
        """
    )
    with gr.Row():
        with gr.Column():
            image_in = gr.Image(type="filepath", label="Portrait image")
            audio_in = gr.Audio(type="filepath", label="Speech audio")
            with gr.Accordion("Advanced", open=False):
                steps_in = gr.Slider(4, 40, value=DEFAULTS["steps"], step=1, label="Diffusion steps")
                audio_cfg_in = gr.Slider(1.0, 6.0, value=DEFAULTS["audio_cfg"], step=0.5, label="Audio guidance")
                text_cfg_in = gr.Slider(1.0, 6.0, value=DEFAULTS["text_cfg"], step=0.5, label="Text guidance")
                prompt_in = gr.Textbox(value="A person is talking to the camera.", label="Prompt")
                seed_in = gr.Number(value=-1, precision=0, label="Seed (-1 = random)")
            run_btn = gr.Button("Generate", variant="primary")
        with gr.Column():
            video_out = gr.Video(label="Result")

    run_btn.click(
        fn=generate,
        inputs=[image_in, audio_in, steps_in, audio_cfg_in, text_cfg_in, prompt_in, seed_in],
        outputs=video_out,
        api_name="generate",
    )
    gr.Examples(
        examples=[["examples/ref_image.png", "examples/1.wav"]],
        inputs=[image_in, audio_in],
        outputs=video_out,
        fn=generate,
        cache_examples=True,
        cache_mode="lazy",
    )

demo.launch(mcp_server=True)
