import numpy as np
import scipy as sp 
import matplotlib.pyplot as plt
import optimize

Lx=17.0; Ly=15.0
Nx=71; Ny=37
xn = np.linspace(0.0, Lx, Nx)
yn  = np.linspace(0.0, Ly, Ny)

X, Y = np.meshgrid(xn, yn)

nodes = np.vstack((np.reshape(X, (Nx*Ny,)), np.reshape(Y, (Nx*Ny,)))).T
NoofNodes = nodes.shape[0]
elements = []
elements_lengths = []
for j in range(Ny - 1):
	for i in range(Nx - 1):
		elements.append([j*Nx + i, j*Nx+i+1, (j+1)*Nx+i+1, (j+1)*Nx + i])
		elements_lengths.append([nodes[j*Nx+i+1][0] - nodes[j*Nx+i][0], nodes[ (j+1)*Nx + i][1] - nodes[ (j)*Nx + i][1]])
elements = np.array(elements)
elements_lengths = np.array(elements_lengths)
jacobians = elements_lengths[:,0]*elements_lengths[:,1]/4.0
# for m in range(0, elements.shape[0]):
# 	for ie, e in enumerate(elements[m]):
# 		line = np.array([nodes[e], nodes[elements[m][(ie+1)%4]]])
# 		plt.plot(line[:,0], line[:, 1], '-k', linewidth=0.5)

E0 = 1000
nu=0.3
D = np.array([[1.0, nu, 0.0], \
	[nu, 1.0, 0.0], \
	[0.0, 0.0, 0.5*(1.0 - nu)]])*E0/(1.0 - nu*nu)

def getdNidxi(xi, eta):
	return np.array([-0.25*(1.0-eta), 0.25*(1.0-eta), 0.25*(1.0+eta), -0.25*(1.0+eta)])

def getdNideta(xi, eta):
	return np.array([-0.25*(1.0-xi), -0.25*(1.0 + xi), 0.25*(1.0+xi), 0.25*(1.0 - xi)])

deg = 10
x, w = np.polynomial.legendre.leggauss(deg)
q = 2.5
def getElementStiffness(xi, eta, element_lengths):
	B_e = np.zeros((3,8))
	dNidXi = getdNidxi(xi, eta)
	dNidEta = getdNideta(xi, eta)
	for n in range(0,4):
		B_e[0, 2*n] = dNidXi[n]*2.0/element_lengths[0]
		B_e[1, 2*n+1] = dNidEta[n]*2.0/element_lengths[1]
		B_e[2, 2*n] = dNidEta[n]*2.0/element_lengths[1]
		B_e[2, 2*n+1] = dNidXi[n]*2.0/element_lengths[0]
		
	return np.dot(B_e.T, np.dot(D, B_e))
k0js = []

for ie, e in enumerate(elements):
	ke = np.zeros((8,8))
	for xi, wi in zip(x, w):
		for eta, wj in zip(x, w):
			ke+= wi*wj*\
			getElementStiffness(xi, eta, \
				elements_lengths[ie])*jacobians[ie]
	k0js.append(ke)

F_G = np.zeros((2*NoofNodes, 1))
u_G = np.zeros((2*NoofNodes, 1))
K_G = np.zeros((2*NoofNodes, 2*NoofNodes))

iF = np.where((np.abs(nodes[:,0] - Lx)< 1.0e-8) & (np.abs(nodes[:,1])<1.0e-8))[0]
F_G[2*iF + 1] = -100.0


nodes_mask = np.abs(nodes[:,0])>1.0e-5
nodes_mask = np.repeat(nodes_mask, 2)

def solveSheet(x_in, node_mask):
	K_G[:] = 0.0
	for ie in range(0, elements.shape[0]):
		element = elements[ie]
		for ii, i in enumerate(element):
			for jj, j in enumerate(element):
				K_G[2*i:2*i+2, 2*j:2*j+2]+= \
				np.power(x_in[ie], q)*k0js[ie][2*ii:2*ii+2, 2*jj:2*jj+2]

	KR = K_G[node_mask][:, node_mask]
	FR = F_G[node_mask]
	uR = sp.linalg.solve(KR, FR)
	u_G[nodes_mask] = uR
	C = np.sum(FR*uR) #FR.dot(uR)
	return u_G, C
def getxStar(lam, Lj, q0j, Area, alpha):
	xt = Lj + np.sqrt(q0j/(lam*Area))
	if(xt < alpha):
		return alpha
	elif(xt > xmax):
		return xmax
	else:
		return xt

def getDual(lam, r0, Vmax, q0s, Ljs, alphas, jacobians, Noofels):
	phi = r0 - lam*Vmax
	for m in range(0, Noofels):
		xStar = getxStar(lam, Ljs[m], q0s[m], jacobians[m]*4.0, alphas[m])
		phi += q0s[m]/(xStar - Ljs[m]) + lam*jacobians[m]*4.0*xStar
	return -phi

xd = np.empty([elements.shape[0],])
xd[:] = 1.0
Vmax = 180.0

xk = np.empty([elements.shape[0],])
q0s = np.zeros(elements.shape[0])
Ljs = np.zeros(elements.shape[0])
Ljk_1 = np.zeros(elements.shape[0])

alphas = np.zeros(elements.shape[0])
xk_1 = np.empty(elements.shape[0])
xk_2 = np.empty(elements.shape[0])
maxk = 200
k_opt = 0
xk[:] = xd[:]
xmax = 5.0
xmin = 0.05
s_init = 1.0e-1
s_faster = 1.15
s_slower = 0.75
mu = 0.15
lam_min = 0.0
lam_max = 1000.0
lam_init = 1.0
lam_star = lam_init
volume = np.sum(xk*jacobians)*4.0
print("saving iteration : ", k_opt, " volume = ", volume)
for m in range(0, elements.shape[0]):
	fill_coords = []
	for ie, e in enumerate(elements[m]):
		line = np.array([nodes[e], nodes[elements[m][(ie+1)%4]]])
		plt.plot(line[:,0], line[:, 1], '-k', linewidth=0.5)
		fill_coords.append(nodes[e])
	fill_coords = np.array(fill_coords)
	cval = 1.0 - (xk[m] - xmin)/(xmax - xmin)
	plt.fill(fill_coords[:,0], fill_coords[:,1], color=str(cval))

plt.savefig("./sheetIterationsSIMP/iter="+"%04d"%k_opt+".png")
plt.close()

k_opt+=1
while(k_opt < maxk):
	u_G, C = solveSheet(xk, nodes_mask)
	r0 = C
	if(k_opt < 2):
		Ljs[:] = xk[:] - s_init*(xmax - xmin)
	else:
		for m in range(0, elements.shape[0]):
			if((xk[m] - xk_1[m])*(xk_1[m] - xk_2[m]) > 0.0):
				Ljs[m] = xk[m] - s_faster*(xk_1[m] - Ljk_1[m])
			else:
				Ljs[m] = xk[m] - s_slower*(xk_1[m] - Ljk_1[m])
	for m in range(0, elements.shape[0]):
		alphas[m] = max(xmin, Ljs[m] + mu*(xk[m] - Ljs[m]))
		um = []
		for n in elements[m]:
			um.append(u_G[2*n])
			um.append(u_G[2*n+1])
		um = np.array(um)

		uTk0u = q*np.power(xk[m], q-1.0)*np.dot(um.T, np.dot(k0js[m], um))[0][0]
		q0s[m] = ((xk[m] - Ljs[m])**2)*uTk0u
		r0 += -(xk[m] - Ljs[m])*uTk0u
		# lam_crit = q0s[m]/(lengths[m]*(xmax - Ljs[m])**2)
		# print(m, lam_crit)
	# goldenSection(f, a_in, b_in, tol, maxiter, *args)
	args_list = [ r0, Vmax, q0s, Ljs, alphas, jacobians, elements.shape[0]]
	lambda_opt = optimize.goldenSection(getDual, lam_min, lam_max, 1.0e-5, 1000, *args_list)
	lam_star = lambda_opt
	negphi_k = getDual(lam_star, *args_list)
	xk_2[:] = xk_1[:]
	xk_1[:] = xk[:]

	Ljk_1[:] = Ljs[:]

	for m in range(0, elements.shape[0]):
		xk[m] = getxStar(lam_star, Ljs[m], q0s[m], jacobians[m]*4.0, alphas[m])
	# print(xk)

	volume = np.sum(xk*jacobians)*4.0
	print("saving iteration : ", k_opt, " volume = ", volume)
	for m in range(0, elements.shape[0]):
		fill_coords = []
		for ie, e in enumerate(elements[m]):
			line = np.array([nodes[e], nodes[elements[m][(ie+1)%4]]])
			plt.plot(line[:,0], line[:, 1], '-k', linewidth=0.5)
			fill_coords.append(nodes[e])
		fill_coords = np.array(fill_coords)
		cval = 1.0 - (xk[m] - xmin)/(xmax - xmin)
		plt.fill(fill_coords[:,0], fill_coords[:,1], color=str(cval))

	plt.savefig("./sheetIterationsSIMP/iter="+"%04d"%k_opt+".png")
	plt.close()
	k_opt+=1
