#!/usr/bin/env python3

################################################################################
# SCRIPT 07: COMPARATIVE REPEAT ANALYSIS - Generate Analysis Plots and Tables
################################################################################

# PURPOSE:
# Performs comprehensive comparative analysis of repeat composition across
# samples . Generates visualizations (bar charts, heatmaps, clustermaps with
# standard scaling) and statistical comparison tables. Identifies differential
# enrichment between all sample pairs with Log2 fold-change calculations.

# INPUT:
# - RepeatMasker .out files from script 05 (auto-detect or explicit)
# - Or repeat_summary_long.tsv from script 06 (long format)
# - Or repeat_summary_pivoted.tsv from script 06 (wide format)

# OUTPUT:
# PLOTS:
# - plot_top_classes_comparison.png - Top classes bar plot (horizontal)
# - plot_top_families_comparison.png - Top families bar plot (horizontal)
# - plot_heatmap_classes.png - Class composition heatmap
# - plot_heatmap_families.png - Family composition heatmap
# - plot_clustermap_classes.png - Hierarchical clustering (WITH standard scaling)
# - plot_clustermap_families.png - Family-based clustering
# - plot_diff_class_<S1>_vs_<S2>.png - Differential plots for each pair
# - plot_diff_family_<S1>_vs_<S2>.png - Family differential plots
#
# TABLES:
# - table_summary_classes.tsv - Complete class statistics
# - table_summary_families.tsv - Complete family statistics
# - table_variation_classes.tsv - Variation metrics (CV, enrichment ratio)
# - table_variation_families.tsv - Family variation metrics
# - table_diff_class_<S1>_vs_<S2>.tsv - Differential analysis for each pair
# - table_diff_family_<S1>_vs_<S2>.tsv - Family differential for each pair

# USAGE:
# python 07_repeat_analysis.py [options]

# OPTIONS:
# # From .out files (auto-detect)
# python 07_repeat_analysis.py
#
# # From .out files (explicit)
# python 07_repeat_analysis.py \
#   --input-files outputs/05_repeatmasker/rca.fasta.out \
#                 outputs/05_repeatmasker/gdna-nt.fasta.out \
#   --sample-names RCA gDNA-NT
#
# # From long-format TSV (recommended)
# python 07_repeat_analysis.py \
#   --input-summary outputs/06_repeat_profiling/repeat_summary_long.tsv
#
# # From pivoted TSV
# python 07_repeat_analysis.py \
#   --input-pivoted outputs/06_repeat_profiling/repeat_summary_pivoted.tsv

# DEPENDENCIES:
# - pandas
# - numpy
# - matplotlib
# - seaborn
# - scipy (for clustering)

################################################################################

import os
import sys
import argparse
import logging
import glob
from pathlib import Path
from collections import defaultdict

import pandas as pd
import numpy as np

import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import seaborn as sns
from matplotlib.patches import Patch

# ============================================================================
# LOGGING SETUP
# ============================================================================

logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)

# ============================================================================
# MATPLOTLIB SETUP
# ============================================================================

plt.rcParams['figure.dpi'] = 150
plt.rcParams['savefig.dpi'] = 300
plt.rcParams['font.size'] = 10
sns.set_style("whitegrid")
sns.set_palette("viridis")

# ============================================================================
# REPEATMASKER PARSING (from script 06, for .out file input support)
# ============================================================================

def parse_repeatmasker_out(out_file):
    """Parse RepeatMasker .out file (same as script 06)."""
    annotations = []
    
    try:
        with open(out_file, 'r') as f:
            for line_num, line in enumerate(f):
                line = line.strip()
                
                # Skip header lines
                if not line or line.startswith(('SW', 'score', 'There', ' ')) or line_num < 3:
                    continue
                
                parts = line.split()
                if len(parts) < 11:
                    continue
                
                try:
                    repeat_class_family = parts[10]
                    
                    # Parse class/family with split('/', 1)
                    if '/' in repeat_class_family:
                        repeat_class = repeat_class_family.split('/', 1)[0]
                        repeat_family = repeat_class_family.split('/', 1)[1]
                    else:
                        repeat_class = repeat_class_family
                        repeat_family = parts[9]
                    
                    length = abs(int(parts[6]) - int(parts[5])) + 1
                    
                    annotations.append({
                        'repeat_class': repeat_class,
                        'repeat_family': repeat_family,
                        'length': length
                    })
                    
                except (ValueError, IndexError):
                    continue
                    
        return annotations
        
    except Exception as e:
        logger.error(f"Error parsing {out_file}: {e}")
        return []

def calculate_composition_from_annotations(annotations, sample_name):
    """Calculate composition percentages from parsed annotations."""
    if not annotations:
        return pd.DataFrame()
    
    df = pd.DataFrame(annotations)
    total_bp = df['length'].sum()
    
    if total_bp == 0:
        return pd.DataFrame()
    
    results = []
    
    # By class
    for cls in df['repeat_class'].unique():
        bp = df[df['repeat_class'] == cls]['length'].sum()
        results.append({
            'sample_name': sample_name,
            'category_type': 'Class',
            'category_name': cls,
            'percentage_of_total_bp': (bp / total_bp) * 100
        })
    
    # By family
    for fam in df['repeat_family'].unique():
        bp = df[df['repeat_family'] == fam]['length'].sum()
        results.append({
            'sample_name': sample_name,
            'category_type': 'Family',
            'category_name': fam,
            'percentage_of_total_bp': (bp / total_bp) * 100
        })
    
    return pd.DataFrame(results)

# ============================================================================
# VARIATION ANALYSIS
# ============================================================================

def analyze_variation(df, category_col, value_col):
    """
    Calculate variation statistics across samples for each category.
    
    Args:
        df (pd.DataFrame): Long-format data with sample_name, category, value
        category_col (str): Column name for categories
        value_col (str): Column name for values (percentages)
        
    Returns:
        pd.DataFrame: Variation statistics sorted by CV
    """
    if df.empty:
        logger.warning(f"No data provided for variation analysis on '{category_col}'")
        return pd.DataFrame()
    
    if value_col not in df.columns or category_col not in df.columns or 'sample_name' not in df.columns:
        logger.error(f"Missing required columns for variation analysis")
        return pd.DataFrame()
    
    try:
        # Ensure value col is numeric
        df[value_col] = pd.to_numeric(df[value_col], errors='coerce')
        df_valid = df.dropna(subset=[value_col])
        
        if df_valid.empty:
            logger.warning(f"No valid numeric data in '{value_col}' for variation analysis")
            return pd.DataFrame()
        
        # Pivot to get samples as columns
        pivot_df = df_valid.pivot_table(
            index=category_col, columns='sample_name', values=value_col, fill_value=0
        )
        
        if pivot_df.empty:
            logger.warning(f"Pivoting resulted in an empty table")
            return pd.DataFrame()
        
        # Calculate statistics
        stats_df = pd.DataFrame(index=pivot_df.index)
        stats_df['Mean'] = pivot_df.mean(axis=1)
        stats_df['Std'] = pivot_df.std(axis=1)
        stats_df['Min'] = pivot_df.min(axis=1)
        stats_df['Max'] = pivot_df.max(axis=1)
        stats_df['Range'] = stats_df['Max'] - stats_df['Min']
        
        # Calculate CV with proper handling
        stats_df['CV'] = (stats_df['Std'] / stats_df['Mean']).replace(
            [np.inf, -np.inf], np.nan
        ).fillna(0) * 100
        
        # Identify max/min samples 
        stats_df['Max_Sample'] = pivot_df.idxmax(axis=1)
        stats_df['Min_Sample'] = pivot_df.idxmin(axis=1)
        
        # Add enrichment ratio with epsilon
        epsilon = 1e-9
        stats_df['Enrichment_Ratio'] = stats_df['Max'] / (stats_df['Min'] + epsilon)
        # Handle cases where Max is also near zero
        stats_df.loc[stats_df['Max'] < epsilon, 'Enrichment_Ratio'] = 1.0
        
        return stats_df.sort_values('CV', ascending=False)
        
    except Exception as e:
        logger.error(f"Error during variation analysis for {category_col}: {e}", exc_info=True)
        return pd.DataFrame()

# ============================================================================
# DIFFERENTIAL ANALYSIS
# ============================================================================

def perform_differential_analysis(df, category_col, value_col):
    """
    Calculate differences between all pairs of samples.
    
    Args:
        df (pd.DataFrame): Long-format data
        category_col (str): Column for categories
        value_col (str): Column for values
        
    Returns:
        list: List of dicts with 'pair' and 'data' keys
    """
    if df.empty:
        logger.warning(f"No data provided for differential analysis on '{category_col}'")
        return []
    
    if value_col not in df.columns or category_col not in df.columns or 'sample_name' not in df.columns:
        logger.error(f"Missing required columns for differential analysis")
        return []
    
    samples = df['sample_name'].unique()
    
    if len(samples) < 2:
        logger.warning("Need at least two samples for differential analysis")
        return []
    
    # Create all pairwise combinations
    pairs = [(s1, s2) for i, s1 in enumerate(samples) for s2 in samples[i+1:]]
    
    results = []
    
    # Ensure value col is numeric
    df[value_col] = pd.to_numeric(df[value_col], errors='coerce')
    
    for s1, s2 in pairs:
        try:
            # Pivot for the pair
            pivot_pair = df[df['sample_name'].isin([s1, s2])].pivot_table(
                index=category_col, columns='sample_name', values=value_col, fill_value=0
            )
            
            # Ensure both samples are present
            if s1 not in pivot_pair.columns:
                pivot_pair[s1] = 0
            if s2 not in pivot_pair.columns:
                pivot_pair[s2] = 0
            
            # Create combined DataFrame
            combined = pivot_pair[[s1, s2]].copy()
            combined['Difference'] = combined[s1] - combined[s2]
            
            # Add epsilon for log2 fold change
            epsilon = 1e-9
            combined['Log2FC'] = np.log2((combined[s1] + epsilon) / (combined[s2] + epsilon))
            
            combined['Abs_Difference'] = abs(combined['Difference'])
            combined['Enriched_In'] = np.where(combined['Difference'] >= 0, s1, s2)
            
            combined = combined.sort_values('Abs_Difference', ascending=False)
            
            results.append({'pair': (s1, s2), 'data': combined})
            
        except Exception as e:
            logger.error(f"Error during differential analysis between {s1} and {s2}: {e}", exc_info=True)
    
    return results

# ============================================================================
# VISUALIZATION FUNCTIONS
# ============================================================================

def plot_comparison_bar(df, category_col, value_col, sample_col, top_n, title, output_file):
    """
    Creates horizontal bar chart comparing top elements across samples.
    """
    logger.info(f"Generating comparison bar plot: {title}")
    
    if df.empty:
        logger.warning(f"No data provided for comparison bar plot '{title}'. Skipping.")
        return
    
    try:
        # Ensure value_col is numeric
        df[value_col] = pd.to_numeric(df[value_col], errors='coerce')
        df.dropna(subset=[value_col], inplace=True)
        
        if df.empty:
            logger.warning(f"No numeric data found for '{value_col}' in bar plot '{title}'. Skipping.")
            return
        
        #  Get top n categories by average percentage 
        top_categories = df.groupby(category_col)[value_col].mean().sort_values(
            ascending=False
        ).head(top_n).index.tolist()
        
        # Filter for only top categories
        plot_df = df[df[category_col].isin(top_categories)].copy()
        
        if plot_df.empty:
            logger.warning(f"No data found for top {top_n} categories. Skipping plot.")
            return
        
        # Sort and create categorical for ordering
        category_means = plot_df.groupby(category_col)[value_col].mean().reindex(
            top_categories
        ).sort_values(ascending=False)
        ordered_categories = category_means.index.tolist()
        
        plot_df[category_col] = pd.Categorical(
            plot_df[category_col],
            categories=ordered_categories,
            ordered=True
        )
        
        plot_df.sort_values(by=[category_col, sample_col], inplace=True)
        
        # Dynamic height
        plt.figure(figsize=(12, max(5, top_n * 0.5)))
        
        sns.barplot(x=value_col, y=category_col, hue=sample_col, 
                   data=plot_df, palette='viridis', dodge=True)
        
        plt.title(title, fontsize=14, fontweight='bold')
        plt.xlabel('Percentage of Total Masked BP')
        plt.ylabel(category_col.replace('category_', '').capitalize())
        plt.legend(title='Sample', bbox_to_anchor=(1.05, 1), loc='upper left')
        plt.tight_layout(rect=[0, 0, 0.85, 1])
        
        plt.savefig(output_file, dpi=300, bbox_inches='tight')
        logger.info(f"  Saved: {output_file}")
        plt.close()
        
    except Exception as e:
        logger.error(f"Failed to generate comparison bar plot '{title}': {e}", exc_info=True)

def plot_heatmap(df_pivot, title, output_file, max_cols=50):
    """
    Creates heatmap from pivoted DataFrame.
    
    """
    logger.info(f"Generating heatmap: {title}")
    
    if df_pivot.empty:
        logger.warning(f"Empty data for heatmap '{title}'. Skipping plot.")
        return
    
    try:
        # Limit number of columns for readability
        if df_pivot.shape[1] > max_cols:
            logger.warning(
                f"Too many categories ({df_pivot.shape[1]}) for heatmap '{title}', "
                f"showing top {max_cols} by mean percentage."
            )
            df_pivot_numeric = df_pivot.apply(pd.to_numeric, errors='coerce').fillna(0)
            top_cols = df_pivot_numeric.mean().sort_values(ascending=False).head(max_cols).index
            df_plot = df_pivot[top_cols]
        else:
            df_plot = df_pivot
        
        if df_plot.empty:
            logger.warning(f"No columns left after filtering for heatmap '{title}'. Skipping.")
            return
        
        plt.figure(figsize=(max(10, df_plot.shape[1] * 0.4), 
                           max(6, df_plot.shape[0] * 0.5)))
        
        sns.heatmap(df_plot.astype(float), cmap='viridis', annot=False, 
                   linewidths=.5, cbar_kws={'label': 'Percentage of Total Masked BP'})
        
        plt.title(title, fontsize=14, fontweight='bold')
        plt.xticks(rotation=90)
        plt.yticks(rotation=0)
        plt.tight_layout()
        
        plt.savefig(output_file, dpi=300, bbox_inches='tight')
        logger.info(f"  Saved: {output_file}")
        plt.close()
        
    except Exception as e:
        logger.error(f"Failed to generate heatmap '{title}': {e}", exc_info=True)

def plot_clustermap(df_pivot, title, output_file, max_cols=50):
    """
    Creates clustermap with standard scaling (Z-score normalization).

    Includes standard_scale=1 for proper hierarchical clustering.
    """
    logger.info(f"Generating clustermap: {title}")
    
    if df_pivot.empty:
        logger.warning(f"Empty data for clustermap '{title}'. Skipping plot.")
        return
    
    # Need at least 2 samples and 2 features
    if df_pivot.shape[0] < 2:
        logger.warning(f"Insufficient samples ({df_pivot.shape[0]}) for clustermap '{title}'. Skipping.")
        return
    
    if df_pivot.shape[1] < 2:
        logger.warning(f"Insufficient features ({df_pivot.shape[1]}) for clustermap '{title}'. Skipping.")
        return
    
    try:
        # Limit number of columns
        if df_pivot.shape[1] > max_cols:
            logger.warning(
                f"Too many categories ({df_pivot.shape[1]}) for clustermap '{title}', "
                f"showing top {max_cols} by mean percentage."
            )
            df_pivot_numeric = df_pivot.apply(pd.to_numeric, errors='coerce').fillna(0)
            top_cols = df_pivot_numeric.mean().sort_values(ascending=False).head(max_cols).index
            df_plot = df_pivot[top_cols].astype(float)
        else:
            df_plot = df_pivot.astype(float)
        
        if df_plot.empty or df_plot.shape[1] < 2:
            logger.warning(f"No columns/features left after filtering for clustermap '{title}'. Skipping.")
            return
        
        # Check for zero variance columns/rows before scaling
        non_zero_var_cols = df_plot.columns[df_plot.std(axis=0) > 0]
        non_zero_var_rows = df_plot.index[df_plot.std(axis=1) > 0]
        
        if len(non_zero_var_cols) < df_plot.shape[1] or len(non_zero_var_rows) < df_plot.shape[0]:
            logger.warning(f"Removing zero-variance rows/columns for standard scaling in clustermap '{title}'.")
            df_plot = df_plot.loc[non_zero_var_rows, non_zero_var_cols]
        
        if df_plot.shape[0] < 2 or df_plot.shape[1] < 2:
            logger.warning(
                f"Insufficient data after removing zero-variance rows/columns "
                f"for clustermap '{title}'. Skipping plot."
            )
            return
        
        # Hierarchical clustering with standard scaling
        # standard_scale=1 normalizes columns (Z-score transformation)
        g = sns.clustermap(
            df_plot,
            cmap='viridis',
            standard_scale=1,  #Column-wise Z-score normalization
            figsize=(max(10, df_plot.shape[1] * 0.3), max(8, df_plot.shape[0] * 0.6)),
            linewidths=.5,
            cbar_kws={'label': 'Scaled Percentage'}
        )
        
        plt.suptitle(title, y=1.02, fontsize=14, fontweight='bold')
        
        plt.savefig(output_file, dpi=300, bbox_inches='tight')
        logger.info(f"  Saved: {output_file}")
        plt.close()
        
    except Exception as e:
        logger.error(f"Failed to generate clustermap '{title}': {e}", exc_info=True)

def plot_differential_enrichment(diff_df, s1, s2, top_n, category_label, title, output_file):
    """
    Plots log2 Fold Change for top differing categories.
    
    """
    logger.info(f"Generating differential enrichment plot: {title}")
    
    if diff_df.empty:
        logger.warning(f"Empty data for differential plot '{title}'. Skipping.")
        return
    
    try:
        # Sort by absolute difference
        plot_data = diff_df.sort_values('Abs_Difference', ascending=False).head(top_n).reset_index()
        
        if plot_data.empty:
            logger.warning(f"No data after filtering for differential plot '{title}'. Skipping.")
            return
        
        # Get category column dynamically
        category_col = plot_data.columns[0]
        
        # Create plotting dataframe
        plot_df = pd.DataFrame({
            'Category': plot_data[category_col],
            'Log2FC': plot_data['Log2FC'],
            'EnrichedIn': plot_data['Enriched_In']
        })
        
        # Reverse order for plotting
        plot_df = plot_df.iloc[::-1]
        
        # Set color based on enriched sample
        colors = ['#e41a1c' if x == s1 else '#377eb8' for x in plot_df['EnrichedIn']]
        
        plt.figure(figsize=(10, max(5, top_n * 0.4)))
        
        bars = plt.barh(y=plot_df['Category'], width=plot_df['Log2FC'], color=colors)
        
        # Add legend
        legend_elements = [
            Patch(facecolor='#e41a1c', label=f'Enriched in {s1}'),
            Patch(facecolor='#377eb8', label=f'Enriched in {s2}')
        ]
        plt.legend(handles=legend_elements, loc='best')
        
        plt.axvline(x=0, color='black', linestyle='-', linewidth=0.5)
        plt.title(f'{title}\n(Top {top_n} differences between {s1} and {s2})', 
                 fontsize=12, fontweight='bold')
        plt.xlabel('Log2 Fold Change (Percentage of Total Masked BP)')
        plt.ylabel(category_label)
        plt.tight_layout()
        
        plt.savefig(output_file, dpi=300, bbox_inches='tight')
        logger.info(f"  Saved: {output_file}")
        plt.close()
        
    except Exception as e:
        logger.error(f"Failed to generate differential enrichment plot '{title}': {e}", exc_info=True)

# ============================================================================
# DATA LOADING FUNCTIONS
# ============================================================================

def load_data_from_long_format(tsv_file):
    """Load data from long-format TSV (output of script 06)."""
    logger.info(f"Loading long-format data from: {tsv_file}")
    
    try:
        df = pd.read_csv(tsv_file, sep='\t')
        
        # Validate columns
        required_cols = ['sample_name', 'category_type', 'category_name', 'percentage_of_total_bp']
        missing = [col for col in required_cols if col not in df.columns]
        
        if missing:
            logger.error(f"Missing required columns: {missing}")
            return pd.DataFrame()
        
        logger.info(f"  Loaded {len(df)} rows")
        return df
        
    except Exception as e:
        logger.error(f"Error loading long-format data: {e}")
        return pd.DataFrame()

def load_data_from_pivoted(tsv_file):
    """Load data from pivoted wide-format TSV and convert to long format."""
    logger.info(f"Loading pivoted data from: {tsv_file}")
    
    try:
        df = pd.read_csv(tsv_file, sep='\t')
        
        if 'sample_name' not in df.columns and 'sample' not in df.columns:
            logger.error("No sample column found in pivoted data")
            return pd.DataFrame()
        
        # Rename 'sample' to 'sample_name' if needed
        if 'sample' in df.columns:
            df.rename(columns={'sample': 'sample_name'}, inplace=True)
        
        # Melt to long format
        id_cols = ['sample_name']
        value_cols = [col for col in df.columns if col not in id_cols]
        
        df_long = df.melt(
            id_vars=id_cols,
            value_vars=value_cols,
            var_name='category_full',
            value_name='percentage_of_total_bp'
        )
        
        # Split 'Class_LINE' into category_type='Class' and category_name='LINE'
        df_long[['category_type', 'category_name']] = df_long['category_full'].str.split('_', n=1, expand=True)
        df_long.drop('category_full', axis=1, inplace=True)
        
        logger.info(f"  Converted to long format: {len(df_long)} rows")
        return df_long
        
    except Exception as e:
        logger.error(f"Error loading pivoted data: {e}")
        return pd.DataFrame()

def auto_detect_input_files():
    """Auto-detect input files from standard pipeline locations."""
    # Try to find long-format output from script 06
    long_path = "outputs/06_repeat_profiling/repeat_summary_long.tsv"
    if os.path.exists(long_path):
        logger.info(f"Auto-detected long-format data: {long_path}")
        return 'long', long_path
    
    # Try pivoted format
    pivot_path = "outputs/06_repeat_profiling/repeat_summary_pivoted.tsv"
    if os.path.exists(pivot_path):
        logger.info(f"Auto-detected pivoted data: {pivot_path}")
        return 'pivoted', pivot_path
    
    # Try .out files from script 05
    out_dir = "outputs/05_repeatmasker"
    if os.path.exists(out_dir):
        out_files = glob.glob(os.path.join(out_dir, "*.out"))
        if out_files:
            logger.info(f"Auto-detected {len(out_files)} .out files in {out_dir}")
            return 'out_files', out_files
    
    return None, None

# ============================================================================
# MAIN ANALYSIS
# ============================================================================

def main():
    parser = argparse.ArgumentParser(
        description=(
            "Comprehensive comparative repeat analysis. Generates visualizations "
            "with statistical analyses including hierarchical clustering with "
            "standard scaling and pairwise differential enrichment."
        ),
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog="""
EXAMPLES:
  # Auto-detect from pipeline outputs
  python 07_repeat_analysis.py
  
  # From long-format TSV (recommended)
  python 07_repeat_analysis.py \\
    --input-summary outputs/06_repeat_profiling/repeat_summary_long.tsv
  
  # From pivoted TSV
  python 07_repeat_analysis.py \\
    --input-pivoted outputs/06_repeat_profiling/repeat_summary_pivoted.tsv
  
  # From .out files (explicit)
  python 07_repeat_analysis.py \\
    --input-files outputs/05_repeatmasker/rca.fasta.out \\
                  outputs/05_repeatmasker/gdna-nt.fasta.out \\
    --sample-names RCA gDNA-NT
  
  # Custom top N and output directory
  python 07_repeat_analysis.py \\
    --input-summary outputs/06_repeat_profiling/repeat_summary_long.tsv \\
    --top-n 25 \\
    --output-dir outputs/07_repeat_analysis

OUTPUT FILES:
  PLOTS:
    - plot_top_classes_comparison.png
    - plot_top_families_comparison.png
    - plot_heatmap_classes.png
    - plot_heatmap_families.png
    - plot_clustermap_classes.png (WITH standard scaling)
    - plot_clustermap_families.png (WITH standard scaling)
    - plot_diff_class_<S1>_vs_<S2>.png (for each sample pair)
    - plot_diff_family_<S1>_vs_<S2>.png (for each sample pair)
  
  TABLES:
    - table_summary_classes.tsv
    - table_summary_families.tsv
    - table_variation_classes.tsv
    - table_variation_families.tsv
    - table_diff_class_<S1>_vs_<S2>.tsv (for each sample pair)
    - table_diff_family_<S1>_vs_<S2>.tsv (for each sample pair)

NEXT STEP:
  Rscript scripts/08_statistical_analysis.R
        """
    )
    
    # Input specification
    input_group = parser.add_mutually_exclusive_group()
    
    input_group.add_argument(
        "--input-summary",
        help="Long-format TSV from script 06 (repeat_summary_long.tsv)"
    )
    
    input_group.add_argument(
        "--input-pivoted",
        help="Pivoted TSV from script 06 (repeat_summary_pivoted.tsv)"
    )
    
    input_group.add_argument(
        "--input-files",
        nargs='+',
        help="RepeatMasker .out files (requires --sample-names)"
    )
    
    parser.add_argument(
        "--sample-names",
        nargs='+',
        help="Sample names (required with --input-files, must match order)"
    )
    
    # Analysis parameters
    parser.add_argument(
        "--top-n",
        type=int,
        default=20,
        help="Number of top categories to show in plots. Default: %(default)s"
    )
    
    parser.add_argument(
        "--max-cols",
        type=int,
        default=50,
        help="Maximum columns in heatmaps/clustermaps. Default: %(default)s"
    )
    
    # Output
    parser.add_argument(
        "--output-dir",
        default="outputs/07_repeat_analysis",
        help="Output directory. Default: %(default)s"
    )
    
    parser.add_argument(
        "--dpi",
        type=int,
        default=300,
        help="Plot DPI for saved figures. Default: %(default)s"
    )
    
    args = parser.parse_args()
    
    logger.info("="*70)
    logger.info("COMPARATIVE REPEAT ANALYSIS (Script 07)")
    logger.info("="*70)
    logger.info("")
    
    # Update matplotlib DPI
    plt.rcParams['savefig.dpi'] = args.dpi
    
    # ========================================================================
    # LOAD DATA
    # ========================================================================
    
    df_long = pd.DataFrame()
    
    if args.input_summary:
        # Load from long-format TSV
        df_long = load_data_from_long_format(args.input_summary)
        
    elif args.input_pivoted:
        # Load from pivoted TSV and convert
        df_long = load_data_from_pivoted(args.input_pivoted)
        
    elif args.input_files:
        # Load from .out files
        if not args.sample_names:
            logger.error("--sample-names required with --input-files")
            sys.exit(1)
        
        if len(args.input_files) != len(args.sample_names):
            logger.error("Number of files must match number of sample names")
            sys.exit(1)
        
        logger.info(f"Processing {len(args.input_files)} .out files...")
        
        all_data = []
        for sample_name, out_file in zip(args.sample_names, args.input_files):
            if not os.path.exists(out_file):
                logger.error(f"File not found: {out_file}")
                sys.exit(1)
            
            logger.info(f"  Parsing: {sample_name}")
            annotations = parse_repeatmasker_out(out_file)
            sample_df = calculate_composition_from_annotations(annotations, sample_name)
            
            if not sample_df.empty:
                all_data.append(sample_df)
        
        if all_data:
            df_long = pd.concat(all_data, ignore_index=True)
        else:
            logger.error("No data loaded from .out files")
            sys.exit(1)
    
    else:
        # Auto-detect
        logger.info("No input specified, attempting auto-detection...")
        input_type, input_path = auto_detect_input_files()
        
        if input_type == 'long':
            df_long = load_data_from_long_format(input_path)
        elif input_type == 'pivoted':
            df_long = load_data_from_pivoted(input_path)
        elif input_type == 'out_files':
            logger.error("Found .out files but cannot auto-determine sample names")
            logger.error("Please use --input-files with --sample-names")
            sys.exit(1)
        else:
            logger.error("Could not auto-detect input files")
            logger.error("Please specify input with --input-summary, --input-pivoted, or --input-files")
            sys.exit(1)
    
    if df_long.empty:
        logger.error("No data loaded")
        sys.exit(1)
    
    logger.info(f"Loaded data: {len(df_long)} rows")
    logger.info(f"Samples: {df_long['sample_name'].nunique()}")
    logger.info("")
    
    # Split into class and family datasets
    df_class = df_long[df_long['category_type'] == 'Class'].copy()
    df_family = df_long[df_long['category_type'] == 'Family'].copy()
    
    logger.info(f"Class entries: {len(df_class)}")
    logger.info(f"Family entries: {len(df_family)}")
    logger.info("")
    
    # ========================================================================
    # CREATE OUTPUT DIRECTORY
    # ========================================================================
    
    os.makedirs(args.output_dir, exist_ok=True)
    logger.info(f"Output directory: {args.output_dir}")
    logger.info("")
    
    # ========================================================================
    # 1. SUMMARY TABLES
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 1: Creating summary tables")
    logger.info("="*70)
    
    # Class summary
    class_summary_file = os.path.join(args.output_dir, "table_summary_classes.tsv")
    df_class.to_csv(class_summary_file, sep='\t', index=False)
    logger.info(f"✓ Class summary: {class_summary_file}")
    
    # Family summary
    family_summary_file = os.path.join(args.output_dir, "table_summary_families.tsv")
    df_family.to_csv(family_summary_file, sep='\t', index=False)
    logger.info(f"✓ Family summary: {family_summary_file}")
    logger.info("")
    
    # ========================================================================
    # 2. VARIATION ANALYSIS
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 2: Calculating variation metrics")
    logger.info("="*70)
    
    # Class variation
    logger.info("Analyzing class variation...")
    var_class = analyze_variation(df_class, 'category_name', 'percentage_of_total_bp')
    
    if not var_class.empty:
        var_class_file = os.path.join(args.output_dir, "table_variation_classes.tsv")
        var_class.to_csv(var_class_file, sep='\t')
        logger.info(f"✓ Class variation: {var_class_file}")
        logger.info(f"  Top 5 most variable classes (by CV):")
        for idx, row in var_class.head(5).iterrows():
            logger.info(f"    {idx}: CV={row['CV']:.2f}%, Enrichment={row['Enrichment_Ratio']:.2f}x")
    
    # Family variation
    logger.info("\nAnalyzing family variation...")
    var_family = analyze_variation(df_family, 'category_name', 'percentage_of_total_bp')
    
    if not var_family.empty:
        var_family_file = os.path.join(args.output_dir, "table_variation_families.tsv")
        var_family.to_csv(var_family_file, sep='\t')
        logger.info(f"✓ Family variation: {var_family_file}")
        logger.info(f"  Top 5 most variable families (by CV):")
        for idx, row in var_family.head(5).iterrows():
            logger.info(f"    {idx}: CV={row['CV']:.2f}%, Enrichment={row['Enrichment_Ratio']:.2f}x")
    
    logger.info("")
    
    # ========================================================================
    # 3. COMPARISON BAR PLOTS
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 3: Generating comparison bar plots")
    logger.info("="*70)
    
    # Class comparison
    out_class_bar = os.path.join(args.output_dir, "plot_top_classes_comparison.png")
    plot_comparison_bar(
        df_class, 'category_name', 'percentage_of_total_bp', 'sample_name',
        args.top_n, f'Top {args.top_n} Repeat Classes Across Samples',
        out_class_bar
    )
    
    # Family comparison
    out_family_bar = os.path.join(args.output_dir, "plot_top_families_comparison.png")
    plot_comparison_bar(
        df_family, 'category_name', 'percentage_of_total_bp', 'sample_name',
        args.top_n, f'Top {args.top_n} Repeat Families Across Samples',
        out_family_bar
    )
    
    logger.info("")
    
    # ========================================================================
    # 4. HEATMAPS
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 4: Generating heatmaps")
    logger.info("="*70)
    
    # Pivot for heatmaps
    df_pivot_class = df_class.pivot_table(
        index='sample_name', columns='category_name',
        values='percentage_of_total_bp', fill_value=0
    )
    
    df_pivot_family = df_family.pivot_table(
        index='sample_name', columns='category_name',
        values='percentage_of_total_bp', fill_value=0
    )
    
    # Class heatmap
    out_heatmap_class = os.path.join(args.output_dir, "plot_heatmap_classes.png")
    plot_heatmap(df_pivot_class, 'Repeat Class Composition Heatmap', 
                 out_heatmap_class, max_cols=args.max_cols)
    
    # Family heatmap
    out_heatmap_family = os.path.join(args.output_dir, "plot_heatmap_families.png")
    plot_heatmap(df_pivot_family, 'Repeat Family Composition Heatmap',
                 out_heatmap_family, max_cols=args.max_cols)
    
    logger.info("")
    
    # ========================================================================
    # 5. CLUSTERMAPS (WITH STANDARD SCALING)
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 5: Generating clustermaps with standard scaling")
    logger.info("="*70)
    logger.info("Using standard_scale=1 for Z-score normalization")
    logger.info("")
    
    # Class clustermap
    out_clustermap_class = os.path.join(args.output_dir, "plot_clustermap_classes.png")
    plot_clustermap(df_pivot_class, 'Hierarchical Clustering - Repeat Classes',
                    out_clustermap_class, max_cols=args.max_cols)
    
    # Family clustermap
    out_clustermap_family = os.path.join(args.output_dir, "plot_clustermap_families.png")
    plot_clustermap(df_pivot_family, 'Hierarchical Clustering - Repeat Families',
                    out_clustermap_family, max_cols=args.max_cols)
    
    logger.info("")
    
    # ========================================================================
    # 6. DIFFERENTIAL ANALYSIS
    # ========================================================================
    
    logger.info("="*70)
    logger.info("STEP 6: Performing pairwise differential analysis")
    logger.info("="*70)
    
    # Class differential
    logger.info("Analyzing class differences between sample pairs...")
    diff_results_class = perform_differential_analysis(
        df_class, 'category_name', 'percentage_of_total_bp'
    )
    
    for result in diff_results_class:
        s1, s2 = result['pair']
        diff_df = result['data']
        
        # Save table
        out_diff_table = os.path.join(args.output_dir, f"table_diff_class_{s1}_vs_{s2}.tsv")
        diff_df.head(args.top_n).to_csv(out_diff_table, sep='\t')
        logger.info(f"  ✓ Table: {s1} vs {s2}")
        
        # Generate plot
        out_diff_plot = os.path.join(args.output_dir, f"plot_diff_class_{s1}_vs_{s2}.png")
        plot_differential_enrichment(
            diff_df, s1, s2, args.top_n, 'Repeat Class',
            'Repeat Class Differential Enrichment', out_diff_plot
        )
    
    # Family differential
    logger.info("\nAnalyzing family differences between sample pairs...")
    diff_results_family = perform_differential_analysis(
        df_family, 'category_name', 'percentage_of_total_bp'
    )
    
    for result in diff_results_family:
        s1, s2 = result['pair']
        diff_df = result['data']
        
        # Save table
        out_diff_table = os.path.join(args.output_dir, f"table_diff_family_{s1}_vs_{s2}.tsv")
        diff_df.head(args.top_n).to_csv(out_diff_table, sep='\t')
        logger.info(f"  ✓ Table: {s1} vs {s2}")
        
        # Generate plot
        out_diff_plot = os.path.join(args.output_dir, f"plot_diff_family_{s1}_vs_{s2}.png")
        plot_differential_enrichment(
            diff_df, s1, s2, args.top_n, 'Repeat Family',
            'Repeat Family Differential Enrichment', out_diff_plot
        )
    
    logger.info("")
    
    # ========================================================================
    # SUMMARY REPORT
    # ========================================================================
    
    logger.info("="*70)
    logger.info("COMPARATIVE REPEAT ANALYSIS COMPLETE")
    logger.info("="*70)
    logger.info("")
    
    # Count outputs
    all_files = sorted(os.listdir(args.output_dir))
    plots = [f for f in all_files if f.startswith('plot_')]
    tables = [f for f in all_files if f.startswith('table_')]
    
    logger.info(f"Generated {len(plots)} plots and {len(tables)} tables")
    logger.info("")
    
    logger.info("OUTPUT FILES:")
    logger.info("  Plots:")
    for f in plots:
        logger.info(f"    - {f}")
    logger.info("  Tables:")
    for f in tables:
        logger.info(f"    - {f}")
    logger.info("")
    
    logger.info("NEXT STEP:")
    logger.info("  Rscript scripts/08_statistical_analysis.R")
    logger.info("")
    logger.info("="*70)

if __name__ == "__main__":
    main()
