Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 17 additions & 22 deletions source/source_hsolver/diago_iter_assist.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,8 @@ namespace hsolver
//----------------------------------------------------------------------
template <typename T, typename Device>
void DiagoIterAssist<T, Device>::diag_subspace(
const hamilt::Hamilt<T, Device>* const pHamilt, // hamiltonian operator carrier
const HPsiFunc& hpsi_func, // applies H to a block of vectors
const SPsiFunc& spsi_func, // applies S to a block of vectors
const psi::Psi<T, Device>& psi, // [in] wavefunction
psi::Psi<T, Device>& evc, // [out] wavefunction, eigenvectors
Real* en, // [out] eigenvalues
Expand Down Expand Up @@ -83,9 +84,7 @@ void DiagoIterAssist<T, Device>::diag_subspace(

T *hpsi = temp;
// do hPsi for all bands
psi::Range all_bands_range(1, psi.get_current_k(), 0, nstart - 1);
hpsi_info hpsi_in(&psi, all_bands_range, hpsi);
pHamilt->ops->hPsi(hpsi_in);
hpsi_func(psi.get_pointer(), hpsi, dmax, psi.get_current_nbas(), nstart);

ModuleBase::gemm_op<T, Device>()('C',
'N',
Expand All @@ -105,7 +104,7 @@ void DiagoIterAssist<T, Device>::diag_subspace(
// Only calculate S_sub if not orthogonal
T *spsi = temp;
// do sPsi for all bands
pHamilt->sPsi(psi.get_pointer(), spsi, dmax, dmin, nstart);
spsi_func(psi.get_pointer(), spsi, dmax, dmin, nstart);

ModuleBase::gemm_op<T, Device>()('C',
'N',
Expand Down Expand Up @@ -176,7 +175,8 @@ void DiagoIterAssist<T, Device>::diag_subspace(

template <typename T, typename Device>
void DiagoIterAssist<T, Device>::diag_subspace_init(
hamilt::Hamilt<T, Device>* pHamilt,
const HPsiFunc& hpsi_func,
const SPsiFunc& spsi_func,
const T* psi,
int psi_nr,
int psi_nc,
Expand All @@ -200,8 +200,8 @@ void DiagoIterAssist<T, Device>::diag_subspace_init(
const int dmax = evc.get_nbasis();
const int dmin = evc.get_current_ngk();

// skip the diagonalization if the operators are not allocated
if (pHamilt->ops == nullptr)
// skip the diagonalization if the caller has no operators allocated
if (!hpsi_func)
{
ModuleBase::WARNING(
"DiagoIterAssist::diag_subspace_init",
Expand Down Expand Up @@ -247,11 +247,9 @@ void DiagoIterAssist<T, Device>::diag_subspace_init(
{
// psi_temp is one band psi, psi is all bands psi, the range always is 1 for the only band in psi_temp
syncmem_complex_op()(ppsi, psi + i * psi_nc, psi_nc);
psi::Range band_by_band_range(true, 0, 0, 0);
hpsi_info hpsi_in(&psi_temp, band_by_band_range, hpsi);

// H|Psi> to get hpsi for target band
pHamilt->ops->hPsi(hpsi_in);
hpsi_func(ppsi, hpsi, psi_temp.get_nbasis(), psi_temp.get_current_nbas(), 1);

// calculate the related elements in hcc <Psi|H|Psi>
ModuleBase::gemv_op<T, Device>()('C', psi_nc, nstart, &one, psi, psi_nc, hpsi, 1, &zero, hcc + i * nstart, 1);
Expand All @@ -262,7 +260,7 @@ void DiagoIterAssist<T, Device>::diag_subspace_init(
for (int i = 0; i < nstart; i++)
{
syncmem_complex_op()(ppsi, psi + i * psi_nc, psi_nc);
pHamilt->sPsi(ppsi, spsi, dmin, dmin, 1);
spsi_func(ppsi, spsi, dmin, dmin, 1);

ModuleBase::gemv_op<T, Device>()('C',
psi_nc,
Expand Down Expand Up @@ -295,15 +293,13 @@ void DiagoIterAssist<T, Device>::diag_subspace_init(

T* hpsi = temp;
// do hPsi for all bands
psi::Range all_bands_range(true, 0, 0, nstart - 1);
hpsi_info hpsi_in(&psi_temp, all_bands_range, hpsi);
pHamilt->ops->hPsi(hpsi_in);
hpsi_func(ppsi, hpsi, psi_temp.get_nbasis(), psi_temp.get_current_nbas(), nstart);

ModuleBase::gemm_op<T, Device>()('C', 'N', nstart, nstart, dmin, &one, ppsi, dmax, hpsi, dmax, &zero, hcc, nstart);

T* spsi = temp;
// do sPsi for all bands
pHamilt->sPsi(ppsi, spsi, psi_temp.get_nbasis(), psi_temp.get_nbasis(), psi_temp.get_nbands());
spsi_func(ppsi, spsi, psi_temp.get_nbasis(), psi_temp.get_nbasis(), psi_temp.get_nbands());

ModuleBase::gemm_op<T, Device>()('C', 'N', nstart, nstart, dmin, &one, ppsi, dmax, spsi, dmax, &zero, scc, nstart);
delmem_complex_op()(temp);
Expand Down Expand Up @@ -488,8 +484,9 @@ void DiagoIterAssist<T, Device>::diag_hegvd(const int nstart,

template <typename T, typename Device>
void DiagoIterAssist<T, Device>::cal_hs_subspace(
const hamilt::Hamilt<T, Device>* pHamilt, // hamiltonian operator carrier
const psi::Psi<T, Device>& psi, // [in] wavefunction
const HPsiFunc& hpsi_func, // applies H to a block of vectors
const SPsiFunc& spsi_func, // applies S to a block of vectors
const psi::Psi<T, Device>& psi, // [in] wavefunction
T* hcc,
T* scc,
const diag_comm_info& diag_comm)
Expand All @@ -511,9 +508,7 @@ void DiagoIterAssist<T, Device>::cal_hs_subspace(

T* hpsi = temp;
// do hPsi for all bands
psi::Range all_bands_range(1, psi.get_current_k(), 0, nstart - 1);
hpsi_info hpsi_in(&psi, all_bands_range, hpsi);
pHamilt->ops->hPsi(hpsi_in);
hpsi_func(psi.get_pointer(), hpsi, dmax, psi.get_current_nbas(), nstart);

ModuleBase::gemm_op<T, Device>()('C',
'N',
Expand All @@ -531,7 +526,7 @@ void DiagoIterAssist<T, Device>::cal_hs_subspace(

T* spsi = temp;
// do sPsi for all bands
pHamilt->sPsi(psi.get_pointer(), spsi, dmax, dmin, nstart);
spsi_func(psi.get_pointer(), spsi, dmax, dmin, nstart);

ModuleBase::gemm_op<T, Device>()('C',
'N',
Expand Down
41 changes: 29 additions & 12 deletions source/source_hsolver/diago_iter_assist.h
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@

#include "source_base/complexmatrix.h"
#include "source_base/macros.h"
#include "source_hamilt/hamilt.h"
#include "source_psi/psi.h"

#include <functional>
Expand All @@ -21,6 +20,18 @@ class DiagoIterAssist
using Real = typename GetTypeReal<T>::type;

public:
/// Apply H to a block of `nvec` vectors laid out with leading dimension
/// `ld_psi`. `current_nbasis` is the number of plane waves *without* npol;
/// it is passed explicitly rather than captured because the callers below
/// drive several different wavefunction layouts through the same functor.
using HPsiFunc = std::function<
void(T* psi_in, T* hpsi_out, const int ld_psi, const int current_nbasis, const int nvec)>;

/// Apply S to a block of `nbands` vectors. The parameter list mirrors
/// hamilt::Hamilt::sPsi exactly, so a caller's functor is a plain forwarder.
using SPsiFunc = std::function<
void(const T* psi_in, T* spsi_out, const int nrow, const int npw, const int nbands)>;

static Real PW_DIAG_THR;
static int PW_DIAG_NMAX;

Expand All @@ -40,14 +51,16 @@ class DiagoIterAssist
*
* @tparam T Data type for computation (e.g., float, double).
* @tparam Device Device type for computation (e.g., CPU, GPU).
* @param pHamilt Pointer to the Hamiltonian object.
* @param hpsi_func Applies H to a block of vectors.
* @param spsi_func Applies S to a block of vectors.
* @param psi Input wavefunction defining the subspace.
* @param evc Output container for computed eigenvectors.
* @param en Output array for computed eigenvalues.
* @param n_band Number of bands (eigenvalues/eigenvectors) to compute. Default is 0 (all).
* @param is_S_orthogonal If true, assumes the input wavefunction is already orthogonalized.
*/
static void diag_subspace(const hamilt::Hamilt<T, Device>* const pHamilt,
static void diag_subspace(const HPsiFunc& hpsi_func,
const SPsiFunc& spsi_func,
const psi::Psi<T, Device>& psi,
psi::Psi<T, Device>& evc,
Real* en,
Expand All @@ -56,7 +69,9 @@ class DiagoIterAssist
const bool is_S_orthogonal = false);

/// @brief use LAPACK to diagonalize the Hamiltonian matrix
/// @param pHamilt interface to hamiltonian
/// @param hpsi_func applies H to a block of vectors; pass an empty functor
/// when the Hamiltonian has no operators allocated yet (see the note below)
/// @param spsi_func applies S to a block of vectors
/// @param psi wavefunction to diagonalize
/// @param psi_nr number of rows (nbands)
/// @param psi_nc number of columns (nbasis)
Expand All @@ -65,10 +80,12 @@ class DiagoIterAssist
/// @param basis_type "lcao", "lcao_in_pw" or "pw"; together with calculation it selects
/// how the rotation matrix is applied to psi
/// @param calculation "scf", "nscf", "md", "relax", ...
/// @note exception handle: if there is no operator initialized in Hamilt, will directly copy value from psi to evc,
/// and return all - zero eigenenergies.
/// @note exception handle: if hpsi_func is empty, meaning the caller has no
/// operators initialized in its Hamiltonian, will directly copy value from
/// psi to evc, and return all - zero eigenenergies.
static void diag_subspace_init(
hamilt::Hamilt<T, Device>* pHamilt,
const HPsiFunc& hpsi_func,
const SPsiFunc& spsi_func,
const T* psi,
int psi_nr,
int psi_nc,
Expand Down Expand Up @@ -96,12 +113,14 @@ class DiagoIterAssist
T *vcc);

/// @brief calculate Hamiltonian and overlap matrix in subspace spanned by nstart states psi
/// @param pHamilt : hamiltonian operator carrier
/// @param hpsi_func : applies H to a block of vectors
/// @param spsi_func : applies S to a block of vectors
/// @param psi : wavefunction
/// @param hcc : Hamiltonian matrix
/// @param scc : overlap matrix
static void cal_hs_subspace(const hamilt::Hamilt<T, Device>* pHamilt, // hamiltonian operator carrier
const psi::Psi<T, Device>& psi, // [in] wavefunction
static void cal_hs_subspace(const HPsiFunc& hpsi_func,
const SPsiFunc& spsi_func,
const psi::Psi<T, Device>& psi, // [in] wavefunction
T* hcc,
T* scc,
const diag_comm_info& diag_comm);
Expand Down Expand Up @@ -132,8 +151,6 @@ class DiagoIterAssist
private:
constexpr static const Device* ctx = {};

using hpsi_info = typename hamilt::Operator<T, Device>::hpsi_info;

using setmem_var_op = base_device::memory::set_memory_op<Real, Device>;
using resmem_var_op = base_device::memory::resize_memory_op<Real, Device>;
using delmem_var_op = base_device::memory::delete_memory_op<Real, Device>;
Expand Down
20 changes: 19 additions & 1 deletion source/source_hsolver/hsolver_lcaopw.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -69,8 +69,26 @@ void HSolverLIP<T>::solve(hamilt::Hamilt<T>* pHamilt, // ESolver_KS_PW::p_hamilt
}
};
#endif
/// An empty hpsi functor tells diag_subspace_init that the Hamiltonian
/// has no operators allocated yet, which used to be the ops == nullptr check.
typename hsolver::DiagoIterAssist<T>::HPsiFunc hpsi_func;
if (pHamilt->ops != nullptr)
{
hpsi_func = [pHamilt](T* psi_in, T* hpsi_out, const int ld_psi, const int current_nbasis, const int nvec) {
auto psi_wrapper = psi::Psi<T>(psi_in, 1, nvec, ld_psi, current_nbasis);
psi::Range bands_range(true, 0, 0, nvec - 1);
using hpsi_info = typename hamilt::Operator<T>::hpsi_info;
hpsi_info info(&psi_wrapper, bands_range, hpsi_out);
pHamilt->ops->hPsi(info);
};
}
auto spsi_func = [pHamilt](const T* psi_in, T* spsi_out, const int nrow, const int npw, const int nbands) {
pHamilt->sPsi(psi_in, spsi_out, nrow, npw, nbands);
};

/// solve eigenvector and eigenvalue for H(k)
hsolver::DiagoIterAssist<T>::diag_subspace_init(pHamilt, // interface to hamilt
hsolver::DiagoIterAssist<T>::diag_subspace_init(hpsi_func,
spsi_func,
transform.get_pointer(), // transform matrix between lcao and pw
transform.get_nbands(),
transform.get_nbasis(),
Expand Down
37 changes: 30 additions & 7 deletions source/source_hsolver/hsolver_pw.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -282,13 +282,36 @@ void HSolverPW<T, Device>::hamiltSolvePsiK(hamilt::Hamilt<T, Device>* hm,
// wrap the subspace_func into a lambda function
// if S_orth is true, then assume psi is S-orthogonal, solve standard eigenproblem
// otherwise, solve generalized eigenproblem
auto subspace_func =
[hm, cur_nbasis, &comm_info](T* psi_in, T* psi_out, const int ld_psi, const int nband, const bool S_orth) {
auto psi_in_wrapper = psi::Psi<T, Device>(psi_in, 1, nband, ld_psi, cur_nbasis);
auto psi_out_wrapper = psi::Psi<T, Device>(psi_out, 1, nband, ld_psi, cur_nbasis);
std::vector<Real> eigen(nband, 0.0);
DiagoIterAssist<T, Device>::diag_subspace(hm, psi_in_wrapper, psi_out_wrapper, eigen.data(), comm_info);
};
// DiagoIterAssist drives more than one wavefunction layout through the
// same functor, so it hands the dimensions over on every call instead of
// relying on a captured set.
auto sub_hpsi_func
= [hm](T* psi_in, T* hpsi_out, const int ld_psi, const int current_nbasis, const int nvec) {
auto psi_wrapper = psi::Psi<T, Device>(psi_in, 1, nvec, ld_psi, current_nbasis);
psi::Range bands_range(true, 0, 0, nvec - 1);
using hpsi_info = typename hamilt::Operator<T, Device>::hpsi_info;
hpsi_info info(&psi_wrapper, bands_range, hpsi_out);
hm->ops->hPsi(info);
};
auto sub_spsi_func
= [hm](const T* psi_in, T* spsi_out, const int nrow, const int npw, const int nbands) {
hm->sPsi(psi_in, spsi_out, nrow, npw, nbands);
};
auto subspace_func = [cur_nbasis, &comm_info, sub_hpsi_func, sub_spsi_func](T* psi_in,
T* psi_out,
const int ld_psi,
const int nband,
const bool S_orth) {
auto psi_in_wrapper = psi::Psi<T, Device>(psi_in, 1, nband, ld_psi, cur_nbasis);
auto psi_out_wrapper = psi::Psi<T, Device>(psi_out, 1, nband, ld_psi, cur_nbasis);
std::vector<Real> eigen(nband, 0.0);
DiagoIterAssist<T, Device>::diag_subspace(sub_hpsi_func,
sub_spsi_func,
psi_in_wrapper,
psi_out_wrapper,
eigen.data(),
comm_info);
};
DiagoCG<T, Device> cg(this->basis_type,
this->calculation_type,
this->need_subspace,
Expand Down
13 changes: 12 additions & 1 deletion source/source_hsolver/test/diago_cg_float_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,18 @@ class DiagoCGPrepare
auto psi_in_wrapper = psi::Psi<std::complex<float>>(psi_in, 1, nband, ld_psi, true);
auto psi_out_wrapper = psi::Psi<std::complex<float>>(psi_out, 1, nband, ld_psi, true);
std::vector<float> eigen(nband, 0.0f);
hsolver::DiagoIterAssist<std::complex<float>>::diag_subspace(ha,
auto sub_hpsi = [ha](std::complex<float>* p, std::complex<float>* hp, const int ld, const int cur_nbas, const int nvec) {
Comment thread
Critsium-xy marked this conversation as resolved.
Outdated
auto w = psi::Psi<std::complex<float>>(p, 1, nvec, ld, cur_nbas);
psi::Range r(true, 0, 0, nvec - 1);
using hpsi_info = typename hamilt::Operator<std::complex<float>>::hpsi_info;
hpsi_info info(&w, r, hp);
ha->ops->hPsi(info);
};
auto sub_spsi = [ha](const std::complex<float>* p, std::complex<float>* sp, const int nrow, const int npw, const int nb) {
ha->sPsi(p, sp, nrow, npw, nb);
};
hsolver::DiagoIterAssist<std::complex<float>>::diag_subspace(sub_hpsi,
sub_spsi,
psi_in_wrapper,
psi_out_wrapper,
eigen.data(),
Expand Down
13 changes: 12 additions & 1 deletion source/source_hsolver/test/diago_cg_real_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,18 @@ class DiagoCGPrepare
auto psi_in_wrapper = psi::Psi<double>(psi_in, 1, nband, ld_psi, true);
auto psi_out_wrapper = psi::Psi<double>(psi_out, 1, nband, ld_psi, true);
std::vector<double> eigen(nband, 0.0);
hsolver::DiagoIterAssist<double>::diag_subspace(ha,
auto sub_hpsi = [ha](double* p, double* hp, const int ld, const int cur_nbas, const int nvec) {
auto w = psi::Psi<double>(p, 1, nvec, ld, cur_nbas);
psi::Range r(true, 0, 0, nvec - 1);
using hpsi_info = typename hamilt::Operator<double>::hpsi_info;
hpsi_info info(&w, r, hp);
ha->ops->hPsi(info);
};
auto sub_spsi = [ha](const double* p, double* sp, const int nrow, const int npw, const int nb) {
ha->sPsi(p, sp, nrow, npw, nb);
};
hsolver::DiagoIterAssist<double>::diag_subspace(sub_hpsi,
sub_spsi,
psi_in_wrapper,
psi_out_wrapper,
eigen.data(),
Expand Down
13 changes: 12 additions & 1 deletion source/source_hsolver/test/diago_cg_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -147,7 +147,18 @@ class DiagoCGPrepare
auto psi_in_wrapper = psi::Psi<std::complex<double>>(psi_in, 1, nband, ld_psi, true);
auto psi_out_wrapper = psi::Psi<std::complex<double>>(psi_out, 1, nband, ld_psi, true);
std::vector<double> eigen(nband, 0.0);
hsolver::DiagoIterAssist<std::complex<double>>::diag_subspace(ha,
auto sub_hpsi = [ha](std::complex<double>* p, std::complex<double>* hp, const int ld, const int cur_nbas, const int nvec) {
auto w = psi::Psi<std::complex<double>>(p, 1, nvec, ld, cur_nbas);
psi::Range r(true, 0, 0, nvec - 1);
using hpsi_info = typename hamilt::Operator<std::complex<double>>::hpsi_info;
hpsi_info info(&w, r, hp);
ha->ops->hPsi(info);
};
auto sub_spsi = [ha](const std::complex<double>* p, std::complex<double>* sp, const int nrow, const int npw, const int nb) {
ha->sPsi(p, sp, nrow, npw, nb);
};
hsolver::DiagoIterAssist<std::complex<double>>::diag_subspace(sub_hpsi,
sub_spsi,
psi_in_wrapper,
psi_out_wrapper,
eigen.data(),
Expand Down
Loading
Loading