#!/usr/bin/env python3
"""
Distributed semantic encoding pipeline for audio dataset.

Encodes audio files to semantic codes using MERT-25 model with:
- Multi-node, multi-GPU support via torch.distributed
- Efficient batching with padding
- Resume support (skip already encoded files)
- Direct mono 24kHz audio loading
"""

import argparse
import json
import os
import time
from typing import Dict, List, Optional, Tuple

import numpy as np
import torch
import torch.distributed as dist
import torchaudio

from suno_utils.audio import Audio

try:
    from tqdm import tqdm
except ImportError:

    def tqdm(iterable, *args, **kwargs):
        return iterable


def setup_distributed() -> Tuple[int, int, int]:
    """
    Initialize distributed process group for SLURM.

    Returns:
        Tuple of (rank, world_size, local_rank)
    """
    # SLURM environment variables
    rank = int(os.environ.get("SLURM_PROCID", 0))
    world_size = int(os.environ.get("SLURM_NTASKS", 1))
    local_rank = int(os.environ.get("SLURM_LOCALID", 0))

    # Initialize process group
    if world_size > 1:
        dist.init_process_group(
            backend="nccl", init_method="env://", rank=rank, world_size=world_size
        )

    # Set device for this process
    torch.cuda.set_device(local_rank)

    return rank, world_size, local_rank


def load_metadata(metadata_path: str) -> List[Dict]:
    """Load metadata from JSONL file."""
    metadata = []
    with open(metadata_path, "r") as f:
        for line in f:
            line = line.strip()
            if line:
                try:
                    metadata.append(json.loads(line))
                except json.JSONDecodeError:
                    continue
    return metadata


def load_audio_list(audio_list_path: str) -> List[Dict]:
    """
    Load list of audio file paths and convert to metadata format.

    Supports multiple formats:
    - Plain text file with one path per line
    - JSON file with list of paths
    - JSON file with list of dicts containing 'path' or 'filepath' keys

    Args:
        audio_list_path: Path to file containing audio paths

    Returns:
        List of metadata dicts with 'id' and 'local_filepath' fields
    """
    metadata = []

    # Try loading as JSON first
    try:
        with open(audio_list_path, "r") as f:
            # Try to load as single JSON object first
            f.seek(0)
            first_line = f.readline().strip()
            f.seek(0)
            
            # Check if it's JSONL format (one JSON per line)
            if first_line and first_line.startswith("{"):
                # JSONL format
                for line in f:
                    line = line.strip()
                    if not line:
                        continue
                    item = json.loads(line)
                    if isinstance(item, dict):
                        audio_path = (
                            item.get("audio_path")
                            or item.get("path")
                            or item.get("filepath")
                            or item.get("local_filepath")
                        )
                        audio_id = (
                            item.get("id")
                            or os.path.splitext(os.path.basename(audio_path))[0]
                        )
                        metadata.append({"id": audio_id, "local_filepath": audio_path})
                return metadata
            else:
                # Regular JSON file
                data = json.load(f)

        if isinstance(data, list):
            for item in data:
                if isinstance(item, str):
                    # Simple list of paths
                    audio_path = item
                    audio_id = os.path.splitext(os.path.basename(audio_path))[0]
                    metadata.append({"id": audio_id, "local_filepath": audio_path})
                elif isinstance(item, dict):
                    # List of dicts - extract path and id
                    audio_path = (
                        item.get("audio_path")
                        or item.get("path")
                        or item.get("filepath")
                        or item.get("local_filepath")
                    )
                    audio_id = (
                        item.get("id")
                        or os.path.splitext(os.path.basename(audio_path))[0]
                    )
                    metadata.append({"id": audio_id, "local_filepath": audio_path})
        return metadata
    except (json.JSONDecodeError, ValueError):
        pass

    # Fall back to plain text file (one path per line)
    with open(audio_list_path, "r") as f:
        for line in f:
            line = line.strip()
            if line and not line.startswith("#"):  # Skip comments
                audio_path = line
                audio_id = os.path.splitext(os.path.basename(audio_path))[0]
                metadata.append({"id": audio_id, "local_filepath": audio_path})

    return metadata


def load_audio_mono_24k(filepath: str) -> Optional[Audio]:
    """
    Load audio file directly as mono 24kHz for semantic encoding.

    Args:
        filepath: Path to audio file

    Returns:
        Mono 24kHz Audio object or None if loading fails
    """
    try:
        # Load with torchaudio
        waveform, sample_rate = torchaudio.load(filepath)

        # Convert to mono by averaging channels
        if waveform.shape[0] > 1:
            waveform = waveform.mean(dim=0, keepdim=True)

        # Resample to 24kHz if needed
        if sample_rate != 24000:
            waveform = torchaudio.functional.resample(waveform, sample_rate, 24000)

        # Convert to Audio object (mono, 24kHz)
        audio = Audio.from_array_float(
            waveform.numpy().squeeze(), sample_rate=24000, max_allowed_val=100
        )
        return audio
    except Exception as e:
        return None


def encode_audio_batch(
    audio_list: List[Audio],
    device: str,
    encode_semantic_fn,
    n_codebooks: Optional[int] = None,
) -> List[np.ndarray]:
    """
    Encode a batch of audio files to semantic codes.

    Args:
        audio_list: List of Audio objects (mono 24kHz)
        device: Device to use for encoding
        encode_semantic_fn: Encoding function from suno_utils
        n_codebooks: Number of codebooks to use (None = use all available)

    Returns:
        List of semantic code arrays
    """
    if len(audio_list) == 0:
        return []

    # Encode each audio individually (encode_semantic expects single audio)
    codes_list = []
    for audio in audio_list:
        try:
            codes = encode_semantic_fn(audio, device=device, n_codebooks=n_codebooks)
            if isinstance(codes, torch.Tensor):
                codes = codes.cpu().numpy()
            codes_list.append(codes.astype(np.int64))
        except Exception as e:
            # If encoding fails, append None and handle in caller
            codes_list.append(None)

    return codes_list


def process_worker(
    rank: int,
    world_size: int,
    local_rank: int,
    metadata: List[Dict],
    output_dir: str,
    batch_size: int,
    overwrite: bool,
    encode_semantic_fn,
    preload_semantic_models_fn,
    n_codebooks: Optional[int] = None,
) -> Dict[str, int]:
    """
    Process assigned subset of dataset on this worker.

    Args:
        rank: Global rank of this worker
        world_size: Total number of workers
        local_rank: Local rank on this node
        metadata: Full metadata list
        output_dir: Output directory for .npz files
        batch_size: Batch size for encoding
        overwrite: Whether to overwrite existing files
        encode_semantic_fn: Encoding function
        preload_semantic_models_fn: Model loading function
        n_codebooks: Number of codebooks to use (None = use all available)

    Returns:
        Dictionary with processing statistics
    """
    device = f"cuda:{local_rank}"

    # Load model on this GPU
    if rank == 0:
        print(f"Loading semantic models on {world_size} workers...")

    # IMPORTANT: Clear any cached models first to ensure we load with correct centroids
    # (suno_utils caches models by device only, not by centroids filepath)
    from suno_utils.tasks.mert_25 import clean_models

    clean_models()

    preload_semantic_models_fn(
        checkpoint_filepath="/app/suno/data/dpo/models/mert_25.pt",
        # this is the 2 codebook model
        # centroids_filepath="/app/suno/data/dpo/models/mert_25_2x4k.npy",
        # this is the 768d model, 64 codebooks
        centroids_filepath="/home/minz/temp/mert_768d_centroids_4000_50.npy",
        device=device,
    )

    # Create output directory
    os.makedirs(output_dir, exist_ok=True)
    error_log_path = os.path.join(output_dir, f"errors_rank{rank}.log")

    # Statistics
    stats = {
        "processed": 0,
        "skipped": 0,
        "errors": 0,
    }

    # Get subset for this worker
    worker_metadata = [metadata[i] for i in range(rank, len(metadata), world_size)]

    if rank == 0:
        print(f"Rank {rank}: Processing {len(worker_metadata)} samples")

    # Process in batches
    batch_audio = []
    batch_ids = []
    batch_filepaths = []

    start_time = time.time()
    processed_count = 0

    for idx, meta in enumerate(worker_metadata):
        audio_id = meta.get("id")
        local_filepath = meta.get("local_filepath")

        if not audio_id or not local_filepath:
            stats["errors"] += 1
            continue

        output_path = os.path.join(output_dir, f"{audio_id}.npz")

        # Check if already exists (resume support)
        if not overwrite and os.path.exists(output_path):
            stats["skipped"] += 1
            continue

        # Load audio as mono 24kHz
        audio_mono = load_audio_mono_24k(local_filepath)
        if audio_mono is None:
            stats["errors"] += 1
            with open(error_log_path, "a") as f:
                f.write(f"Failed to load: {audio_id} - {local_filepath}\n")
            continue

        # Add to batch
        batch_audio.append(audio_mono)
        batch_ids.append(audio_id)
        batch_filepaths.append(local_filepath)

        # Process batch when full or at end
        if len(batch_audio) >= batch_size or idx == len(worker_metadata) - 1:
            # Encode batch
            try:
                codes_list = encode_audio_batch(
                    batch_audio,
                    device,
                    encode_semantic_fn,
                    n_codebooks=n_codebooks,
                )

                # Save individual files
                for bid, codes in zip(batch_ids, codes_list):
                    if codes is not None:
                        output_path = os.path.join(output_dir, f"{bid}.npz")
                        np.savez(output_path, codes=codes)
                        stats["processed"] += 1
                        processed_count += 1
                    else:
                        stats["errors"] += 1
                        with open(error_log_path, "a") as f:
                            f.write(f"Failed to encode: {bid}\n")

                # Log progress every 100 samples
                if processed_count % 100 == 0 and processed_count > 0:
                    elapsed = time.time() - start_time
                    rate = processed_count / elapsed if elapsed > 0 else 0
                    if rank == 0:
                        print(
                            f"Rank {rank}: Processed {processed_count}/{len(worker_metadata)} "
                            f"({rate:.2f} samples/sec)"
                        )

            except Exception as e:
                stats["errors"] += len(batch_ids)
                with open(error_log_path, "a") as f:
                    f.write(f"Batch encoding failed: {batch_ids}\n")
                    f.write(f"Error: {str(e)}\n")

            # Clear batch
            batch_audio = []
            batch_ids = []
            batch_filepaths = []

    # Final statistics
    elapsed = time.time() - start_time
    rate = stats["processed"] / elapsed if elapsed > 0 else 0

    if rank == 0 or stats["processed"] > 0:
        print(f"\nRank {rank} Summary:")
        print(f"  Processed: {stats['processed']}")
        print(f"  Skipped:   {stats['skipped']}")
        print(f"  Errors:    {stats['errors']}")
        print(f"  Rate:      {rate:.2f} samples/sec")
        print(f"  Time:      {elapsed:.1f}s")

    return stats


def main():
    """Main entry point for semantic encoding."""
    parser = argparse.ArgumentParser(
        description="Distributed semantic encoding pipeline"
    )
    parser.add_argument(
        "--metadata_path",
        type=str,
        default=None,
        help="Path to JSONL metadata file (format: one JSON object per line with 'id' and 'local_filepath' fields)",
    )
    parser.add_argument(
        "--audio_list",
        type=str,
        default=None,
        help="Path to audio list file (plalsin text with one path per line, or JSON list of paths)",
    )
    parser.add_argument(
        "--output_dir",
        type=str,
        default="/app2/suno/data/semantic_code/sft",
        help="Output directory for .npz files",
    )
    parser.add_argument(
        "--batch_size", type=int, default=16, help="Batch size for encoding"
    )
    parser.add_argument(
        "--overwrite", action="store_true", help="Overwrite existing files"
    )
    parser.add_argument(
        "--n_codebooks",
        type=int,
        default=None,
        help="Number of codebooks to use (default: use all available from centroids file)",
    )

    args = parser.parse_args()

    # Setup distributed processing first
    rank, world_size, local_rank = setup_distributed()

    # Validate input arguments
    if args.metadata_path is None and args.audio_list is None:
        # Set default if neither provided
        args.metadata_path = (
            "/home/tony/Data/Preference/RealGen/metas_v5_val_filtered_clean.jsonl"
        )

    if args.metadata_path and args.audio_list:
        if rank == 0:
            print("Error: Cannot specify both --metadata_path and --audio_list")
        return

    if rank == 0:
        print("=" * 80)
        print("SEMANTIC ENCODING PIPELINE")
        print("=" * 80)
        if args.metadata_path:
            print(f"Input mode:  JSONL metadata")
            print(f"Metadata:    {args.metadata_path}")
        else:
            print(f"Input mode:  Audio list")
            print(f"Audio list:  {args.audio_list}")
        print(f"Output dir:  {args.output_dir}")
        print(f"Batch size:  {args.batch_size}")
        print(f"Overwrite:   {args.overwrite}")
        print(
            f"N codebooks: {args.n_codebooks if args.n_codebooks else 'all available'}"
        )
        print(f"World size:  {world_size}")
        print(f"Nodes:       {world_size // 8}")
        print(f"GPUs/node:   8")
        print("=" * 80)
        print()

    # Load metadata
    if rank == 0:
        print("Loading input data...")

    if args.metadata_path:
        metadata = load_metadata(args.metadata_path)
    else:
        metadata = load_audio_list(args.audio_list)

    if rank == 0:
        print(f"Loaded {len(metadata):,} metadata entries")

        # Count existing files
        if os.path.exists(args.output_dir):
            existing_count = len(
                [f for f in os.listdir(args.output_dir) if f.endswith(".npz")]
            )
            print(f"Found {existing_count:,} existing .npz files")
            if not args.overwrite and existing_count > 0:
                print("Will skip existing files (use --overwrite to re-encode)")
        print()

    # Import encoding functions
    try:
        from suno_utils.tasks.mert_25 import (
            preload_models as preload_semantic_models_,
            encode as encode_semantic,
        )
    except ImportError as e:
        if rank == 0:
            print(f"Error: Failed to import suno_utils: {e}")
            print("Make sure suno_utils is installed and accessible")
        return

    # Process worker's subset
    stats = process_worker(
        rank=rank,
        world_size=world_size,
        local_rank=local_rank,
        metadata=metadata,
        output_dir=args.output_dir,
        batch_size=args.batch_size,
        overwrite=args.overwrite,
        encode_semantic_fn=encode_semantic,
        preload_semantic_models_fn=preload_semantic_models_,
        n_codebooks=args.n_codebooks,
    )

    # Gather statistics from all workers
    if world_size > 1:
        # Synchronize
        dist.barrier()

        # Collect stats on rank 0
        if rank == 0:
            total_stats = {
                "processed": stats["processed"],
                "skipped": stats["skipped"],
                "errors": stats["errors"],
            }

            for i in range(1, world_size):
                # Simple gathering (could use dist.gather for more efficiency)
                pass

            print("\n" + "=" * 80)
            print("FINAL SUMMARY (Rank 0 only, see individual logs for full stats)")
            print("=" * 80)
            print(f"Total processed (rank 0): {total_stats['processed']:,}")
            print(f"Total skipped (rank 0):   {total_stats['skipped']:,}")
            print(f"Total errors (rank 0):    {total_stats['errors']:,}")
            print("=" * 80)
    else:
        print("\n" + "=" * 80)
        print("FINAL SUMMARY")
        print("=" * 80)
        print(f"Total processed: {stats['processed']:,}")
        print(f"Total skipped:   {stats['skipped']:,}")
        print(f"Total errors:    {stats['errors']:,}")
        print("=" * 80)

    # Cleanup
    if world_size > 1:
        dist.destroy_process_group()

    if rank == 0:
        print("\nEncoding complete!")


if __name__ == "__main__":
    main()
