import nlopt import numpy as np import shelve import os import time import nlopt from functools import partial,update_wrapper import jax from jax import value_and_grad as value_and_grad from jax import jit import scipy import jax.numpy as jnp import scipy import scipy.sparse.linalg as spla import scipy.sparse as sp from jax import custom_vjp from typing import Callable import matinverse as mi jax.config.update('jax_platform_name', 'gpu') jax.config.update("jax_log_compiles",0) jax.config.update("jax_enable_x64",1) @jax.jit def compute_kappa_and_gradient(x,data,dir_dep): ki,kj,P_adj = dir_dep #[gm_direct,gp_direct,RHSB,t_flat,d,_,_,_,_,dt_sparse,dt_sparseT,i_mat,j_mat,gm,gp,D,GG,_,_,_,_,RR,reflection,f] = data [gm_direct,gp_direct,RHSB,t_flat,d,_,_,dt_sparse,dt_sparseT,i_mat,j_mat,gm,gp,D,GG,_,RR,reflection,f] = data (nq,n_elems) = P_adj.shape T_mat = x.reshape((nq,n_elems)) #reflection G_mat = jnp.zeros_like(T_mat) for i in range(nq): G_mat = G_mat.at[i].set(-T_mat[reflection[i]]) #------ #Compute kappa-------- kappa = -jnp.einsum('n,qn ->',t_flat[ki],gm[:,ki],optimize=True) kappa += jnp.einsum('uc,uc->',P_adj,T_mat) #--------------------- #Perturbation--- g = jnp.zeros(n_elems) tmp = jnp.einsum('qi,qi->i',G_mat[:,i_mat[ki]],gm[:,ki]) g = g.at[i_mat[ki]].add(-dt_sparse[ki]*tmp) g = g.at[j_mat[ki]].add(-dt_sparseT[ki]*tmp) tmp = -jnp.einsum('qi,qi->i',G_mat[:,i_mat[kj]],gm[:,kj]) g = g.at[i_mat[kj]].add(-dt_sparse[kj]*tmp) g = g.at[j_mat[kj]].add(-dt_sparseT[kj]*tmp) #---------------- #Matrix tmp = jnp.einsum('qk,qk,qk->k',gm,T_mat[:,j_mat],G_mat[:,i_mat]) g = g.at[i_mat].add(-tmp*dt_sparse) g = g.at[j_mat].add(-tmp*dt_sparseT) #Boundary tmp = jnp.einsum('kqu,uk,qk->k',RR,T_mat[:,i_mat],G_mat[:,i_mat]) g = g.at[i_mat].add(-tmp*dt_sparse) g = g.at[j_mat].add(-tmp*dt_sparseT) #From kappa tmp = gm[:,ki].sum(axis=0) + jnp.einsum('uk,uk->k',T_mat[:,i_mat[ki]],gp[:,ki]) + jnp.einsum('uk,uk->k',T_mat[:,j_mat[ki]],gm[:,ki]) g = g.at[i_mat[ki]].add(-dt_sparse[ki] *tmp) g = g.at[j_mat[ki]].add(-dt_sparseT[ki]*tmp) return [kappa,g] def get_common(**options): #Parse-- directions =options['directions'] n_dir,dim = np.array(directions).shape grid = options['grid'] if dim == 3: N = int(grid**3) factor = 1/grid n_theta = options['n_theta'] else: N = int(grid**2) factor = 1 aux = mi.get_grid(grid,dim) i_mat = aux['i'] j_mat = aux['j'] ind_extremes = aux['ind_extremes'] n_phi = options['n_phi'] Knt = options['Knt'] Kn = Knt*grid #--------------- #Momentum space-- Dphi = 2*np.pi/n_phi phi = np.linspace(0,2.0*np.pi-Dphi,n_phi,endpoint=True) if dim == 2: polar = jnp.array([np.cos(phi),np.sin(phi)]).T #Remember to revert back fphi = np.sinc(Dphi/2.0/np.pi) S = polar*fphi else: #-------------------------- S = jnp.zeros((3,n_phi,n_theta)) Dtheta = jnp.pi/n_theta theta = jnp.linspace(Dtheta/2,jnp.pi-Dtheta/2,n_theta,endpoint=True) tmp = (Dtheta-jnp.cos(2*theta)*jnp.sin(Dtheta))*jnp.sin(Dphi/2) Sz = Dphi*jnp.sin(theta)*jnp.cos(theta) S = S.at[0].set(jnp.outer(jnp.cos(phi),tmp)) S = S.at[1].set(jnp.outer(jnp.sin(phi),tmp)) S = S.at[2].set(Sz[jnp.newaxis,:]) DeltaOmega = 2*Dphi*jnp.sin(theta)*jnp.sin(Dtheta/2) #correct S = jnp.einsum('ijk,k->ijk',S,1/DeltaOmega) S = S.reshape((n_theta*n_phi,3)) #---------------------------- M = len(S) #New---------------------------- G = Kn*jnp.einsum('qj,nj->qn',S,aux['normals'],optimize=True) gp = G.clip(min=0); gm = G.clip(max=0) GG = jnp.einsum('qn,n->qn',gp,1/gm.sum(axis=0)) D = jnp.zeros((N,M)).at[aux['i']].add(gp.T).T RR = jnp.einsum('uk,vk->kuv',gm,GG) #------------------------------- #print(S) reflection = np.zeros(M) for a,sa1 in enumerate(S): sa1 /= np.linalg.norm(sa1) found = False for b,sa2 in enumerate(S): sa2 /= np.linalg.norm(sa1) if abs(np.dot(sa1,sa2) + 1) < 1e-3: reflection[a] = int(b) found = True break if not found: print('No reflection found') quit() f = Kn**2*M/2 #----------------------------------- reflection = jnp.array(reflection,int) return [aux['i'],aux['j'],gm,gp,D,GG,aux['ind_extremes'],RR,reflection,f,(N,M)] @jit def sparse_dense_product_jax(i,j,data,X): return jnp.zeros_like(X.T).at[i].add(data.T * X.T[j]).T @partial(jax.jit) def get_rho_dependent(rho,aux2): [i_mat,j_mat,gm,gp,D,GG,ind_extremes,RR,reflection,f,N] = aux2 dim = len(ind_extremes) k0 = 1e-12;k1 = 1 (nq,n_elems) = D.shape rho = k0 + rho*(k1-k0) t_flat = 2*rho[i_mat] * rho[j_mat]/(rho[i_mat] + rho[j_mat]) dt_sparse = (k1-k0)*0.5*jnp.power(t_flat/rho[i_mat],2) dt_sparseT = (k1-k0)*0.5*jnp.power(t_flat/rho[j_mat],2) gm_direct = jnp.einsum('un,n->un',gm,t_flat) gp_direct = jnp.einsum('un,n->un',gp,t_flat) RHSB = jnp.zeros((n_elems,nq,nq)).at[i_mat].add(jnp.einsum('un,vn->nuv',gm_direct-gm,GG)) - 1/nq d = (D + 1 - RHSB[:,np.arange(nq),np.arange(nq)].T).flatten() P_vec = jnp.zeros((dim,n_elems,nq)) P_adj_vec = jnp.zeros((dim,n_elems,nq)) for i in range(dim): ii_1 = ind_extremes[i,1] ii_0 = ind_extremes[i,0] P_vec = P_vec.at[(i,i_mat[ii_1])].add((-gm_direct[:,ii_1]).T ) P_vec = P_vec.at[(i,i_mat[ii_0])].add(( gm_direct[:,ii_0]).T ) P_adj_vec = P_adj_vec.at[(i,i_mat[ii_1])].add((-gp_direct[:,ii_1]).T ) P_adj_vec = P_adj_vec.at[(i,i_mat[ii_0])].add(( gp_direct[:,ii_0]).T ) return gm_direct,gp_direct,RHSB,t_flat,d,P_adj_vec,P_vec,dt_sparse,dt_sparseT,i_mat,j_mat,gm,gp,D,GG,ind_extremes,RR,reflection,f #Preconditioner @jax.jit def PREC(x,d): return x#/d #Operator @jax.jit def L(X,aux3): D,i_mat,j_mat,gm_direct,RHSB = aux3 X = X.reshape(D.shape) return (X + jnp.multiply(D,X) + sparse_dense_product_jax(i_mat,j_mat,gm_direct,X) + jnp.einsum('cuq,qc->uc',RHSB,X)).flatten() def bte(**options)->Callable: common = get_common(**options) directions = options['directions'] n_dir,dim = np.array(directions).shape (N,M) = common[-1] def func(rho): kappa = np.zeros(n_dir) jacobian = np.zeros((n_dir,N)) x0 = np.zeros((n_dir,M*N)) for n,direction in enumerate(directions): aux = get_rho_dependent(rho,common) [gm_direct,gp_direct,RHSB,t_flat,d,P_adj_vec,P_vec,dt_sparse,dt_sparseT,i_mat,j_mat,gm,gp,D,GG,ind_extremes,RR,reflection,f] = aux #For operator aux3 = D,i_mat,j_mat,gm_direct,RHSB P = jnp.einsum('i,iuc->cu',direction,P_vec) P_adj = jnp.einsum('i,iuc->cu',direction,P_adj_vec) if direction == [1,0,0]: kj = ind_extremes[0,0] ki = ind_extremes[0,1] if direction == [0,1,0]: kj = ind_extremes[1,0] ki = ind_extremes[1,1] if direction == [0,0,1]: kj = ind_extremes[2,0] ki = ind_extremes[2,1] if direction==[1,0]: kj = ind_extremes[0,0] ki = ind_extremes[0,1] if direction==[0,1]: kj = ind_extremes[1,0] ki = ind_extremes[1,1] dir_dep = ki,kj,P_adj (kappa[n],jacobian[n]),(T_mat,call_count) = mi.gmres_wrapper(partial(L,aux3=aux3),P.flatten(),x0[n],\ partial(compute_kappa_and_gradient,data=aux,dir_dep=dir_dep),verbose=False,early_termination=True) x0[n] = T_mat #print(func.total_call) func.total_call += call_count #Compute flux (nq,n_elems) = D.shape #flux = jnp.einsum('uc,ui->ci',T_mat.reshape(D.shape),sigma) return (kappa/f,T_mat.reshape((nq,n_elems))),jacobian/f func.total_call = 0 return func