use ndarray::prelude::*;
use num_complex::Complex;
use crate::thermodynamics::{Occupation, fermi_derivative_from_width, fermi_from_width};
use super::quadrature::*;
use super::types::TrackedSimplex;
#[inline]
pub(crate) fn fermi(e: f64, mu: f64, thermal_width: f64) -> f64 {
fermi_from_width(e, mu, thermal_width)
}
#[inline]
#[allow(dead_code)]
pub(crate) fn fermi_deriv(e: f64, mu: f64, thermal_width: f64) -> f64 {
if thermal_width == 0.0 {
0.0
} else {
fermi_derivative_from_width(e, mu, thermal_width)
}
}
pub(crate) fn eval_berry_kernel(
band_q: &[f64],
k_ab_q: &Array2<Complex<f64>>,
eta: f64,
nsta: usize,
) -> (Array1<f64>, Array1<f64>) {
let mut metric = Array1::<f64>::zeros(nsta);
let mut berry = Array1::<f64>::zeros(nsta);
let eta2 = eta * eta;
for n in 0..nsta {
let mut g_sum = Complex::new(0.0, 0.0);
for m in 0..nsta {
if m == n {
continue;
}
let de = band_q[n] - band_q[m];
let denom = de * de + eta2;
if denom < 1e-30 {
continue;
}
g_sum += k_ab_q[[n, m]] / denom;
}
metric[n] = g_sum.re;
berry[n] = -2.0 * g_sum.im;
}
(metric, berry)
}
#[allow(dead_code)]
pub(crate) fn eval_berry_band_at_lam(
n: usize,
bands: &[Vec<f64>],
kmats: &[Array2<Complex<f64>>],
lam: &[f64],
eta: f64,
nsta: usize,
) -> f64 {
let mut e_q = vec![0.0; nsta];
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_q[m] += bands[v][m] * lv;
}
}
let mut k_row = vec![Complex::new(0.0, 0.0); nsta];
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_row[m] += kmats[v][[n, m]] * lv;
}
}
let eta2 = eta * eta;
let mut g_sum = Complex::new(0.0, 0.0);
for m in 0..nsta {
if m == n {
continue;
}
let de = e_q[n] - e_q[m];
let denom = de * de + eta2;
if denom < 1e-30 {
continue;
}
g_sum += k_row[m] / denom;
}
-2.0 * g_sum.im
}
#[inline]
pub(crate) fn eval_berry_band_at_lam_buf(
n: usize,
bands: &[&[f64]],
kmats: &[&Array2<Complex<f64>>],
lam: &[f64],
eta: f64,
nsta: usize,
e_buf: &mut [f64],
k_buf: &mut [Complex<f64>],
) -> f64 {
e_buf[..nsta].fill(0.0);
k_buf[..nsta].fill(Complex::new(0.0, 0.0));
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_buf[m] += bands[v][m] * lv;
}
}
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_buf[m] += kmats[v][[n, m]] * lv;
}
}
let eta2 = eta * eta;
let mut g_sum = Complex::new(0.0, 0.0);
for m in 0..nsta {
if m == n {
continue;
}
let de = e_buf[n] - e_buf[m];
let denom = de * de + eta2;
if denom < 1e-30 {
continue;
}
g_sum += k_buf[m] / denom;
}
-2.0 * g_sum.im
}
#[allow(dead_code)]
pub(crate) fn eval_berry_complex_at_lam(
n: usize,
bands: &[Vec<f64>],
kmats: &[Array2<Complex<f64>>],
lam: &[f64],
eta: f64,
nsta: usize,
) -> (f64, f64) {
let mut e_q = vec![0.0; nsta];
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_q[m] += bands[v][m] * lv;
}
}
let mut k_row = vec![Complex::new(0.0, 0.0); nsta];
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_row[m] += kmats[v][[n, m]] * lv;
}
}
let eta2 = eta * eta;
let mut g_sum = Complex::new(0.0, 0.0);
for m in 0..nsta {
if m == n {
continue;
}
let de = e_q[n] - e_q[m];
let denom = de * de + eta2;
if denom < 1e-30 {
continue;
}
g_sum += k_row[m] / denom;
}
(g_sum.re, -2.0 * g_sum.im)
}
#[inline]
pub(crate) fn eval_berry_complex_at_lam_buf(
n: usize,
bands: &[&[f64]],
kmats: &[&Array2<Complex<f64>>],
lam: &[f64],
eta: f64,
nsta: usize,
e_buf: &mut [f64],
k_buf: &mut [Complex<f64>],
) -> (f64, f64) {
e_buf[..nsta].fill(0.0);
k_buf[..nsta].fill(Complex::new(0.0, 0.0));
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_buf[m] += bands[v][m] * lv;
}
}
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_buf[m] += kmats[v][[n, m]] * lv;
}
}
let eta2 = eta * eta;
let mut g_sum = Complex::new(0.0, 0.0);
for m in 0..nsta {
if m == n {
continue;
}
let de = e_buf[n] - e_buf[m];
let denom = de * de + eta2;
if denom < 1e-30 {
continue;
}
g_sum += k_buf[m] / denom;
}
(g_sum.re, -2.0 * g_sum.im)
}
#[inline]
#[allow(dead_code)]
pub(crate) fn eval_intrinsic_G_at_lam(
n: usize,
bands: &[Vec<f64>],
kmats: &[Array2<Complex<f64>>],
lam: &[f64],
nsta: usize,
) -> f64 {
let mut e_q = vec![0.0; nsta];
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_q[m] += bands[v][m] * lv;
}
}
let mut k_row = vec![Complex::new(0.0, 0.0); nsta];
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_row[m] += kmats[v][[n, m]] * lv;
}
}
let mut g_sum = 0.0f64;
for m in 0..nsta {
if m == n {
continue;
}
let de = e_q[n] - e_q[m];
let de3 = de * de * de;
if de3.abs() < 1e-30 {
continue;
}
g_sum += k_row[m].re / de3;
}
g_sum
}
#[inline]
#[allow(dead_code)]
pub(crate) fn eval_intrinsic_G_at_lam_buf(
n: usize,
bands: &[&[f64]],
kmats: &[&Array2<Complex<f64>>],
lam: &[f64],
nsta: usize,
e_buf: &mut [f64],
k_buf: &mut [Complex<f64>],
) -> f64 {
e_buf[..nsta].fill(0.0);
k_buf[..nsta].fill(Complex::new(0.0, 0.0));
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_buf[m] += bands[v][m] * lv;
}
}
for v in 0..kmats.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
k_buf[m] += kmats[v][[n, m]] * lv;
}
}
let mut g_sum = 0.0f64;
for m in 0..nsta {
if m == n {
continue;
}
let de = e_buf[n] - e_buf[m];
let de3 = de * de * de;
if de3.abs() < 1e-30 {
continue;
}
g_sum += k_buf[m].re / de3;
}
g_sum
}
#[inline]
pub(crate) fn eval_intrinsic_G3_at_lam_buf(
n: usize,
bands: &[&[f64]],
kmat_ab: &[&Array2<Complex<f64>>],
kmat_bc: &[&Array2<Complex<f64>>],
kmat_ac: &[&Array2<Complex<f64>>],
lam: &[f64],
nsta: usize,
e_buf: &mut [f64],
k_buf: &mut [Complex<f64>],
) -> (f64, f64, f64) {
e_buf[..nsta].fill(0.0);
let (k_ab_row, rest) = k_buf.split_at_mut(nsta);
let (k_bc_row, k_ac_row) = rest.split_at_mut(nsta);
k_ab_row.fill(Complex::new(0.0, 0.0));
k_bc_row.fill(Complex::new(0.0, 0.0));
k_ac_row.fill(Complex::new(0.0, 0.0));
for v in 0..bands.len() {
let lv = lam[v];
if lv == 0.0 {
continue;
}
for m in 0..nsta {
e_buf[m] += bands[v][m] * lv;
k_ab_row[m] += kmat_ab[v][[n, m]] * lv;
k_bc_row[m] += kmat_bc[v][[n, m]] * lv;
k_ac_row[m] += kmat_ac[v][[n, m]] * lv;
}
}
let en = e_buf[n];
let mut g_ab = 0.0f64;
let mut g_bc = 0.0f64;
let mut g_ac = 0.0f64;
for m in 0..nsta {
if m == n {
continue;
}
let de = en - e_buf[m];
let de3 = de * de * de;
if de3.abs() < 1e-30 {
continue;
}
g_ab += k_ab_row[m].re / de3;
g_bc += k_bc_row[m].re / de3;
g_ac += k_ac_row[m].re / de3;
}
(g_ab, g_bc, g_ac)
}
pub(crate) fn eval_optical_kernel(
band_q: &[f64],
k_ab_q: &Array2<Complex<f64>>,
omega: f64,
eta: f64,
mu: f64,
thermal_width: f64,
nsta: usize,
) -> Complex<f64> {
let mut total = Complex::new(0.0, 0.0);
let w_plus_ieta = Complex::new(omega, eta);
let denom_shift = w_plus_ieta * w_plus_ieta;
for n in 0..nsta {
let fn_val = fermi(band_q[n], mu, thermal_width);
for m in 0..nsta {
if m == n {
continue;
}
let fm_val = fermi(band_q[m], mu, thermal_width);
let df = fn_val - fm_val;
if df.abs() < 1e-30 {
continue;
}
let d = band_q[n] - band_q[m];
let denom = d * d - denom_shift;
if denom.norm_sqr() < 1e-30 {
continue;
}
total += df * k_ab_q[[n, m]] / denom;
}
}
total
}
#[allow(dead_code)]
pub(crate) fn quadrature_berry_simplex(sim: &TrackedSimplex, eta: f64) -> (f64, f64) {
let d = sim.vertices.len() - 1;
let nsta = sim.vertices[0].band.len();
let nv = d + 1;
let bands: Vec<Vec<f64>> = (0..nv).map(|v| sim.vertices[v].band.to_vec()).collect();
let kmats: Vec<Array2<Complex<f64>>> = (0..nv).map(|v| sim.vertices[v].k_ab.clone()).collect();
let mut total_g = 0.0;
let mut total_o = 0.0;
if d == 2 {
for iq in 0..3 {
let lam = TRI_QUAD_PTS_3[iq].as_slice();
let w = TRI_QUAD_WTS_3[iq];
let band_q = bary_interp_band(&bands, lam, nsta);
let k_ab_q = bary_interp_matrix(&kmats, lam);
let (g_n, o_n) = eval_berry_kernel(&band_q, &k_ab_q, eta, nsta);
total_g += w * g_n.iter().copied().sum::<f64>();
total_o += w * o_n.iter().copied().sum::<f64>();
}
} else {
for iq in 0..4 {
let lam = TET_QUAD_PTS_4[iq].as_slice();
let w = TET_QUAD_WTS_4[iq];
let band_q = bary_interp_band(&bands, lam, nsta);
let k_ab_q = bary_interp_matrix(&kmats, lam);
let (g_n, o_n) = eval_berry_kernel(&band_q, &k_ab_q, eta, nsta);
total_g += w * g_n.iter().copied().sum::<f64>();
total_o += w * o_n.iter().copied().sum::<f64>();
}
}
(total_g * sim.volume, total_o * sim.volume)
}
pub(crate) fn quadrature_occupied_geometry_simplex(
sim: &TrackedSimplex,
eta: f64,
chemical_potentials: &Array1<f64>,
occupation: Occupation,
) -> (Array1<f64>, Array1<f64>) {
let dimension = sim.vertices.len() - 1;
let nsta = sim.vertices[0].band.len();
let vertex_count = dimension + 1;
let bands: Vec<Vec<f64>> = (0..vertex_count)
.map(|vertex| sim.vertices[vertex].band.to_vec())
.collect();
let kernels: Vec<Array2<Complex<f64>>> = (0..vertex_count)
.map(|vertex| sim.vertices[vertex].k_ab.clone())
.collect();
let mut metric = Array1::<f64>::zeros(chemical_potentials.len());
let mut berry = Array1::<f64>::zeros(chemical_potentials.len());
let mut accumulate = |lambda: &[f64], weight: f64| {
let energies = bary_interp_band(&bands, lambda, nsta);
let kernel = bary_interp_matrix(&kernels, lambda);
let (metric_n, berry_n) = eval_berry_kernel(&energies, &kernel, eta, nsta);
for (index, &mu) in chemical_potentials.iter().enumerate() {
for band in 0..nsta {
let f = occupation.value_unchecked(energies[band], mu);
metric[index] += weight * f * metric_n[band];
berry[index] += weight * f * berry_n[band];
}
}
};
if dimension == 2 {
for index in 0..TRI_QUAD_PTS_3.len() {
accumulate(&TRI_QUAD_PTS_3[index], TRI_QUAD_WTS_3[index]);
}
} else {
for index in 0..TET_QUAD_PTS_4.len() {
accumulate(&TET_QUAD_PTS_4[index], TET_QUAD_WTS_4[index]);
}
}
(metric * sim.volume, berry * sim.volume)
}
pub(crate) fn quadrature_optical_simplex(
sim: &TrackedSimplex,
omega: f64,
eta: f64,
mu: f64,
thermal_width: f64,
) -> Complex<f64> {
let d = sim.vertices.len() - 1;
let nsta = sim.vertices[0].band.len();
let nv = d + 1;
let bands: Vec<Vec<f64>> = (0..nv).map(|v| sim.vertices[v].band.to_vec()).collect();
let kmats: Vec<Array2<Complex<f64>>> = (0..nv).map(|v| sim.vertices[v].k_ab.clone()).collect();
let mut total = Complex::new(0.0, 0.0);
if d == 2 {
for iq in 0..3 {
let lam = TRI_QUAD_PTS_3[iq].as_slice();
let w = TRI_QUAD_WTS_3[iq];
let band_q = bary_interp_band(&bands, lam, nsta);
let k_ab_q = bary_interp_matrix(&kmats, lam);
total += w * eval_optical_kernel(&band_q, &k_ab_q, omega, eta, mu, thermal_width, nsta);
}
} else {
for iq in 0..4 {
let lam = TET_QUAD_PTS_4[iq].as_slice();
let w = TET_QUAD_WTS_4[iq];
let band_q = bary_interp_band(&bands, lam, nsta);
let k_ab_q = bary_interp_matrix(&kmats, lam);
total += w * eval_optical_kernel(&band_q, &k_ab_q, omega, eta, mu, thermal_width, nsta);
}
}
total * sim.volume
}