"""
Tool for analyzing formula components during training.
Tracks and visualizes how formula components evolve.
"""

import torch
import numpy as np
import json
from typing import Dict, List, Optional
from collections import defaultdict
import os


class FormulaComponentTracker:
    """
    Tracks formula component values during training for analysis.
    """
    
    def __init__(self, save_dir: str = './analysis'):
        self.save_dir = save_dir
        os.makedirs(save_dir, exist_ok=True)
        
        self.epoch_data = defaultdict(list)
        self.step_data = defaultdict(list)
        self.current_epoch = 0
        self.current_step = 0
        
    def log_step(self, components: Dict[str, torch.Tensor], step: Optional[int] = None):
        """Log formula components at a training step."""
        if step is not None:
            self.current_step = step
        
        step_entry = {}
        for key, value in components.items():
            if value is None:
                continue
            if isinstance(value, torch.Tensor):
                # Convert to numpy for storage
                if value.numel() == 1:
                    step_entry[key] = float(value.item())
                else:
                    step_entry[key] = {
                        'mean': float(value.mean().item()),
                        'std': float(value.std().item()),
                        'min': float(value.min().item()),
                        'max': float(value.max().item()),
                    }
            else:
                step_entry[key] = value
        
        self.step_data[self.current_step].append(step_entry)
        
    def log_epoch(self, components: Dict[str, torch.Tensor], epoch: Optional[int] = None):
        """Log formula components at end of epoch."""
        if epoch is not None:
            self.current_epoch = epoch
        
        epoch_entry = {}
        for key, value in components.items():
            if value is None:
                continue
            if isinstance(value, torch.Tensor):
                if value.numel() == 1:
                    epoch_entry[key] = float(value.item())
                else:
                    epoch_entry[key] = {
                        'mean': float(value.mean().item()),
                        'std': float(value.std().item()),
                        'min': float(value.min().item()),
                        'max': float(value.max().item()),
                        'median': float(value.median().item()),
                    }
            elif isinstance(value, dict):
                # Handle nested dicts (e.g., weights)
                epoch_entry[key] = {k: float(v) if isinstance(v, torch.Tensor) else v 
                                    for k, v in value.items()}
            else:
                epoch_entry[key] = value
        
        self.epoch_data[self.current_epoch].append(epoch_entry)
        
    def save_epoch_summary(self, epoch: int):
        """Save summary statistics for an epoch."""
        if epoch not in self.epoch_data:
            return
        
        epoch_entries = self.epoch_data[epoch]
        
        # Aggregate statistics
        summary = {
            'epoch': epoch,
            'num_logs': len(epoch_entries),
        }
        
        # Aggregate each component across all logs
        component_keys = set()
        for entry in epoch_entries:
            component_keys.update(entry.keys())
        
        for key in component_keys:
            values = [entry.get(key) for entry in epoch_entries if key in entry]
            if not values:
                continue
                
            if isinstance(values[0], (int, float)):
                summary[key] = {
                    'mean': float(np.mean(values)),
                    'std': float(np.std(values)),
                    'min': float(np.min(values)),
                    'max': float(np.max(values)),
                }
            elif isinstance(values[0], dict) and 'mean' in values[0]:
                # Already aggregated values
                means = [v['mean'] for v in values if isinstance(v, dict) and 'mean' in v]
                if means:
                    summary[key] = {
                        'mean_of_means': float(np.mean(means)),
                        'std_of_means': float(np.std(means)),
                    }
                # Also keep individual entries
                summary[f'{key}_all'] = values
        
        # Save to file
        summary_path = os.path.join(self.save_dir, f'epoch_{epoch}_formula_summary.json')
        with open(summary_path, 'w') as f:
            json.dump(summary, f, indent=2)
        
        return summary
    
    def save_training_summary(self):
        """Save complete training summary with all epochs."""
        all_summaries = {}
        
        for epoch in sorted(self.epoch_data.keys()):
            summary = self.save_epoch_summary(epoch)
            if summary:
                all_summaries[epoch] = summary
        
        # Save complete summary
        summary_path = os.path.join(self.save_dir, 'training_formula_summary.json')
        with open(summary_path, 'w') as f:
            json.dump(all_summaries, f, indent=2)
        
        return all_summaries
    
    def get_component_trends(self, component_name: str) -> Dict[str, List[float]]:
        """Get trend of a component across epochs."""
        trends = {
            'mean': [],
            'std': [],
            'min': [],
            'max': [],
        }
        
        for epoch in sorted(self.epoch_data.keys()):
            epoch_entries = self.epoch_data[epoch]
            if not epoch_entries:
                continue
            
            # Extract values for this component
            values = []
            for entry in epoch_entries:
                comp = entry.get(component_name)
                if comp is None:
                    continue
                if isinstance(comp, dict) and 'mean' in comp:
                    values.append(comp['mean'])
                elif isinstance(comp, (int, float)):
                    values.append(float(comp))
            
            if values:
                trends['mean'].append(float(np.mean(values)))
                trends['std'].append(float(np.std(values)))
                
                if isinstance(epoch_entries[0].get(component_name), dict):
                    all_vals = []
                    for entry in epoch_entries:
                        comp = entry.get(component_name)
                        if isinstance(comp, dict):
                            if 'min' in comp:
                                all_vals.append(comp['min'])
                            if 'max' in comp:
                                all_vals.append(comp['max'])
                    if all_vals:
                        trends['min'].append(float(np.min(all_vals)))
                        trends['max'].append(float(np.max(all_vals)))
        
        return trends


def analyze_formula_components(model, tokenizer, data_loader, device='cuda'):
    """
    Analyze formula components on a dataset.
    
    Returns:
        Dictionary with aggregated component statistics
    """
    model.eval()
    tracker = FormulaComponentTracker()
    
    all_components = defaultdict(list)
    
    with torch.no_grad():
        for batch_idx, batch in enumerate(data_loader):
            input_ids = batch['input_ids'].to(device)
            attention_mask = batch.get('attention_mask')
            if attention_mask is not None:
                attention_mask = attention_mask.to(device)
            
            # Get model output with components
            output = model(input_ids=input_ids, attention_mask=attention_mask, return_components=True)
            
            if 'components' in output:
                components = output['components']
                
                # Aggregate components
                for key, value in components.items():
                    if value is not None and isinstance(value, torch.Tensor):
                        # Convert to numpy
                        if value.numel() > 0:
                            all_components[key].append(value.cpu().numpy())
            
            # Limit to avoid memory issues
            if batch_idx >= 100:
                break
    
    # Aggregate statistics
    summary = {}
    for key, values_list in all_components.items():
        if values_list:
            all_values = np.concatenate([v.flatten() for v in values_list])
            summary[key] = {
                'mean': float(np.mean(all_values)),
                'std': float(np.std(all_values)),
                'min': float(np.min(all_values)),
                'max': float(np.max(all_values)),
                'median': float(np.median(all_values)),
            }
    
    return summary


def print_formula_analysis(summary: Dict):
    """Print a formatted analysis of formula components."""
    print("\n" + "=" * 70)
    print("FORMULA COMPONENT ANALYSIS")
    print("=" * 70)
    
    component_order = ['novelty', 'retention', 'payoff', 'continuity', 'fatigue', 'weighted_sum']
    
    for comp_name in component_order:
        if comp_name in summary:
            stats = summary[comp_name]
            print(f"\n{comp_name.upper()}:")
            print(f"  Mean:   {stats['mean']:.6f} ± {stats['std']:.6f}")
            print(f"  Range:  [{stats['min']:.6f}, {stats['max']:.6f}]")
            print(f"  Median: {stats['median']:.6f}")
    
    # Print weights if available
    if 'weights' in summary:
        weights = summary['weights']
        if isinstance(weights, dict):
            print(f"\nFORMULA WEIGHTS:")
            for w_name, w_value in weights.items():
                if isinstance(w_value, (int, float)):
                    print(f"  {w_name}: {w_value:.6f}")
    
    print("\n" + "=" * 70)


