import numpy as np
from copy import deepcopy

from devito import Constant, Eq, Operator
from recipes import ElasticTTI, WaveModel, ElasticIsotropic, AcousticIsotropic
from scipy.ndimage import rotate

import matplotlib.pyplot as plt

shape = (2001, 2001)
origin = (0, 0)
spacing = (10, 10)
space_order = 8
dt_comp = 1.0

mid = shape[0] // 2

vp = Constant("vp", value=3.0)
vs = Constant("vs", value=3.0/1.5)
epsilon = Constant("epsilon", value=0.2)
delta = Constant("delta", value=0.1)
theta = Constant("theta", value=np.pi/4)
b = Constant("b", value=1.0)
zeps = Constant("epsilon", value=0.0)
zdel = Constant("delta", value=0.0)
vszero = Constant("vs", value=0.0)

modelvti = WaveModel(origin, spacing, shape, space_order, vp, nbl=10,
                     bcs="mask", b=b, vs=vs, epsilon=epsilon, delta=delta, dt=dt_comp)

modeltti = WaveModel(origin, spacing, shape, space_order, vp, nbl=10,
                     bcs="mask", b=b, vs=vs, epsilon=epsilon, delta=delta,
                     theta=theta, dt=dt_comp)

opt = {'nt': 2500, 'space_order': space_order, 'f0': 0.005, 'measurements': False}

solver_vti = ElasticTTI(modelvti, opt, fd='rsfd')
solver_vti_fd = ElasticTTI(modelvti, opt, fd='ssg')
solver_tti = ElasticTTI(modeltti, opt, fd='rsfd')
solver_elas = ElasticIsotropic(modelvti, opt)
solver_acou = AcousticIsotropic(modelvti, opt)
solver_tti_fd = ElasticTTI(modeltti, opt, fd='ssg')


src = solver_vti.src
src.coordinates.data[:] = np.array(modelvti.domain_size) * .5

src2 = deepcopy(src)._rebuild(name="src2")
Operator(Eq(src2, src.dt))(dt=dt_comp)
src2.coordinates.data[:] = np.array(modelvti.domain_size) * .5

#################################
# Elas vs Acou
#################################

solver_acou.forward(src=src2)
solver_elas.forward(src=src, vs=vszero)

u0 = (2*modelvti.critical_dt)*solver_acou.pressure_data
u1 = solver_elas.pressure_data

clip = 5e-2
plt.figure(figsize=(10, 5))
plt.subplot(131)
plt.imshow(u0.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("Acoustic")
plt.subplot(132)
plt.imshow(u1.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("Elastic zero vs")
plt.subplot(133)
plt.imshow(u1.T - u0.T, cmap='seismic', vmin=-clip/100, vmax=clip/100)
plt.title("Difference X100")
plt.tight_layout()
plt.savefig("./images/elas_vs_acou.png", bbox_inches="tight")

plt.figure(figsize=(10, 6))
plt.subplot(211)
plt.plot(u0[mid, mid:], label="Acoustic")
plt.plot(u1[mid, mid:], label="Elastic zero vs")
plt.legend(loc='upper left')
plt.subplot(212)
plt.plot(u1[mid, mid:]-u0[mid, mid:], label="Difference")
plt.legend(loc='best')
plt.tight_layout()
plt.savefig("./images/elas_vs_acou_trace.png", bbox_inches="tight")


#################################
# AnisoElas vs Elas
#################################

solver_elas.forward(src=src)
solver_vti.forward(src=src, epsilon=zeps, delta=zdel)

u2 = solver_elas.pressure_data
u3 = solver_vti.pressure_data

clip = 5e-2
plt.figure(figsize=(10, 5))
plt.subplot(131)
plt.imshow(u2.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("Isotropic")
plt.subplot(132)
plt.imshow(u3.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("Zero Anisotropy")
plt.subplot(133)
plt.imshow(u2.T - u3.T, cmap='seismic', vmin=-clip/100, vmax=clip/100)
plt.title("Difference X100")
plt.tight_layout()
plt.savefig("./images/elas_vs_anisoelas.png", bbox_inches="tight")

plt.figure(figsize=(10, 6))
plt.subplot(211)
plt.plot(u2[mid, mid:], label="Elastic Isotropic")
plt.plot(u3[mid, mid:], label="Zero Anisotropy Elastic Anisotropic")
plt.legend(loc='upper left')
plt.subplot(212)
plt.plot(u2[mid, mid:]-u3[mid, mid:], label="Difference")
plt.legend(loc='best')
plt.tight_layout()
plt.savefig("./images/elas_vs_anisoelas_trace.png", bbox_inches="tight")

#################################
# VTI vs TTI
#################################

solver_vti.forward(src=src)
solver_vti_fd.forward(src=src)
solver_tti.forward(src=src)
solver_tti_fd.forward(src=src)

u4 = solver_vti.pressure_data
u4_sg = solver_vti_fd.pressure_data
u5 = solver_tti.pressure_data
u5_sg = solver_tti_fd.pressure_data

u6 = rotate(u5, -45, reshape=False)
u6_sg = rotate(u5_sg, -45, reshape=False)

clip = 1e-3
plt.figure(figsize=(10, 10))
plt.subplot(231)
plt.imshow(u5.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("TTI RSFD")
plt.subplot(234)
plt.imshow(u5_sg.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("TTI SG")
plt.subplot(232)
plt.imshow(u6.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("TTI RSFD projected to VTI")
plt.subplot(235)
plt.imshow(u6_sg.T, cmap='seismic', vmin=-clip, vmax=clip)
plt.title("TTI SG projected to VTI")
plt.subplot(233)
plt.imshow(u6.T - u4.T, cmap='seismic', vmin=-clip/100, vmax=clip/100)
plt.title("Difference X10")
plt.subplot(236)
plt.imshow(u6_sg.T - u4_sg.T, cmap='seismic', vmin=-clip/100, vmax=clip/100)
plt.title("Difference X10")
plt.tight_layout()
plt.savefig("./images/vti_vs_tti.png", bbox_inches="tight")

plt.figure(figsie=(10, 6))
plt.subplot(211)
plt.plot(u4[mid, mid:], label="VTI")
plt.plot(u4_sg[mid, mid:], label="VTI SG")
plt.plot(u6[mid, mid:], label="TTI RSFD projected to vti")
plt.plot(u6_sg[mid, mid:], label="TTI SG projected to vti")
plt.legend(loc='upper left')
plt.subplot(212)
plt.plot(u6[mid, mid:]-u4[mid, mid:], label="TTI RSFD diff")
plt.plot(u6_sg[mid, mid:]-u4_sg[mid, mid:], label="TTI SG diff")
plt.legend()
plt.savefig("./images/vti_vs_tti_trace.png", bbox_inches="tight")


#################################
# SSG vs RSG
#################################
solver_tti_fd = ElasticTTI(modeltti, opt, fd='ssg')

# Stability w.r.t theta

vs = Constant("vs", value=3.0/2.0)

plt.figure(figsize=(16, 8))

thetas = [0, 20, 40, 60]

for (i, theta) in enumerate(thetas):
    rad = np.deg2rad(theta)
    thetaloc = Constant(name="theta", value=rad)

    solver_tti_fd.forward(src=src, theta=thetaloc, vs=vs)
    solver_tti.forward(src=src, theta=thetaloc, vs=vs)

    uloc_fd = solver_tti_fd.pressure_data
    uloc = solver_tti.pressure_data

    clip = 1e-3
    plt.subplot(2, 4, i+1)
    plt.imshow(uloc_fd.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"SSG {theta}°")
    plt.subplot(2, 4, i+5)
    plt.imshow(uloc.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"RSG {theta}°")


plt.tight_layout()
plt.savefig("./images/fd_rsfd_poisson2.png", bbox_inches="tight")

# Stability w.r.t vs

vs = Constant("vs", value=3.0/1.5)

plt.figure(figsize=(16, 8))

for (i, theta) in enumerate(thetas):
    rad = np.deg2rad(theta)
    thetaloc = Constant(name="theta", value=rad)

    solver_tti_fd.forward(src=src, theta=thetaloc, vs=vs)
    solver_tti.forward(src=src, theta=thetaloc, vs=vs)

    uloc_fd = solver_tti_fd.pressure_data
    uloc = solver_tti.pressure_data

    clip = 5e-3
    plt.subplot(2, 4, i+1)
    plt.imshow(uloc_fd.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"SSG {theta}°")
    plt.subplot(2, 4, i+5)
    plt.imshow(uloc.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"RSG {theta}°")


plt.tight_layout()
plt.savefig("./images/fd_rsfd_poisson15.png", bbox_inches="tight")

# Stability w.r.t dt

dts = [1.0, 1.5, 2.0, 2.25]
plt.figure(figsize=(16, 8))

for (i, dt) in enumerate(dts):
    nt = int(np.ceil(1250 / dt))
    opt = {'nt': nt, 'space_order': space_order, 'f0': 0.005, 'measurements': False}
    modeltti._dt = dt
    solver_tti = ElasticTTI(modeltti, opt, fd='rsfd')
    solver_tti_fd = ElasticTTI(modeltti, opt, fd='ssg')
    src = solver_tti.src
    src.coordinates.data[:] = [5000., 5000.]

    solver_tti_fd.forward(src=src, dt=dt)
    solver_tti.forward(src=src, dt=dt)

    uloc_fd = solver_tti_fd.pressure_data
    uloc = solver_tti.pressure_data

    clip = 5e-3
    plt.subplot(2, 4, i+1)
    plt.imshow(uloc_fd.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"SSG dt={dt}ms")
    plt.subplot(2, 4, i+5)
    plt.imshow(uloc.T, cmap='seismic', vmin=-clip, vmax=clip)
    plt.title(f"RSG dt{dt}ms")


plt.tight_layout()
plt.savefig("./images/fd_rsfd_dt.png", bbox_inches="tight")
plt.show()
