# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import argparse
import re
import time
from typing import Any

import numpy as np
from PIL import Image
from vllm.multimodal.media.audio import load_audio

from vllm_omni.diffusion.data import DiffusionParallelConfig
from vllm_omni.diffusion.utils.media_utils import mux_video_audio_bytes
from vllm_omni.entrypoints.omni import Omni
from vllm_omni.inputs.data import OmniDiffusionSamplingParams

IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".gif"}
AUDIO_EXTENSIONS = {".wav", ".mp3", ".flac", ".m4a", ".aac", ".ogg"}


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Offline inference for DreamID-Omni (video + audio).")
    parser.add_argument("--model", required=True, help="DreamID ckpt root directory.")
    parser.add_argument("--model-type", default="dreamid-omni", help="Model type.")
    parser.add_argument("--prompt", default=None, help="Text prompt.")

    parser.add_argument("--image-path", type=str, nargs="+", help="list of image-path")
    parser.add_argument("--audio-path", type=str, nargs="+", help="list of audio-path")
    parser.add_argument("--prompt-file", type=str, default=None, help="Text prompt in json format.")

    parser.add_argument("--height", type=int, default=704, help="Video height.")
    parser.add_argument("--width", type=int, default=1280, help="Video width.")
    parser.add_argument("--num-inference-steps", type=int, default=45, help="Sampling steps.")
    parser.add_argument("--solver-name", default="unipc", help="Solver name: unipc|dpm++|euler.")
    parser.add_argument("--shift", type=float, default=5.0, help="Scheduler shift.")
    parser.add_argument("--seed", type=int, default=103, help="Random seed for reproducible generation.")
    parser.add_argument(
        "--cfg-parallel-size",
        type=int,
        default=1,
        choices=[1, 2, 3, 4],
        help="Number of GPUs used for classifier free guidance parallel size (max 4 branches).",
    )
    parser.add_argument(
        "--tensor-parallel-size",
        type=int,
        default=1,
        help=("Number of GPUs used for tensor parallelism (TP) inside the DiT."),
    )
    parser.add_argument(
        "--video-negative-prompt",
        default="jitter, bad hands, blur, distortion",
        help="Negative prompt for video.",
    )
    parser.add_argument(
        "--audio-negative-prompt",
        default="robotic, muffled, echo, distorted",
        help="Negative prompt for audio.",
    )
    parser.add_argument("--output", default="dreamid_output.mp4", help="Output video path.")
    parser.add_argument(
        "--quantization",
        type=str,
        default=None,
        choices=["fp8", "int8"],
        help=("Online (dynamic) quantization method for the DreamID-Omni transformer. "),
    )
    parser.add_argument(
        "--enable-cpu-offload",
        action="store_true",
        default=False,
        help="Enable CPU offloading for diffusion models.",
    )
    parser.add_argument(
        "--enable-layerwise-offload",
        action="store_true",
        help="Enable layerwise (blockwise) offloading on DiT modules.",
    )
    parser.add_argument("--cache-backend", type=str, default=None, choices=["cache_dit"], help="Cache backend.")
    parser.add_argument(
        "--use-hsdp",
        action="store_true",
        help="Enable HSDP for DreamID-Omni fused transformer blocks.",
    )
    parser.add_argument(
        "--hsdp-shard-size",
        type=int,
        default=1,
        help="Number of GPUs used for HSDP sharding when HSDP is enabled.",
    )
    parser.add_argument(
        "--hsdp-replicate-size",
        type=int,
        default=1,
        help="Number of HSDP replica groups. Default 1 means pure sharding.",
    )
    return parser.parse_args()


def load_image_and_audio(image_paths, audio_paths):
    image = []
    audio = []

    for path in image_paths:
        with Image.open(path) as img:
            img = img.convert("RGB")
            image.append(img)

    for path in audio_paths:
        audio_array, sr = load_audio(path, sr=16000)
        audio_array = audio_array[int(sr * 1) : int(sr * 3)]
        audio.append(audio_array)
    return image, audio


def _extract_peak_memory_mb(result: Any) -> float:
    """Pull worker-reported peak VRAM (MiB) generation result.

    Mirrors vllm_omni/entrypoints/openai/serving_video.py:_extract_peak_memory_mb.
    """
    if isinstance(result, list):
        result = result[0] if result else None
    if result is None:
        return 0.0
    val = getattr(result, "peak_memory_mb", 0.0)
    if not val:
        inner = getattr(result, "request_output", None)
        if isinstance(inner, list):
            inner = inner[0] if inner else None
        val = getattr(inner, "peak_memory_mb", 0.0)
    try:
        return float(val or 0.0)
    except (TypeError, ValueError):
        return 0.0


def main() -> None:
    args = parse_args()
    if args.prompt is None and args.prompt_file is None:
        raise ValueError("Either --prompt or --prompt-file must be provided.")

    text_prompt = args.prompt
    if args.prompt_file:
        import json

        with open(args.prompt_file) as f:
            text_prompt = json.load(f)
            text_prompt = re.sub(
                r"\[SPEAKER_TIMESTAMPS_START\].*?\[SPEAKER_TIMESTAMPS_END\]", "", text_prompt, flags=re.DOTALL
            ).strip()
            text_prompt = re.sub(
                r"\[AUDIO_DESCRIPTION_START].*?\[AUDIO_DESCRIPTION_END]", "", text_prompt, flags=re.DOTALL
            ).strip()
            text_prompt = re.sub(r"\[[A-Z_]+\]", "", text_prompt)
            text_prompt = re.sub(r"\n\s*\n", "\n", text_prompt).strip()

    image, audio = load_image_and_audio(args.image_path, args.audio_path)

    prompt = {
        "prompt": text_prompt,
        "video_negative_prompt": args.video_negative_prompt,
        "audio_negative_prompt": args.audio_negative_prompt,
        "multi_modal_data": {"image": image, "audio": audio},
    }

    sampling_params = OmniDiffusionSamplingParams(
        height=args.height,
        width=args.width,
        num_inference_steps=args.num_inference_steps,
        seed=args.seed,
        extra_args={
            "solver_name": args.solver_name,
            "shift": args.shift,
        },
    )

    parallel_config = DiffusionParallelConfig(
        cfg_parallel_size=args.cfg_parallel_size,
        tensor_parallel_size=args.tensor_parallel_size,
        use_hsdp=args.use_hsdp,
        hsdp_shard_size=args.hsdp_shard_size,
        hsdp_replicate_size=args.hsdp_replicate_size,
    )

    cache_config = None
    if args.cache_backend == "cache_dit":
        cache_config = {
            "Fn_compute_blocks": 1,
            "Bn_compute_blocks": 0,
            "max_warmup_steps": 4,
            "max_cached_steps": 20,
            "residual_diff_threshold": 0.24,
            "max_continuous_cached_steps": 3,
            "enable_taylorseer": False,
            "taylorseer_order": 1,
            "scm_steps_mask_policy": None,
            "scm_steps_policy": "dynamic",
        }

    omni_kwargs = dict(
        model=args.model,
        parallel_config=parallel_config,
        model_type=args.model_type,
        enable_cpu_offload=args.enable_cpu_offload,
        enable_layerwise_offload=args.enable_layerwise_offload,
        cache_backend=args.cache_backend,
        cache_config=cache_config,
    )
    if args.quantization is not None:
        omni_kwargs["quantization"] = args.quantization
    omni = Omni(**omni_kwargs)
    start = time.perf_counter()
    outputs = omni.generate(prompt, sampling_params)
    elapsed = time.perf_counter() - start

    if not outputs:
        raise RuntimeError("No output returned from DreamID-Omni.")
    result = outputs[0]
    if not result.images:
        raise RuntimeError("No video frames found in DreamID-Omni output.")
    generated_video = result.images[0]
    mm = result.multimodal_output or {}
    generated_audio = mm.get("audio")
    fps = int(mm.get("fps", 24))
    sample_rate = int(mm.get("audio_sample_rate", 16000))

    # DreamID-Omni returns video as (C, F, H, W) float32 in [-1, 1].
    # mux_video_audio_bytes expects (F, H, W, C) uint8.
    if not isinstance(generated_video, np.ndarray) or generated_video.ndim != 4:
        raise RuntimeError(f"Unexpected video shape: {getattr(generated_video, 'shape', None)}")
    frames = generated_video.transpose(1, 2, 3, 0)
    frames = (np.clip((frames + 1.0) / 2.0, 0.0, 1.0) * 255.0).round().astype(np.uint8)

    audio_np = None
    if generated_audio is not None:
        audio_np = np.squeeze(np.asarray(generated_audio)).astype(np.float32)

    output_path = args.output
    video_bytes = mux_video_audio_bytes(
        frames,
        audio_np,
        fps=float(fps),
        audio_sample_rate=sample_rate,
    )
    with open(output_path, "wb") as f:
        f.write(video_bytes)
    print(f"Saved generated video to {output_path}")
    print(f"Total time: {elapsed:.2f}s")
    peak_mb = _extract_peak_memory_mb(outputs)
    if peak_mb:
        print(f"Worker peak GPU memory (reserved): {peak_mb:.2f} MiB ({peak_mb / 1024:.2f} GiB)")


if __name__ == "__main__":
    main()
