from scipy import constants as c
from matplotlib import pyplot as plt
from math import factorial, pi
from itertools import permutations
#from goto import goto, label
from scipy.linalg import eigh, det, inv
from scipy.optimize import brentq

import os.path

import random

import numpy as np
#========================================================================================================
def alpha_param_initialize_all(Uinv_matrix) :

    number_of_alpha_params = int((npart+1)*npart/2)
    alpha_params           = np.zeros(number_of_alpha_params)
    
    check = False
    
    while not check :
    
        for i in range(0,number_of_alpha_params) :
        
            xxx = random.random()

            alpha_params[i] = np.exp( np.log(bmin) + xxx * ( np.log(bmax)-np.log(bmin) ))
            
        check = pos_def_check(Acalc(alpha_params,Uinv_matrix),npart-1)
        
    return(alpha_params)
#========================================================================================================
def A_matrix_from_alpha_params_change_one(alpha_params,n_nl_param,Uinv_matrix) :
    
    check = False
    
    while not check :
    
        xxx = random.random()

        alpha_params[n_nl_param] = np.exp( np.log(bmin) + xxx * ( np.log(bmax)-np.log(bmin) ))
    
        A_matrix = Acalc(alpha_params,Uinv_matrix)    
        check = pos_def_check(A_matrix,npart-1)
        
    return(A_matrix)
#========================================================================================================
def Acalc(alpha_params, Uinv_matrix) :

    k  = 0
    Xn = np.zeros((npart,npart))
    G  = np.zeros((npart,npart))
    
    for i in range(0,npart) :
        for j in range(i+1,npart) :

            Xn[i,j] = np.abs( 1. / ( alpha_params[k]**2 ) )
            Xn[j,i] = Xn[i,j]
            k+=1
        
        Xn[i,i]=0

    #print(Xn)
    for i in range(0,npart) :
        for j in range(i+1,npart) :
            G[i,j] = -Xn[i,j]
            G[j,i] = G[i,j]
        
        ss=0

        for k in range(0,npart) :
            ss += Xn[i,k]
       
        G[i,i] = ss

    #Transpose[Uinv].G.Uinv
    A_matrix_aux = np.transpose(Uinv_matrix).dot(G).dot(Uinv_matrix)
    
    #neglecting center of mass component
    A_matrix = np.zeros((npart-1,npart-1))

    for i in range(0,npart-1) :
        for j in range(0,npart-1) :
        
            A_matrix[i,j] = A_matrix_aux[i,j]

    return(A_matrix)

#========================================================================================================
def pos_def_check(matrix, dim) :

    p = np.zeros(dim)
    matrix_aux = np.zeros((dim,dim))

    for i in range(0,dim) :
        for j in range(0,dim) :
            matrix_aux[i,j]=matrix[i,j]

    for i in range(0,dim) :
        for j in range(i,dim) :
        
            suma = matrix_aux[i,j]

            for k in reversed(range(0,i)) :
                suma-=matrix_aux[i,k]*matrix_aux[j,k]

            if i==j :
                if suma <= 0 :
                    return(False)
          
                p[i] = np.sqrt(suma)
            
            else :
                matrix_aux[j,i] = suma/p[i]

    return(True)
#========================================================================================================
def lin_dep_check(n_matrix, n) :

    for m in range(0,n) :
    
        aux = np.abs(n_matrix[m,n] / np.sqrt (n_matrix[m,m]*n_matrix[n,n] ) )

        if aux < 0.00001 and n+1 > 100 :
            return(False)

        if aux > 0.99 :
            return(False) 
        
    return(True)
#========================================================================================================
def generalized_eigenvalue_problem_full(h_matrix, n_matrix, dim) :

    h_matrix_small = np.zeros((dim,dim))
    n_matrix_small = np.zeros((dim,dim))
    
    for i in range(0,dim) :
        for j in range(0,dim) :
            h_matrix_small[i,j] = h_matrix[i,j]
            n_matrix_small[i,j] = n_matrix[i,j]
            
    eigvals, eigvecs = eigh(h_matrix_small, n_matrix_small, eigvals_only=False)
    
    return(eigvals,eigvecs)
#========================================================================================================
def spect_pol(e,a,q,eigvals) :
    
    result=a-e
    
    for i in range(0,len(q)):
        result-=q[i]**2/(eigvals[i]-e)
    
    return(result)    
    
#========================================================================================================
def energy_spect_pol(n_matrix, h_matrix, n, eigvals, eigvecs) :

    if n==0 :
        return(h_matrix[n,n]/n_matrix[n,n]) 


    #overlap of the new basis state with the previously selected orthogonalized
    #basis
    psi = np.zeros(n)
    
    for i in range(0,n) :
        for j in range(0,n) :
            psi[i]+=eigvecs[j,i]*n_matrix[j,n]
        
    #norm of the new basis state
    norm = n_matrix[n,n]
    for i in range(0,n) :
        norm-=psi[i]**2
    norm = np.sqrt(norm)
   
    #matrix element of the Hamiltonian of the new basis state with the previously selected
    #orthogonalized basis
    hpsi = np.zeros(n)
    
    for i in range(0,n) :
        for j in range(0,n) :
            hpsi[i]+=eigvecs[j,i]*h_matrix[j,n]
    
    #q coefficient
    q = np.zeros(n)
    q[:] = (hpsi[:]-eigvals[:]*psi[:])/norm
    
    #a coefficient
    a = h_matrix[n,n]
    for i in range(0,n) :
        a += eigvals[i]*psi[i]**2 - 2.*hpsi[i]*psi[i]
    a=a/norm**2
    
    eps = 0.1**12
    
    if opt_ener_level==0 or opt_ener_level+1 > n :
        if n == 1 :
            brac=0.05
            e1 = eigvals[0] - brac
            e2 = eigvals[0] - eps
        else :
            brac = np.abs(eigvals[n-1]-eigvals[n-2])/1000.
            e1   = eigvals[0] - brac
            e2   = eigvals[0] - eps
  
        while spect_pol(e1,a,q,eigvals)*spect_pol(e2,a,q,eigvals) > 0 :
            brac = 2.*brac
            e1   = eigvals[0] - brac
            if brac > 1000. :
                return(0)
    else :
        e1 = eigvals[opt_ener_level-1] + eps
        e2 = eigvals[opt_ener_level] - eps
       
    energy, status = brentq(spect_pol, e1, e2, args=(a,q,eigvals), xtol=2e-12, rtol=8.881784197001252e-16, maxiter=100, full_output=True, disp=False)

    if status.converged == False :
        return(0)

    return(energy)
#========================================================================================================
def add_basis_state_SVM(h_matrix, n_matrix, wave_funct_A_params, permut_matrices, spiso_elem_id, \
                        spiso_elem_two_body, Uinv_matrix, lmbd_matrix, W_vectors, npermut, eigvals, eigvecs,  n) :
                        
    A_best = np.zeros((npart-1,npart-1))
    
    h_matrix_elements_best = np.zeros(n+1) 
    n_matrix_elements_best = np.zeros(n+1)
    
    check = False
    energy_lowest  = 100000000.

    #stochastic optimization of correlated gausssian basis state with randomly chosen A
    
    #random initial parameters alpha
    number_of_alpha_params = int((npart+1)*npart/2)
    alpha_params   = alpha_param_initialize_all(Uinv_matrix)


    for mm in range(0,mm0) :
       
        for n_nl_param in range(0,number_of_alpha_params) :

            for kk in range(0,kk0) :
                             
                cnt=0
                linear_dependency_check = False
                
                while not linear_dependency_check :

                    A_matrix = A_matrix_from_alpha_params_change_one(alpha_params,n_nl_param,Uinv_matrix)
                    
                    wave_funct_A_params[n,:,:] = A_matrix[:,:]
                                
                    add_elements(h_matrix, n_matrix, wave_funct_A_params, permut_matrices, lmbd_matrix, W_vectors, \
                             spiso_elem_id, spiso_elem_two_body, npermut, n)

                    cnt+=1
                    
                    linear_dependency_check = lin_dep_check(n_matrix,n)
                        
                energy = energy_spect_pol(n_matrix, h_matrix, n, eigvals, eigvecs)

                if energy < energy_lowest :
                    check = True
                    energy_lowest  = energy
                    
                    #saving so far the best variational parameters - A matrix
                    A_best[:,:] = wave_funct_A_params[n,:,:]

                    #saving matrix elements of the best parameter selection up to now
                    for i in range(0,n+1) :
                        h_matrix_elements_best[i] = h_matrix[n,i] 
                        n_matrix_elements_best[i] = n_matrix[n,i]
                    

    #saving the best parameters acquired for the n-th basis state using SVM procedure
    wave_funct_A_params[n,:,:] = A_best[:,:]

    for i in range(0,n+1) :
        h_matrix[n,i] = h_matrix_elements_best[i]
        n_matrix[n,i] = n_matrix_elements_best[i]
        h_matrix[i,n] = h_matrix_elements_best[i]
        n_matrix[i,n] = n_matrix_elements_best[i]

    eigvals, eigvecs = generalized_eigenvalue_problem_full(h_matrix, n_matrix, n+1)
                   
    return(eigvals, eigvecs)                        

#========================================================================================================
def svm_run() :

    npairs = int(npart*(npart-1)/2)

    #=========================Jacobi coordinates===================================
    lmbd_matrix   = np.zeros(npart-1)
    U_matrix      = np.zeros((npart,npart))
    Uinv_matrix   = np.zeros((npart,npart))
    W_vectors     = np.zeros((npairs,npart-1))

    init_jacobi(lmbd_matrix, U_matrix, Uinv_matrix, W_vectors)
    
    
    #======================Permutation matrices etc.================================
    npermut          = factorial(npart)   
    permut_particles = np.zeros((npermut,npart), dtype=np.int8)
    permut_signs     = np.zeros(npermut, dtype=np.int8)
    permut_matrices   = np.zeros((npermut,npart-1,npart-1))
    
    init_permutations(npermut, permut_particles, permut_signs, permut_matrices, U_matrix, Uinv_matrix)
    
    
    #====================Spin-isospin matrix elements===============================
    spiso_elem_id       = np.zeros(npermut)
    spiso_elem_two_body = np.zeros((5,npermut,npairs))
    
    init_spin_isospin_elements(spiso_elem_id, spiso_elem_two_body, npermut, permut_particles, permut_signs)
    
    #====Running step-by-step stochastic selection of correlated Gaussian basis====
    h_matrix = np.zeros((mnb,mnb)) # Hamiltonian matrix
    n_matrix = np.zeros((mnb,mnb)) # Overlap matrix
    
    wave_funct_A_params = np.zeros((mnb,npart-1,npart-1))
    
    
        #===================Checking for already calculated basis states===========
    conv_file_exists   = os.path.exists(name+".conv")
    A_mats_file_exists = os.path.exists(name+".res")
    
        #===================There are already calculated basis states===========
    if conv_file_exists and A_mats_file_exists :
    
        #loading parameters of already calculated basis states
        A_matrices_load = np.loadtxt(name+".res",float)
        
        number_of_calc_basis_states = len(A_matrices_load[:,0]) 
        
        if number_of_calc_basis_states > mnb :
           line="Number of precalculated basis states is larger than maximum number of basis states." 
           print(line)
           exit(0)
           
        for n in range(0,number_of_calc_basis_states) :
            for i in range(0,npart-1) :
                for j in range(0,npart-1) :
                    wave_funct_A_params[n,i,j] = A_matrices_load[n,1+i*(npart-1)+j]
        
        del A_matrices_load

        #calculating matrix elements corresponing to already calculated basis states
        for n in range(0,number_of_calc_basis_states) :
            add_elements(h_matrix, n_matrix, wave_funct_A_params, permut_matrices, lmbd_matrix, W_vectors, \
                             spiso_elem_id, spiso_elem_two_body, npermut, n)
                             
        eigvals, eigvecs =generalized_eigenvalue_problem_full(h_matrix, n_matrix, number_of_calc_basis_states)
        
        print("Found "+str(number_of_calc_basis_states)+" precalculated basis states giving energy " \
              +"{0:.8f}".format(eigvals[0]))
        
        #Starting SVM selection for remaining (mnb - number_of_prec_basis_states) basis states
        for n in range(number_of_calc_basis_states,mnb) :
        
            output_convergence   = open(name+".conv", "a")
            output_cg_A_matrices = open(name+".res", "a")
        
            eigvals, eigvecs = add_basis_state_SVM(h_matrix, n_matrix, wave_funct_A_params, permut_matrices, spiso_elem_id, \
            spiso_elem_two_body, Uinv_matrix, lmbd_matrix, W_vectors, npermut, eigvals, eigvecs, n)
            
            if opt_ener_level>n :
                state = n
            else :
                state = opt_ener_level
                
            line = str(n+1)
            for i in range(0,state+1) :
                line+=" \t " + "{0:.8f}".format(eigvals[i])                         
            print(line)
            
            #save to file bound state enegy
            output_convergence.write(line+"\n") 
            
            #save to file parameters of n-th corr. Gaussian A matrix
            line=str(n)
            
            for i in range(0,npart-1) :
                for j in range(0,npart-1) :
                    line += "\t" + "{0:.12e}".format(wave_funct_A_params[n,i,j]) 
                
            output_cg_A_matrices.write(line+"\n")
            
            output_convergence.close()
            output_cg_A_matrices.close()
    
       #===============Starting SVM selection of basis states from scratch=========
    else :
     
        eigvals = []
        eigvecs = []
     
        for n in range(0,mnb) :
        
            output_convergence   = open(name+".conv", "a")
            output_cg_A_matrices = open(name+".res", "a")
        
            eigvals, eigvecs = add_basis_state_SVM(h_matrix, n_matrix, wave_funct_A_params, permut_matrices, spiso_elem_id,
                                  spiso_elem_two_body, Uinv_matrix, lmbd_matrix, W_vectors, npermut, eigvals, eigvecs, n)

            if opt_ener_level>n :
                state = n
            else :
                state = opt_ener_level
            
            line = str(n+1)
            for i in range(0,state+1) :
                line+=" \t " + "{0:.8f}".format(eigvals[i])                         
            print(line)
            
            #save to file bound state enegy
            output_convergence.write(line+"\n") 
            
            #save to file parameters of n-th corr. Gaussian A matrix
            line=str(n)
            
            for i in range(0,npart-1) :
                for j in range(0,npart-1) :
                    line += "\t" + "{0:.12e}".format(wave_funct_A_params[n,i,j]) 
                
            output_cg_A_matrices.write(line+"\n")
            
            output_convergence.close()
            output_cg_A_matrices.close()
        

    
#========================================================================================================
def add_elements(h_matrix,n_matrix,wave_funct_A_params,permut_matrices, lmbd_matrix, W_vectors, \
                 spiso_elem_two_id, spiso_elem_two_body, npermut, n) :


    A_matrix     = np.zeros((npart-1,npart-1))
    A_matrix_aux = np.zeros((npart-1,npart-1))
    Perm_matrix  = np.zeros((npart-1,npart-1))

    h_matrix[n,:] = 0
    n_matrix[n,:] = 0

    #loop over different permutations
    for permut in range(0,npermut) :

        A_matrix_aux[:,:] = wave_funct_A_params[n,:,:]
        Perm_matrix[:,:]  = permut_matrices[permut,:,:]
        
        Aperm_matrix      = np.transpose(Perm_matrix).dot(A_matrix_aux).dot(Perm_matrix)

        for basis_state in range(0,n+1) :
        
            A_matrix[:,:] = wave_funct_A_params[basis_state,:,:]
            
            B_matrix      = np.add(A_matrix,Aperm_matrix)
            detB15          = (det(B_matrix))**(-1.5)
            Binv          = inv(B_matrix)

            
            #Overlap matrix element
            n_matrix[n,basis_state] += overlap_elem(detB15)*spiso_elem_two_id[permut]            
            
            #Calcuate Trace for kinetic energy
            trace = np.trace( Binv.dot(Aperm_matrix).dot( np.diag(lmbd_matrix) ).dot(A_matrix) )
            
            #Kinetic energy matrix element
            h_matrix[n,basis_state] += kin_ener_elem(trace,detB15)*spiso_elem_two_id[permut]


            
            #Two-body interaction
            for pair in range(0,int(npart*(npart-1)/2)) :
            
                p = ( W_vectors[pair,:].dot(Binv).dot(W_vectors[pair,:]) )
                
                for operator_term in range(0,4) :
                
                    for term in range(0,npt) :
                    
                        h_matrix[n,basis_state] += gauss_pot_elem(vpot[operator_term,term],apot[operator_term,term],p,detB15) \
                                         * spiso_elem_two_body[operator_term,permut,pair]
                
                if Coulomb_pp == True :
                
                    h_matrix[n,basis_state] += coulomb_elem(p,detB15) * spiso_elem_two_body[4,permut,pair]
                                         
                if HO_trap == True :

                    h_matrix[n,basis_state] += ho_elem(p,detB15) * spiso_elem_two_body[0,permut,pair]
        
    #Symmetric matrices
    for i in range(0,n+1) :
        h_matrix[i,n] = h_matrix[n,i]
        n_matrix[i,n] = n_matrix[n,i]

#========================================================================================================
def overlap_elem(detB15) :

    ##########################################
    #                                        #
    #                COMPLETE                #
    #                                        #
    ##########################################

    return( )
#========================================================================================================
def kin_ener_elem(trace,detB15) :
    return(1.5 * trace * detB15)
#========================================================================================================
def gauss_pot_elem(vpot,apot,p,detB15) :

    ##########################################
    #                                        #
    #                COMPLETE                #
    #                                        #
    ##########################################

    return( )
#========================================================================================================
def coulomb_elem(p,detB15) :

    ##########################################
    #                                        #
    #                COMPLETE                #
    #                                        #
    ##########################################

    return( )
#========================================================================================================
def ho_elem(p,detB15) :
    return(2.*hh/(npart*HO_trap_length**4)*3.*p*detB15)
#========================================================================================================
def init_jacobi(lmbd_matrix, U_matrix, Uinv_matrix, W_vectors) :

    #===================================================
    # Calculation of kinetic energy Lambda diagonal matrix
    aux1 = 0
    aux2 = mass[0]
    
    for i in range(0,npart-1) :

        aux1 += mass[i]
        aux2 += mass[i+1]
        
        lmbd_matrix[i]=aux2/aux1
        
    for i in range(0,npart-1) :
        lmbd_matrix[i]*=hh/mass[i+1]
    #===================================================
    # Calculation of Jacobi transformation matrix U 
    aux1 = 0
    aux2 = mass[0]
    
    for i in range(0,npart-1) :

        aux1+=mass[i]
        aux2+=mass[i+1]


        for j in range(0,i+1) :
            U_matrix[i,j]=-mass[j]/aux1

        U_matrix[i,i+1]=1.

    for i in range(0,npart) :
        U_matrix[npart-1,i] = mass[i]/aux2
        
    #===================================================
    #Calculation of inverse of Jacobi transformation matrix U^-1
    Uinverse_aux = inv(U_matrix)
    Uinv_matrix[:,:] = Uinverse_aux[:,:]    
    #===================================================
    #Calculation of auxiliary w vector (needed for two-body potential matrix elements)
    pair=0;
    for i in range(0, npart) :
        for j in range(i+1, npart) :
    
            for k in range(0, npart-1) :
                W_vectors[pair,k]=Uinv_matrix[i,k] - Uinv_matrix[j,k]

            pair+=1
#========================================================================================================
def init_permutations(npermut, permut_particles, permut_signs, permut_matrices, U_matrix, Uinv_matrix) :

    #====================================== Particle permutations =======================================

    for permut in range(0,npermut) :
    
        iv        = np.zeros(npart+1, dtype=np.int8 )
        perm_aux  = np.zeros(npart+1, dtype=np.int8 )
        
        ipx=permut+1;

        for i in range(1,npart+1) :
            iv[i]=i

        i1=ipx-1;
        for m in reversed(range(1,npart)) :
            i2=int(i1/factorial(m))+1
            i1=i1%factorial(m)
            perm_aux[npart-m]=iv[i2]

            for i in range(i2,m+1) :
                iv[i]=iv[i+1]

        perm_aux[npart]=iv[1];

        for i in range(0,npart) :
            permut_particles[permut,i]=perm_aux[i+1]-1

    #====================================== Signs of permutations =======================================
    #( note that for bosons all signs are automatically +; symmetrization )
    
    if ibf==1 :
        #bosons
        for permut in range(0,npermut) :
            permut_signs[permut]=1
    
    else :
        #fermions
        int_aux = np.zeros(npart+1, dtype=np.int8)
        
        for permut in range(0,npermut) :
        
            for i in range(0,npart) :
                int_aux[i+1]=permut_particles[permut,i]+1
        
            ii=0
            for i in range(1,npart+1) :
                for j in range(1,npart+1) :

                    if i==int_aux[j]  and  i!=j :
                        mm         = int_aux[i]
                        int_aux[i] = int_aux[j]
                        int_aux[j] = mm
                        ii+=1

            permut_signs[permut] = (-1)**ii
    
    #====================================== Permutation matrices ========================================

    C             = np.zeros((npart,npart))
    Aux_matrix    = np.zeros((npart,npart))

    for permut in range(0,npermut) :

        #Permutation matrix C
        for k in range(0,npart) :
            for j in range(0,npart) :

                if j==permut_particles[permut,k] :
                    C[j,k]=1.
                else :
                    C[j,k]=0

        #Permutation matrix U.C.U^-1
        Aux_matrix = U_matrix.dot(C).dot(Uinv_matrix)
        
        for i in range(0,npart-1) :
            for j in range(0,npart-1) :
                permut_matrices[permut,i,j]=Aux_matrix[i,j]

#========================================================================================================
def init_spin_isospin_elements(spiso_elem_id, spiso_elem_two_body, npermut, permut_particles, permut_signs) :


    #=========================== Spin-isospin matrix elements identity ==================================   
    for permut in range(0,npermut) :

        #isospin part
        isospin_elem_id=0
        for isos_term1 in range(0,nisc) :
            for isos_term2 in range(0,nisc) :
            
                aux=cisc[isos_term1]*cisc[isos_term2]

                for i in range(0,npart) :

                    if iso[isos_term1,permut_particles[permut,i]] != iso[isos_term2,i] :
                        aux*=0
                    else :
                        aux*=1.
                        
                isospin_elem_id+=aux
        
        #spin part
        spin_elem_id=0
        for spin_term1 in range(0,nspc) :
            for spin_term2 in range(0,nspc) :
            
                aux=cspc[spin_term1]*cspc[spin_term2]

                for i in range(0,npart) :

                    if isp[spin_term1,permut_particles[permut,i]] != isp[spin_term2,i] :
                        aux*=0
                    else :
                        aux*=1.
                        
                spin_elem_id+=aux
                        
        spiso_elem_id[permut] = isospin_elem_id * spin_elem_id * permut_signs[permut]
    
    #========================== Two-body spin-isospin matrix elements ==================================== 
    npairs = int(npart*(npart-1)/2)
    
    #Two-body spin-isospin matrix elements WIGNER
    for permut in range(0,npermut) :
       
        pair=0
        
        for part1 in range(0,npart) :
            for part2 in range(part1+1,npart) :
        
                isospin_elem_id_two_body=0
                isospin_elem_Pt_two_body=0
                isospin_elem_pp_two_body=0
            
                for isos_term1 in range(0,nisc) :
                    for isos_term2 in range(0,nisc) :
            
                        aux=cisc[isos_term1]*cisc[isos_term2]
                
                        for part in range(0,npart) :
                            if part==part1 or part==part2 :
                                aux*=1.
                            else :                   
                                if iso[isos_term1,permut_particles[permut,part]] != iso[isos_term2,part] :
                                    aux*=0
                                else :
                                    aux*=1.
                                    
                        #Isospin identity id
                        if (iso[isos_term1,permut_particles[permut,part1]] == iso[isos_term2,part1]) and \
                           (iso[isos_term1,permut_particles[permut,part2]] == iso[isos_term2,part2]) :
                           
                            isospin_elem_id_two_body+=aux            
                            
                        #Isospin exchange Pt
                        if (iso[isos_term1,permut_particles[permut,part1]] == iso[isos_term2,part2]) and \
                           (iso[isos_term1,permut_particles[permut,part2]] == iso[isos_term2,part1]) :
                           
                            isospin_elem_Pt_two_body+=aux            
                            
                        #proton-proton pair (Coulomb)
                        if (iso[isos_term1,permut_particles[permut,part1]] == 1) and (iso[isos_term2,part1]==1) and \
                           (iso[isos_term1,permut_particles[permut,part2]] == 1) and (iso[isos_term2,part2]==1) :
                            isospin_elem_pp_two_body+=aux            
            
                spin_elem_id_two_body=0
                spin_elem_Ps_two_body=0

                for spin_term1 in range(0,nspc) :
                    for spin_term2 in range(0,nspc) :
            
                        aux=cspc[spin_term1]*cspc[spin_term2]
                
                        for part in range(0,npart) :
                            if part==part1 or part==part2 :
                                aux*=1.
                            else :                   
                                if isp[spin_term1,permut_particles[permut,part]] != isp[spin_term2,part] :
                                    aux*=0
                                else :
                                    aux*=1.
                                    
                        #Isospin identity id
                        if (isp[spin_term1,permut_particles[permut,part1]] == isp[spin_term2,part1]) and \
                           (isp[spin_term1,permut_particles[permut,part2]] == isp[spin_term2,part2]) :
                           
                            spin_elem_id_two_body+=aux            
                            
                        #Isospin exchange Pt
                        if (isp[spin_term1,permut_particles[permut,part1]] == isp[spin_term2,part2]) and \
                           (isp[spin_term1,permut_particles[permut,part2]] == isp[spin_term2,part1]) :
                           
                            spin_elem_Ps_two_body+=aux             
        

                #Two-body spin-isospin matrix elements WIGNER        
                spiso_elem_two_body[0,permut,pair] = isospin_elem_id_two_body * spin_elem_id_two_body * permut_signs[permut]
            
                #Two-body spin-isospin matrix elements MAJORANA
                spiso_elem_two_body[1,permut,pair] = -isospin_elem_Pt_two_body * spin_elem_Ps_two_body * permut_signs[permut]
            
                #Two-body spin-isospin matrix elements BARTLETT
                spiso_elem_two_body[2,permut,pair] = isospin_elem_id_two_body * spin_elem_Ps_two_body * permut_signs[permut]
            
                #Two-body spin-isospin matrix elements HAISENBERG
                spiso_elem_two_body[3,permut,pair] = -isospin_elem_Pt_two_body * spin_elem_id_two_body * permut_signs[permut]
            
                #Two-body spin-isospin matrix elements COULOMB
                spiso_elem_two_body[4,permut,pair] = isospin_elem_pp_two_body * spin_elem_id_two_body * permut_signs[permut]

                pair+=1
#========================================================================================================


#########################################################################################################
#                                                                                                       #
#       PYTHON STOCHASTIC VARIATIONAL METHOD SCRIPT (M. Schafer; July 2022; TALENT School@MITP)         #
#                                                                                                       #
#########################################################################################################


######################################### INPUT PARAMETERS ##############################################

name = "3H_Minnesota"

#number of particles
npart=3

# masses of particles
mass = np.zeros(npart)

mass[0] = 1.
mass[1] = 1.
mass[2] = 1.        

#================ isospin configurations ================
nisc = 2

cisc = np.zeros(nisc)
iso  = np.zeros((nisc,npart))

#iso=2 (neutron), iso=1 (proton)

#cisc[0]=1.

#iso[0,0]=1
#iso[0,1]=2

#cisc[1]=-1.

#iso[1,0]=2
#iso[1,1]=1


cisc[0]=1.


iso[0,0]=1
iso[0,1]=1
iso[0,2]=2

cisc[1]=-1.

iso[1,0]=1
iso[1,1]=2
iso[1,2]=1
#================== spin configurations =================
nspc = 2

cspc = np.zeros(nspc)
isp  = np.zeros((nspc,npart))

#cspc[0]=1.

#isp[0,0]=2
#isp[0,1]=2
#isp[0,2]=2


cspc[0]=1.

isp[0,0]=1
isp[0,1]=2
isp[0,2]=2

cspc[1]=-1.

isp[1,0]=2
isp[1,1]=1
isp[1,2]=2
#==================== other parameters ==================
hh=41.47     # (hbarc)^2/m_N 

ibf=2        # ibf=1 (bosons), ibf=2 (fermions)

mm0=2        # optimization parameter (alpha matrix loop)
kk0=10       # optimization parameter (individual element loop)

mnb=150      # maximum number of basis states

bmin=0.1     # lower limit on nonlinear correlated gaussian parameters (elements of A matrix) 
bmax=15.     # upper limit

opt_ener_level = 0 #state with respect to which SVM should optimize the correl. gaussian basis
                   #0 is the ground state, 1 is the first excited state, etc. 

#Harmonic oscillator (HO) trap
HO_trap = False
HO_trap_length = 10.0

#Coulomb interaction between protons
Coulomb_pp = False

#================== POTENTIAL INPUT MINNESOTA======================

npt=3   #number of potential terms 
   
  
vpot = np.zeros ((4,npt))
apot = np.zeros ((4,npt))

# ========== Wigner terms ==========
vpot[0,0]= 100.000
apot[0,0]=1.487 

vpot[0,1]= -44.50
apot[0,1]=0.639

vpot[0,2]= -22.9625
apot[0,2]=0.465

# ========== Majorana terms ==========
vpot[1,0]= 100.000 
apot[1,0]=1.487

vpot[1,1]= -44.50
apot[1,1]=0.639

vpot[1,2]= -22.9625
apot[1,2]=0.465

# ========== Bartlett terms ==========
vpot[2,0]=   0.000
apot[2,0]=1.487

vpot[2,1]= -44.50
apot[2,1]=0.639

vpot[2,2]= +22.9625
apot[2,2]=0.465

# ========== Haisenberg terms ==========
vpot[3,0]=   0.000
apot[3,0]=1.487

vpot[3,1]= -44.50   
apot[3,1]=0.639 

vpot[3,2]= +22.9625
apot[3,2]=0.465

#===================POTENTIAL INPUT VOLKOV======================
#npt=2   #number of potential terms 
   
  
#vpot = np.zeros ((4,npt))
#apot = np.zeros ((4,npt))

# Wigner terms
#vpot[0,0] = 144.86
#apot[0,0] = 1./0.82**2

#vpot[0,1] = -83.34
#apot[0,1] = 1./1.60**2

#########################################################################################################

# Running Stochastic Variational Method 
svm_run()
