import numpy as np
import matplotlib.pyplot as plt
from scipy.constants import hbar, m_e, pi

# -----------------------------
# パラメータ設定
# -----------------------------
a = 5.0e-10          # 格子定数 [m]
meff = 0.2           # 有効質量 (単位 m_e)
b = -1.0e-58          # 四次項係数 [J·m^4]
nk = 40              # k点分割数 (各方向)
sigma = 0.01 * 1.6e-19  # DOS計算用 smearing 幅 [J]

# エネルギー範囲
Emin = 0.0
Emax = 1.0 * 1.6e-19   # 1 eV
nE = 200
Egrid = np.linspace(Emin, Emax, nE)

# -----------------------------
# 分散関数 E(k)
# -----------------------------
def E_k(k, meff, b):
    return (hbar**2 / (2 * meff * m_e)) * k**2 + b * k**4

# -----------------------------
# 解析式DOS
# -----------------------------
def DOS_analytic(E, meff, b):
    # a = hbar^2 / (2 meff m_e)
    a_coeff = hbar**2 / (2 * meff * m_e)

    # k^2(E) 解く
    disc = a_coeff**2 + 4*b*E
    k2 = (-a_coeff + np.sqrt(disc)) / (2*b)
    k2 = np.where(k2 > 0, k2, 0.0)
    k = np.sqrt(k2)

    # DOS formula: (1/(4π^2)) * k / (a + 2 b k^2)
    denom = a_coeff + 2*b*k2
    DOS = (1.0 / (4*pi**2)) * np.divide(k, denom, out=np.zeros_like(k), where=denom>0)
    return DOS

# -----------------------------
# 数値積分によるDOS (BZ内kサンプリング)
# -----------------------------
# BZの範囲 [-pi/a, pi/a]
kmax = pi / a
klist = np.linspace(-kmax, kmax, nk)
dk = (2*kmax / nk)

# k点メッシュ
kx, ky, kz = np.meshgrid(klist, klist, klist, indexing="ij")
kpts = np.stack([kx, ky, kz], axis=-1).reshape(-1, 3)
k_norm = np.linalg.norm(kpts, axis=1)

Ek = E_k(k_norm, meff, b)

# smearing でDOS計算
DOS_num = np.zeros_like(Egrid)
pref = 1.0 / (2*pi)**3  # DOS prefactor
dV = dk**3              # k-space volume element

for iE, E in enumerate(Egrid):
    weight = np.exp(-((E - Ek)**2) / (2*sigma**2)) / (np.sqrt(2*pi) * sigma)
    DOS_num[iE] = pref * np.sum(weight) * dV

# -----------------------------
# 計算結果をプロット
# -----------------------------
plt.figure(figsize=(6,4))
plt.plot(Egrid/1.6e-19, DOS_analytic(Egrid, meff, b), label="Analytic DOS")
plt.plot(Egrid/1.6e-19, DOS_num, label="Numerical DOS (Gaussian smearing)", linestyle="--")
plt.xlabel("Energy [eV]")
plt.ylabel("DOS [states/eV·m³]")
plt.legend()
plt.title("Comparison of Analytic vs Numerical DOS")
plt.tight_layout()
plt.show()
