import pandas as pd
import matplotlib.pyplot as plt
import numpy as np
from scipy.stats import pearsonr, spearmanr

# Load the data
mutation_file = 'mutation_summary.csv'  # Replace with the actual file path
fraction_file = 'fraction_genome_altered.csv'  # Replace with the actual file path

mutation_data = pd.read_csv(mutation_file)
fraction_data = pd.read_csv(fraction_file)

# Merge the datasets on the common field 'Sample_ID'
data = pd.merge(mutation_data, fraction_data, on='Sample_ID')

# Extract relevant columns
x = data['Fraction_Genome_Altered']
y = data['Mutation_Count']

# Calculate correlations
pearson_corr, pearson_p = pearsonr(x, y)
spearman_corr, spearman_p = spearmanr(x, y)

# Create the scatter plot with gradient color
colors = np.log1p(y)  # Apply a log transformation for better color scaling
plt.figure(figsize=(10, 6))
sc = plt.scatter(x, y, c=colors, cmap='viridis', alpha=0.7, edgecolor='k', label=f'Samples: {len(x)}')

# Add colorbar
cbar = plt.colorbar(sc)
cbar.set_label('Log(Mutation Count)', fontsize=12)

# Add annotations
plt.text(0.05, max(y) * 0.9, f'Pearson: {pearson_corr:.4f}\np={pearson_p:.2g}', fontsize=12)
plt.text(0.05, max(y) * 0.8, f'Spearman: {spearman_corr:.4f}\np={spearman_p:.2g}', fontsize=12)

# Enhance the plot
plt.title('Mutation Count vs Fraction Genome Altered', fontsize=14)
plt.xlabel('Fraction Genome Altered', fontsize=12)
plt.ylabel('Mutation Count', fontsize=12)
plt.grid(alpha=0.3)
plt.legend()
plt.tight_layout()

# Save and show the plot
plt.savefig('Mutation_Count_vs_Fraction_Genome_Altered.png', dpi=300)
plt.show()
