import numpy as np
from numpy import matlib
from mpl_toolkits import mplot3d 
import matplotlib.pyplot as plt
from mpl_toolkits import mplot3d
from mpl_toolkits.mplot3d import Axes3D
from mpl_toolkits.mplot3d import proj3d
import scipy.sparse
import scipy.sparse.linalg
from scipy.sparse import csr_matrix
from scipy.sparse import spdiags
import timeit
import sys

# exact solution for purposes of comparison
def uex_fun(x,y):
    val= np.sin(np.pi*x)*np.sin(np.pi*y)
    return val

# right-hand side function f
def f_fun(x,y):
    val= 2*np.pi**2*np.sin(np.pi*x)*np.sin(np.pi*y)
    return val

def FDM_coords(n):
    N=(n+1)**2
    coord_x=np.reshape(matlib.repmat(1/n*np.arange(0,n+1),1,n+1),(N,1),order='F')
    coord_y=np.reshape(matlib.repmat(1/n*np.arange(0,n+1),n+1,1),(N,1),order='F')
    return coord_x, coord_y

def FDM_data(n,f):
    h=1/n
    N=(n+1)**2
    data = np.array([-np.ones(N),-np.ones(N),4*np.ones(N),-np.ones(N),-np.ones(N)])
    diags = np.array([-n-1,-1,0, 1,n+1])
    A=spdiags(data, diags, N, N).tocsr();
    coord_x, coord_y=FDM_coords(n)
    b=h**2*np.reshape(f(coord_x,coord_y),(n+1)**2)
    return A, b, h

def restrict2dof(n):
    N=(n+1)**2
    bottom=np.arange(0,n+1)
    top=np.arange(n*(n+1),N)
    left=(n+1)*np.arange(1,n)
    right=left+n
    bdry=np.concatenate((bottom,right,top,left))
    dof=np.setdiff1d(range(0,N),bdry)
    ndof=np.size(dof)
    R=csr_matrix((np.ones(ndof),(dof,np.arange(0,ndof))),shape = (N,ndof))
    return R

def prolongation(n):
    N=2*n
    m=(n+1)**2;
    M=(N+1)**2
    oldnode=np.zeros(m)
    for j in range(0,n+1):
        oldnode[range(j*(n+1),(j+1)*(n+1))]=np.arange(2*j*(N+1),(2*j+1)*(N+1),2)
    oldnode=np.reshape(np.int_(oldnode),(n+1,n+1))
    I=np.zeros((n+1,n+1,9)); J=np.zeros((n+1,n+1,9)); V=np.zeros((n+1,n+1,9))
    dummy=M
    val=np.asarray([.5,.5,0,.5,1,.5,0,.5,.5])
    for j in range(0,n+1):
        for k in range(0,n+1):
            J[j,k,:]=(j*(n+1)+k)*np.ones(9)
            V[j,k,:]=val
            p=oldnode[j,k]
            nw=p+N;   north=p+N+1; ne=p+N+2
            west=p-1;              east=p+1
            sw=p-N-2; south=p-N-1; se=p-N
            if j==0: sw=dummy; south=dummy; se=dummy
            if j==n: nw=dummy; north=dummy; ne=dummy
            if k==0: nw=dummy; west=dummy; sw=dummy
            if k==n: ne=dummy; east=dummy; se=dummy
            I[j,k,:]=np.asarray([sw,south,se,west,p,east,nw,north,ne])
    I=np.reshape(I,(9*m,1)).T
    J=np.reshape(J,(9*m,1)).T
    V=np.reshape(V,(9*m,1)).T
    P=csr_matrix((V[0,:],(I[0,:],J[0,:])),shape = (M+1,m))
    R=csr_matrix((np.ones(M),(np.arange(0,M),np.arange(0,M))),shape = (M,M+1))
    P=R*P
    return P

def relax(A, b, u,n_smooth):
    for _ in range(n_smooth):
        u = u - 1/8 * (A*u-b)
    return u

def Wcycle(coarse,fine,A,b,x,n_smooth):
    S=restrict2dof(2**fine)
    P=prolongation(2**(fine-1))
    A_inner=(S.transpose()@A)@S
    b_inner=S.transpose()@b
    if coarse==fine:
        x=S*scipy.sparse.linalg.spsolve(A_inner,b_inner)
    else:
        for _ in range(0,2):
            x=S*relax(A_inner, b_inner, S.transpose()*x,n_smooth)
            r=b-A*x;
            q=Wcycle(coarse,fine-1,P.transpose()@(A@P),P.transpose()*r,0*P.transpose()*r,n_smooth)       
            x=x+P*q
            x=S*relax(A_inner, b_inner, S.transpose()*x,n_smooth)
    return x
    
def FDM_mg(coarse,fine,f,n_iter,n_smooth):
    n=2**fine
    A, b, h = FDM_data(n,f)
    x=np.zeros((n+1)**2)
    for m in range(0,n_iter):
        x = Wcycle(coarse,fine,A,b,x,n_smooth)
    return x, h

#-------- the numerical experiment --------------------------
u_exact=np.vectorize(uex_fun)
f=np.vectorize(f_fun)

coarse=1 #2^coarse intervals per axis
fine=10
n_iter=3
n_smooth=2
n_meshes=fine-coarse+1
hlist=np.zeros(n_meshes)
errlist=np.zeros(n_meshes)
for j in range(0,n_meshes):
    lvl=coarse+j
    n=2**(lvl)
    print("%10.3E" %(n+1)**2,' nodes')
    start = timeit.default_timer()
    x, h=FDM_mg(coarse,lvl,f,n_iter,n_smooth)
    stop = timeit.default_timer()
    coord_x, coord_y=FDM_coords(n)
    uex4nodes=np.reshape(u_exact(coord_x,coord_y),(n+1)**2)
    e=uex4nodes-x
    maxerr = np.max(np.abs(e))
    hlist[j]=h
    errlist[j]=maxerr
    print(h)
    print('mesh ',j+1,'\b/',n_meshes,'',"%10.3E" %(n+1)**2,' nodes',
          '  h=',"%10.3E" %h,'  max error=',"%10.3E" %maxerr,'  Time: ', "%10.3E" %(stop - start))

#-----------------------------------------------------------------------
# Visualization
#-----------------------------------------------------------------------
# convergence history
fig = plt.figure()
ax = plt.axes()
ax.loglog(1./hlist,errlist,label='error',marker='x')
ax.loglog(1./hlist,errlist[0]/hlist[0]*hlist**2,label='slope -2',linestyle='dotted')
ax.set_xlabel('1/h')
ax.set_title('max norm error u_{exact}-u_{multigrid}')
ax.legend()
plt.tight_layout()
plt.show()
