#!/usr/bin/env python3
"""
Sample diverse subset using TF-IDF on tags.

Strategy:
1. Keep ALL non-English songs with lyrics
2. Sample English songs with text using TF-IDF scores (high diversity)
3. Sample no-text songs using TF-IDF scores (high diversity)
"""

import json
import math
import os
import random
import statistics
from collections import Counter
from typing import Dict, List, Tuple

try:
    import numpy as np

    HAS_NUMPY = True
except ImportError:
    HAS_NUMPY = False

    # Simple numpy replacements
    class np:
        @staticmethod
        def array(x):
            return list(x)

        @staticmethod
        def mean(x):
            return statistics.mean(x)

        @staticmethod
        def median(x):
            return statistics.median(x)

        @staticmethod
        def percentile(x, p):
            sorted_x = sorted(x)
            k = (len(sorted_x) - 1) * p / 100
            f = math.floor(k)
            c = math.ceil(k)
            if f == c:
                return sorted_x[int(k)]
            return sorted_x[int(f)] * (c - k) + sorted_x[int(c)] * (k - f)

        @staticmethod
        def random_choice(n, size, replace, p):
            # Weighted sampling without numpy
            indices = list(range(n))
            sampled = []
            for _ in range(size):
                r = random.random()
                cumsum = 0
                for i, prob in enumerate(p):
                    cumsum += prob
                    if r <= cumsum:
                        sampled.append(indices[i])
                        if not replace:
                            indices.pop(i)
                            p = list(p)
                            p.pop(i)
                            # Renormalize
                            total = sum(p)
                            p = [x / total for x in p]
                        break
            return sampled


try:
    from tqdm import tqdm
except ImportError:

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


def load_metadata(jsonl_path: str) -> List[Dict]:
    """Load metadata from JSONL file."""
    print(f"Loading metadata from: {jsonl_path}")
    metadata = []

    with open(jsonl_path, "r") as f:
        for line in tqdm(f, desc="Loading"):
            line = line.strip()
            if line:
                try:
                    metadata.append(json.loads(line))
                except json.JSONDecodeError:
                    continue

    print(f"Loaded {len(metadata):,} entries")
    return metadata


def calculate_tag_idf(metadata: List[Dict]) -> Dict[str, float]:
    """
    Calculate IDF (Inverse Document Frequency) for each tag.

    IDF = log(total_docs / docs_containing_tag)
    """
    print("\nCalculating IDF scores for tags...")

    tag_doc_counts = Counter()
    total_docs = len(metadata)

    # Count how many documents contain each tag
    for entry in tqdm(metadata, desc="Counting tags"):
        tags = entry.get("tags", [])
        unique_tags = set(tags)  # Each tag counted once per doc
        tag_doc_counts.update(unique_tags)

    # Calculate IDF
    tag_idf = {}
    for tag, doc_count in tag_doc_counts.items():
        tag_idf[tag] = math.log(total_docs / doc_count)

    print(f"Calculated IDF for {len(tag_idf):,} unique tags")

    # Show some examples
    sorted_tags = sorted(tag_idf.items(), key=lambda x: x[1], reverse=True)
    print("\nMost distinctive tags (highest IDF):")
    for tag, idf in sorted_tags[:5]:
        print(f"  {tag}: {idf:.4f}")
    print("\nMost common tags (lowest IDF):")
    for tag, idf in sorted_tags[-5:]:
        print(f"  {tag}: {idf:.4f}")

    return tag_idf


def calculate_tfidf_score(entry: Dict, tag_idf: Dict[str, float]) -> float:
    """
    Calculate TF-IDF score for a single entry.

    For simplicity, TF = 1 for each tag (binary: tag present or not)
    Score = sum of IDF values for all tags in the entry
    """
    tags = entry.get("tags", [])

    if not tags:
        return 0.0

    # Sum of IDF scores for all tags
    score = sum(tag_idf.get(tag, 0.0) for tag in tags)

    return score


def analyze_tfidf_distributions(
    entries: List[Dict], tag_idf: Dict[str, float], category_name: str
) -> List[float]:
    """
    Analyze TF-IDF score distribution for a category.

    Returns:
        List of TF-IDF scores
    """
    print(f"\n{category_name}:")
    print(f"  Total entries: {len(entries):,}")

    scores = []
    for entry in entries:
        score = calculate_tfidf_score(entry, tag_idf)
        scores.append(score)

    print(f"  TF-IDF Statistics:")
    print(f"    Min:    {min(scores):.2f}")
    print(f"    Max:    {max(scores):.2f}")
    print(f"    Mean:   {np.mean(scores):.2f}")
    print(f"    Median: {np.median(scores):.2f}")
    if HAS_NUMPY:
        print(f"    Std:    {statistics.stdev(scores):.2f}")
    print(f"    P25:    {np.percentile(scores, 25):.2f}")
    print(f"    P75:    {np.percentile(scores, 75):.2f}")
    print(f"    P95:    {np.percentile(scores, 95):.2f}")

    return scores


def sample_by_tfidf(
    entries: List[Dict],
    tag_idf: Dict[str, float],
    n_samples: int,
    strategy: str = "top",
) -> List[Dict]:
    """
    Sample entries based on TF-IDF scores.

    Args:
        entries: List of metadata entries
        tag_idf: IDF scores for tags
        n_samples: Number of samples to draw
        strategy: "top" for highest scores, "weighted" for weighted random sampling
    """
    if n_samples >= len(entries):
        return entries

    print(f"  Calculating TF-IDF scores for {len(entries):,} entries...")

    # Calculate TF-IDF score for each entry
    scored_entries = []
    for entry in tqdm(entries, desc="  Scoring", disable=True):
        score = calculate_tfidf_score(entry, tag_idf)
        scored_entries.append((entry, score))

    if strategy == "top":
        # Sort by score and take top N
        scored_entries.sort(key=lambda x: x[1], reverse=True)
        sampled = [entry for entry, score in scored_entries[:n_samples]]

        # Print score statistics
        scores = [score for _, score in scored_entries[:n_samples]]
        print(f"  TF-IDF score range: {min(scores):.2f} - {max(scores):.2f}")
        print(f"  Mean TF-IDF: {np.mean(scores):.2f}")

    elif strategy == "weighted":
        # Weighted random sampling (favor high scores but add randomness)
        entries_list = [entry for entry, score in scored_entries]
        scores = [score for entry, score in scored_entries]

        # Normalize scores to probabilities
        min_score = min(scores)
        scores_shifted = [s - min_score + 1e-6 for s in scores]  # Ensure all positive
        total = sum(scores_shifted)
        probabilities = [s / total for s in scores_shifted]

        if HAS_NUMPY:
            indices = np.random.choice(
                len(entries_list), size=n_samples, replace=False, p=probabilities
            )
        else:
            indices = np.random_choice(
                len(entries_list), size=n_samples, replace=False, p=probabilities
            )
        sampled = [entries_list[i] for i in indices]

    else:
        raise ValueError(f"Unknown strategy: {strategy}")

    return sampled


def sample_diverse_subset(
    metadata: List[Dict], target_size: int = 298000, sampling_strategy: str = "top"
) -> List[Dict]:
    """
    Sample diverse subset using TF-IDF.

    Args:
        metadata: Full metadata list
        target_size: Target number of samples
        sampling_strategy: "top" or "weighted"
    """
    print(f"\n{'='*80}")
    print(f"SAMPLING STRATEGY")
    print(f"{'='*80}")
    print(f"Target size: {target_size:,}")
    print(f"Sampling strategy: {sampling_strategy}")
    print()

    # Split into categories
    print("Categorizing entries...")
    non_english_with_text = []
    english_with_text = []
    no_text = []

    for entry in tqdm(metadata, desc="Categorizing"):
        lang = entry.get("lang")
        lang = lang.lower() if lang else ""
        text = entry.get("text")

        if text and lang != "en":
            non_english_with_text.append(entry)
        elif text and lang == "en":
            english_with_text.append(entry)
        else:
            no_text.append(entry)

    print(f"\nCategory breakdown:")
    print(f"  Non-English with text: {len(non_english_with_text):>10,}")
    print(f"  English with text:     {len(english_with_text):>10,}")
    print(f"  No text (all langs):   {len(no_text):>10,}")
    print(f"  Total:                 {len(metadata):>10,}")

    # Calculate global tag IDF
    tag_idf = calculate_tag_idf(metadata)

    # Analyze TF-IDF distributions for each category
    print(f"\n{'='*80}")
    print(f"TF-IDF SCORE DISTRIBUTIONS BY CATEGORY")
    print(f"{'='*80}")

    _ = analyze_tfidf_distributions(english_with_text, tag_idf, "English with text")
    _ = analyze_tfidf_distributions(
        non_english_with_text, tag_idf, "Non-English with text"
    )
    _ = analyze_tfidf_distributions(no_text, tag_idf, "No text (all languages)")

    # Sample selection
    sampled = []

    # 1. Keep ALL non-English with text
    print(f"\n{'='*80}")
    print(f"SAMPLING STEP 1: Keep ALL non-English songs with lyrics")
    print(f"{'='*80}")
    sampled.extend(non_english_with_text)
    print(f"  Added: {len(non_english_with_text):,}")

    remaining_budget = target_size - len(sampled)
    print(f"  Remaining budget: {remaining_budget:,}")

    # 2. Sample English with text using TF-IDF
    print(f"\n{'='*80}")
    print(f"SAMPLING STEP 2: Sample English songs with lyrics (TF-IDF)")
    print(f"{'='*80}")

    # Allocate roughly 70% of remaining to English
    n_english = int(remaining_budget * 0.70)
    n_english = min(n_english, len(english_with_text))

    english_sampled = sample_by_tfidf(
        english_with_text, tag_idf, n_english, strategy=sampling_strategy
    )
    sampled.extend(english_sampled)
    print(f"  Added: {len(english_sampled):,}")

    remaining_budget = target_size - len(sampled)
    print(f"  Remaining budget: {remaining_budget:,}")

    # 3. Sample no-text using TF-IDF
    print(f"\n{'='*80}")
    print(f"SAMPLING STEP 3: Sample songs without lyrics (TF-IDF)")
    print(f"{'='*80}")

    n_no_text = min(remaining_budget, len(no_text))

    no_text_sampled = sample_by_tfidf(
        no_text, tag_idf, n_no_text, strategy=sampling_strategy
    )
    sampled.extend(no_text_sampled)
    print(f"  Added: {len(no_text_sampled):,}")

    print(f"\n{'='*80}")
    print(f"FINAL SAMPLE SIZE: {len(sampled):,}")
    print(f"{'='*80}")

    return sampled


def analyze_sample(sampled: List[Dict]) -> None:
    """Print analysis of the sampled data."""
    print(f"\n{'='*80}")
    print(f"SAMPLE ANALYSIS")
    print(f"{'='*80}")

    # Language distribution
    lang_counts = Counter(entry.get("lang", "unknown") for entry in sampled)
    print(f"\nTop 10 languages:")
    for lang, count in lang_counts.most_common(10):
        pct = 100 * count / len(sampled)
        print(f"  {lang:<10s} {count:>8,} ({pct:>5.2f}%)")

    # Text statistics
    has_text = sum(1 for e in sampled if e.get("text"))
    print(f"\nText statistics:")
    print(f"  Has text: {has_text:,} ({100*has_text/len(sampled):.2f}%)")
    print(
        f"  No text:  {len(sampled)-has_text:,} ({100*(len(sampled)-has_text)/len(sampled):.2f}%)"
    )

    # Tag statistics
    all_tags = []
    for entry in sampled:
        all_tags.extend(entry.get("tags", []))

    tag_counts = Counter(all_tags)
    print(f"\nTag statistics:")
    print(f"  Total tags: {len(all_tags):,}")
    print(f"  Unique tags: {len(tag_counts):,}")
    print(f"  Avg tags per entry: {len(all_tags)/len(sampled):.2f}")

    print(f"\nTop 10 tags:")
    for tag, count in tag_counts.most_common(10):
        print(f"  {tag:<30s} {count:>8,}")

    # Duration statistics
    durations = [e.get("duration_s") for e in sampled if e.get("duration_s")]
    if durations:
        print(f"\nDuration statistics:")
        print(f"  Mean: {np.mean(durations):.2f}s")
        print(f"  Median: {np.median(durations):.2f}s")
        print(f"  Min: {min(durations):.2f}s")
        print(f"  Max: {max(durations):.2f}s")


def write_sampled_metadata(sampled: List[Dict], output_path: str) -> None:
    """Write sampled metadata to JSONL file."""
    print(f"\nWriting sampled metadata to: {output_path}")

    os.makedirs(os.path.dirname(output_path), exist_ok=True)

    with open(output_path, "w") as f:
        for entry in tqdm(sampled, desc="Writing"):
            f.write(json.dumps(entry) + "\n")

    file_size = os.path.getsize(output_path)
    print(f"Written: {len(sampled):,} entries ({file_size / (1024**3):.2f} GB)")


def main():
    """Main function."""
    # Configuration
    input_path = "/home/tony/Work/tony/RealGen/metas_v5_val_filtered_clean.jsonl"
    output_dir = "/home/tony/Work/tony/RealGen"
    output_filename = "metas_v5_val_sampled_diverse.jsonl"
    output_path = os.path.join(output_dir, output_filename)

    target_size = 298000  # ~half of 595,633
    sampling_strategy = "top"  # or "weighted" for more randomness

    print("🎵 Diverse Subset Sampling with TF-IDF")
    print(f"Input:  {input_path}")
    print(f"Output: {output_path}")
    print()

    # Set random seed for reproducibility
    random.seed(42)
    if HAS_NUMPY:
        np.random.seed(42)

    # Load metadata
    metadata = load_metadata(input_path)

    # Sample diverse subset
    sampled = sample_diverse_subset(
        metadata, target_size=target_size, sampling_strategy=sampling_strategy
    )

    # Analyze sample
    analyze_sample(sampled)

    # Write output
    write_sampled_metadata(sampled, output_path)

    print("\n✅ Sampling completed successfully!")


if __name__ == "__main__":
    main()
