"""Claim 3: Embedding geometry analysis.
Run staged training for both Claim 1 (unfrozen) and Claim 2 (frozen),
then extract embeddings and compute similarity heatmaps + UMAP.
"""
import sys
sys.path.insert(0, '/tmp/geometric_memory')

import os
import numpy as np
import torch
from pathlib import Path

OUTPUT_DIR = '/tmp/claim3_analysis'
os.makedirs(OUTPUT_DIR, exist_ok=True)

def run_training(freeze_embeddings=False, epochs=20):
    """Run staged training via the actual train_in_weights.py script."""
    import subprocess
    
    cli_args = [
        'python', '/tmp/geometric_memory/train_in_weights.py',
        '--training_recipe', 'staged_full_path',
        '--model_family', 'gpt',
        '--graph_type', 'star',
        '--star_degree', '10',
        '--star_subtree_degree', '10',
        '--path_length', '4',
        '--total_nodes', '-1',
        '--add_forward_edges',
        '--add_backward_edges',
        '--edge_memorization_epochs', str(epochs),
        '--path_finetuning_epochs', str(epochs),
        '--edge_memorization_batch_size', '64',
        '--path_finetuning_batch_size', '64',
        '--edge_memorization_learning_rate', '0.01',
        '--path_finetuning_learning_rate', '0.0005',
        '--optimizer_weight_decay', '0.0',
        '--edge_memorization_warmup_steps', '0',
        '--path_finetuning_warmup_steps', '0',
        '--disable_edge_memorization_lr_decay',
        '--disable_path_finetuning_lr_decay',
        '--edge_memorization_eval_interval_epochs', '5',
        '--path_finetuning_eval_interval_epochs', '1',
        '--embedding_dimension', '384',
        '--attention_head_count', '8',
        '--transformer_layer_count', '12',
        '--dropout_rate', '0.0',
        '--experiment_log_root', OUTPUT_DIR,
        '--dataset_root', '/tmp/geometric_memory/data/datasets/in_weights_graphs',
        '--dataset_name', 'in_weights',
        '--train_split_ratio', '0.75',
        '--path_prefix_pause_token_count', '1',
        '--exclude_task_token_in_prefix',
        '--track_embedding_evolution',
        '--no-enable_wandb',
    ]
    
    if freeze_embeddings:
        cli_args.append('--freeze_token_embeddings')
    
    print(f"Running training (freeze={freeze_embeddings}, epochs={epochs})...")
    result = subprocess.run(cli_args, cwd='/tmp/geometric_memory', capture_output=True, text=True, timeout=600)
    
    # Find the final model checkpoint
    run_dir = None
    for root, dirs, files in os.walk(OUTPUT_DIR):
        for f in files:
            if f.endswith('_final_model.pt'):
                checkpoint_path = os.path.join(root, f)
                print(f"Found checkpoint: {checkpoint_path}")
                run_dir = root
                break
        if run_dir:
            break
    
    return result, run_dir

def extract_embeddings_from_checkpoint(checkpoint_path, tokenizer, device):
    """Extract embeddings from a trained model checkpoint."""
    from geometric_memory.models import get_model
    from geometric_memory.in_weights.data_loader import EdgeMemorizationDataset
    from torch.utils.data import DataLoader
    import argparse
    
    # Build args for model construction
    args = argparse.Namespace()
    args.model_family = 'gpt'
    args.total_nodes = 11
    args.embedding_dimension = 384
    args.attention_head_count = 8
    args.transformer_layer_count = 12
    args.block_size = 64
    args.bias = True
    args.dropout_rate = 0.0
    args.use_flash = False
    args.teacherless_token = 0
    args.use_attention = True
    args.use_residual_connections = True
    args.freeze_token_embeddings = False
    args.use_layer_norm = True
    args.use_positional_encoding = True
    args.use_mlp_only_blocks = False
    args.tie_input_output_embeddings = False
    args.weight_init_mode = 'default'
    args.vocab_size = 110
    args.use_layernorm = True
    
    model = get_model(args)
    state_dict = torch.load(checkpoint_path, map_location=device, weights_only=False)
    model.load_state_dict(state_dict)
    model = model.to(device)
    model.eval()
    
    # Get tokenizer
    from geometric_memory.tokenizing import get_tokenizer
    tokenizer = get_tokenizer(args)
    
    # Extract embeddings from edge dataset
    dataset = EdgeMemorizationDataset(
        tokenizer=tokenizer,
        data_path=str(Path('/tmp/geometric_memory/data/datasets/in_weights_graphs/star_graphs_randomized') / 
                       'star_deg_10_deg_tree_10_path_4_nodes_1111_sd_00_fb_11_selfedge_0_pretrain.txt'),
        device=device,
        teacherless_token_id=0,
        drop_pause_token=True,
        include_task_token_in_prefix=False,
    )
    
    loader = DataLoader(dataset, batch_size=64, shuffle=False)
    
    all_embeddings = []
    all_inputs = []
    
    with torch.no_grad():
        for batch in loader:
            if isinstance(batch, tuple):
                inputs, targets = batch
            else:
                inputs = batch
                targets = None
            
            inputs = inputs.to(device)
            
            # Get embeddings from the model
            if hasattr(model, 'transformer') and hasattr(model.transformer, 'wte'):
                embeddings = model.transformer.wte(inputs)
            elif hasattr(model, 'wte'):
                embeddings = model.wte(inputs)
            else:
                for name, module in model.named_modules():
                    if 'embed' in name.lower() and hasattr(module, 'weight'):
                        embeddings = module(inputs)
                        break
                else:
                    print("Warning: Could not find embedding layer")
                    continue
            
            # Take the last token embedding for each sequence
            last_embeddings = embeddings[:, -1, :]  # (batch_size, embed_dim)
            
            all_embeddings.append(last_embeddings.cpu().numpy())
            all_inputs.append(inputs.cpu().numpy())
    
    all_embeddings = np.concatenate(all_embeddings, axis=0)
    all_inputs = np.concatenate(all_inputs, axis=0)
    
    print(f"Extracted {len(all_embeddings)} embeddings, shape: {all_embeddings.shape}")
    return all_embeddings, all_inputs

def compute_similarity_heatmap(embeddings, output_path):
    """Compute and save embedding similarity heatmap."""
    import matplotlib
    matplotlib.use('Agg')
    import matplotlib.pyplot as plt
    
    # Compute dot-product similarity matrix
    sim_matrix = embeddings @ embeddings.T
    # Normalize to [0, 1]
    sim_matrix = (sim_matrix - sim_matrix.min()) / (sim_matrix.max() - sim_matrix.min())
    
    # Plot heatmap
    fig, ax = plt.subplots(figsize=(12, 10))
    im = ax.imshow(sim_matrix, cmap='viridis', aspect='auto')
    ax.set_title('Embedding Similarity Heatmap')
    ax.set_xlabel('Token Index')
    ax.set_ylabel('Token Index')
    plt.colorbar(im, ax=ax)
    plt.tight_layout()
    plt.savefig(output_path, dpi=150, bbox_inches='tight')
    plt.close()
    print(f"Saved similarity heatmap to {output_path}")

def run_umap_projection(embeddings, output_path):
    """Run UMAP projection on embeddings."""
    try:
        import umap
        reducer = umap.UMAP(n_neighbors=15, min_dist=0.1, metric='cosine')
        embedding_2d = reducer.fit_transform(embeddings)
        
        import matplotlib
        matplotlib.use('Agg')
        import matplotlib.pyplot as plt
        
        fig, ax = plt.subplots(figsize=(10, 8))
        scatter = ax.scatter(embedding_2d[:, 0], embedding_2d[:, 1], c=range(len(embeddings)), 
                           cmap='viridis', s=10, alpha=0.7)
        ax.set_title('UMAP Projection of Embeddings')
        ax.set_xlabel('UMAP 1')
        ax.set_ylabel('UMAP 2')
        plt.colorbar(scatter, ax=ax)
        plt.tight_layout()
        plt.savefig(output_path, dpi=150, bbox_inches='tight')
        plt.close()
        print(f"Saved UMAP projection to {output_path}")
    except ImportError:
        print("UMAP not available, skipping projection")

def main():
    print("=" * 60)
    print("Claim 3: Embedding Geometry Analysis")
    print("=" * 60)
    
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"Device: {device}")
    
    # ── Run Claim 1: unfrozen embeddings ──
    print("\n" + "=" * 60)
    print("Claim 1: Running staged training (unfrozen embeddings)")
    print("=" * 60)
    
    result1, run_dir1 = run_training(freeze_embeddings=False, epochs=20)
    print(result1.stdout[-1000:] if len(result1.stdout) > 1000 else result1.stdout)
    if result1.returncode != 0:
        print(f"Claim 1 training failed: {result1.stderr}")
        return
    
    # ── Run Claim 2: frozen embeddings ──
    print("\n" + "=" * 60)
    print("Claim 2: Running staged training (frozen embeddings)")
    print("=" * 60)
    
    result2, run_dir2 = run_training(freeze_embeddings=True, epochs=20)
    print(result2.stdout[-1000:] if len(result2.stdout) > 1000 else result2.stdout)
    if result2.returncode != 0:
        print(f"Claim 2 training failed: {result2.stderr}")
        return
    
    # ── Extract embeddings from checkpoints ──
    print("\n" + "=" * 60)
    print("Extracting embeddings from trained models")
    print("=" * 60)
    
    # Find final model checkpoints
    import argparse
    from geometric_memory.tokenizing import get_tokenizer
    
    args = argparse.Namespace()
    args.model_family = 'gpt'
    args.total_nodes = 11
    args.embedding_dimension = 384
    args.attention_head_count = 8
    args.transformer_layer_count = 12
    args.block_size = 64
    args.bias = True
    args.dropout_rate = 0.0
    args.use_flash = False
    args.teacherless_token = 0
    args.use_attention = True
    args.use_residual_connections = True
    args.freeze_token_embeddings = False
    args.use_layer_norm = True
    args.use_positional_encoding = True
    args.use_mlp_only_blocks = False
    args.tie_input_output_embeddings = False
    args.weight_init_mode = 'default'
    args.vocab_size = 110
    args.use_layernorm = True
    
    tokenizer = get_tokenizer(args)
    
    # Extract from Claim 1 checkpoint
    checkpoint1 = None
    for root, dirs, files in os.walk(OUTPUT_DIR):
        for f in files:
            if 'claim1' in f.lower() and f.endswith('_final_model.pt'):
                checkpoint1 = os.path.join(root, f)
                break
        if checkpoint1:
            break
    
    if not checkpoint1:
        # Try to find any claim1 checkpoint
        for root, dirs, files in os.walk(OUTPUT_DIR):
            for f in files:
                if f.endswith('_final_model.pt') and 'claim1' in root.lower():
                    checkpoint1 = os.path.join(root, f)
                    break
            if checkpoint1:
                break
    
    if checkpoint1:
        print(f"\nExtracting from Claim 1 checkpoint: {checkpoint1}")
        embeddings_c1, inputs_c1 = extract_embeddings_from_checkpoint(checkpoint1, tokenizer, device)
        np.save(f'{OUTPUT_DIR}/claim1_embeddings.npy', embeddings_c1)
        np.save(f'{OUTPUT_DIR}/claim1_inputs.npy', inputs_c1)
        compute_similarity_heatmap(embeddings_c1, f'{OUTPUT_DIR}/claim1_similarity_heatmap.png')
        run_umap_projection(embeddings_c1, f'{OUTPUT_DIR}/claim1_umap_projection.png')
    else:
        print("Warning: Claim 1 checkpoint not found")
    
    # Extract from Claim 2 checkpoint
    checkpoint2 = None
    for root, dirs, files in os.walk(OUTPUT_DIR):
        for f in files:
            if 'claim2' in f.lower() and f.endswith('_final_model.pt'):
                checkpoint2 = os.path.join(root, f)
                break
        if checkpoint2:
            break
    
    if not checkpoint2:
        # Try to find any claim2 checkpoint
        for root, dirs, files in os.walk(OUTPUT_DIR):
            for f in files:
                if f.endswith('_final_model.pt') and 'claim2' in root.lower():
                    checkpoint2 = os.path.join(root, f)
                    break
            if checkpoint2:
                break
    
    if checkpoint2:
        print(f"\nExtracting from Claim 2 checkpoint: {checkpoint2}")
        embeddings_c2, inputs_c2 = extract_embeddings_from_checkpoint(checkpoint2, tokenizer, device)
        np.save(f'{OUTPUT_DIR}/claim2_embeddings.npy', embeddings_c2)
        np.save(f'{OUTPUT_DIR}/claim2_inputs.npy', inputs_c2)
        compute_similarity_heatmap(embeddings_c2, f'{OUTPUT_DIR}/claim2_similarity_heatmap.png')
        run_umap_projection(embeddings_c2, f'{OUTPUT_DIR}/claim2_umap_projection.png')
    else:
        print("Warning: Claim 2 checkpoint not found")
    
    # ── Compare embeddings ──
    print("\n" + "=" * 60)
    print("Comparison: Unfrozen vs Frozen Embeddings")
    print("=" * 60)
    
    if checkpoint1 and checkpoint2:
        embeddings_c1 = np.load(f'{OUTPUT_DIR}/claim1_embeddings.npy')
        embeddings_c2 = np.load(f'{OUTPUT_DIR}/claim2_embeddings.npy')
        
        if len(embeddings_c1) == len(embeddings_c2):
            # Compute correlation between embeddings
            corr = np.corrcoef(embeddings_c1.flatten(), embeddings_c2.flatten())[0, 1]
            print(f"Embedding correlation (unfrozen vs frozen): {corr:.4f}")
            
            # Compute mean pairwise similarity
            sim_c1 = (embeddings_c1 @ embeddings_c1.T).diagonal().mean()
            sim_c2 = (embeddings_c2 @ embeddings_c2.T).diagonal().mean()
            print(f"Mean self-similarity (unfrozen): {sim_c1:.4f}")
            print(f"Mean self-similarity (frozen): {sim_c2:.4f}")
            
            # Compute distance between corresponding embeddings
            dist = np.linalg.norm(embeddings_c1 - embeddings_c2, axis=1).mean()
            print(f"Mean distance between corresponding embeddings: {dist:.4f}")
    
    print("\n" + "=" * 60)
    print("Claim 3 Analysis Complete")
    print("=" * 60)
    print(f"Output directory: {OUTPUT_DIR}")
    print("Files generated:")
    for f in os.listdir(OUTPUT_DIR):
        print(f"  - {OUTPUT_DIR}/{f}")

if __name__ == '__main__':
    main()
