import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import pycwt as wavelet

# Load your data from CSV
file_path = 'E:/data_new.csv'
data = pd.read_csv(file_path)

# Extract Time and GR columns
time = data['Time'].values   # Time axis (kyr)
gr = data['GR'].values       # Gamma-ray proxy data

# Verify time range and calculate sampling interval
t_start, t_end = time[0], time[-1]
dt = np.mean(np.diff(time))  # Sampling interval
total_time = t_end - t_start

print(f"Time range: {t_start:.2f} to {t_end:.2f} kyr, dt: {dt:.3f} kyr, Total time: {total_time:.2f} kyr")

# Standardize the GR data
gr = (gr - np.mean(gr)) / np.std(gr)

# Define periods in kyr
max_period = total_time / 2
periods_base = np.array([0.5, 1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 5000, 8192, 16384])
periods_kyr = periods_base[periods_base <= max_period]

# Wavelet transform parameters
mother = wavelet.Morlet(6)
s0 = periods_kyr[0] * dt
dj = 0.1
J = int(np.log2(periods_kyr[-1] / periods_kyr[0]) / dj) + 1

# Perform the continuous wavelet transform
wave, scales, freqs, coi, fft, fftfreqs = wavelet.cwt(gr, dt, dj, s0, J, mother)

# Convert scales to periods in kyr
periods = scales

# Compute the power spectrum
power = np.abs(wave) ** 2

# Significance testing
alpha, _, _ = wavelet.ar1(gr)
signif, fft_theor = wavelet.significance(1.0, dt, scales, 0, alpha,
                                         significance_level=0.95, wavelet=mother)
sig95 = power / signif[:, np.newaxis]

# Plotting
plt.figure(figsize=(12, 8))

# Wavelet power spectrum
plt.subplot(2, 1, 1)
plt.contourf(time, periods, np.log2(power), cmap='bwr')
plt.colorbar(label='log2(Power)')
plt.contour(time, periods, sig95, [1], colors='k', linewidths=1)
plt.plot(time, coi, 'k--', linewidth=1.5)
plt.title('Wavelet Power Spectrum (GR Proxy)')
plt.xlabel('Time (kyr)')
plt.ylabel('Period (kyr)')
plt.yscale('log')
plt.gca().invert_yaxis()

# Original GR signal
plt.subplot(2, 1, 2)
plt.plot(time, gr, 'b-', linewidth=1)
plt.title('Gamma-Ray (GR) Proxy')
plt.xlabel('Time (kyr)')
plt.ylabel('Standardized GR')

plt.tight_layout()
plt.show()

# Global wavelet spectrum
global_power = power.mean(axis=1)
plt.figure(figsize=(6, 4))
plt.plot(global_power, periods, 'b-')
plt.plot(signif, periods, 'k--', label='95% significance')
plt.title('Global Wavelet Spectrum')
plt.xlabel('Power')
plt.ylabel('Period (kyr)')
plt.yscale('log')
plt.gca().invert_yaxis()
plt.legend()
plt.show()
