Skip to content

Commit 5033b73

Browse files
author
dyzheng
committed
Feat(psi): integrate k-point paging into HSolverPW and force/stress paths
- hsolver_pw.cpp: load_k_to_gpu/store_k_from_gpu around k-point solve loop Skip k-parent propagation in PAGED_GPU mode - psi_prepare.cpp: PAGED_GPU memory handling in initialize_psi - setup_psi_pw.cpp: detect paging mode, log GPU memory mode, copy_d2h paging path - forces/stress onsite/nl/kin: ensure_k_on_gpu before k-point loops - onsite_proj.cpp: ensure_k_on_gpu in cal_occupations
1 parent cfe3238 commit 5033b73

9 files changed

Lines changed: 86 additions & 9 deletions

File tree

source/source_hsolver/hsolver_pw.cpp

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -109,13 +109,14 @@ void HSolverPW<T, Device>::solve(hamilt::Hamilt<T, Device>* pHamilt,
109109
// update H(k) for each k point
110110
pHamilt->updateHk(ik);
111111

112-
112+
psi.load_k_to_gpu(ik);
113113

114114
// update psi pointer for each k point
115115
psi.fix_k(ik);
116116

117117
// If using k-point continuity and not first k-point, propagate from parent
118-
if (ik > 0 && count == 0 && k_parent.find(ik) != k_parent.end()) {
118+
if (ik > 0 && count == 0 && k_parent.find(ik) != k_parent.end()
119+
&& psi.get_storage_mode() != psi::PsiStorageMode::PAGED_GPU) {
119120
propagate_psi(psi, k_parent[ik], ik);
120121
}
121122

@@ -136,6 +137,8 @@ void HSolverPW<T, Device>::solve(hamilt::Hamilt<T, Device>* pHamilt,
136137
// solve eigenvector and eigenvalue for H(k)
137138
this->hamiltSolvePsiK(pHamilt, psi, precondition, eigenvalues.data() + ik * psi.get_nbands(), this->wfc_basis->nks);
138139

140+
psi.store_k_from_gpu(ik);
141+
139142
if (skip_charge)
140143
{
141144
GlobalV::ofs_running << " Average iterative diagonalization steps for k-points " << ik
@@ -152,7 +155,7 @@ void HSolverPW<T, Device>::solve(hamilt::Hamilt<T, Device>* pHamilt,
152155
// update H(k) for each k point
153156
pHamilt->updateHk(ik);
154157

155-
158+
psi.load_k_to_gpu(ik);
156159

157160
// update psi pointer for each k point
158161
psi.fix_k(ik);
@@ -174,6 +177,8 @@ void HSolverPW<T, Device>::solve(hamilt::Hamilt<T, Device>* pHamilt,
174177
// solve eigenvector and eigenvalue for H(k)
175178
this->hamiltSolvePsiK(pHamilt, psi, precondition, eigenvalues.data() + ik * psi.get_nbands(), this->wfc_basis->nks);
176179

180+
psi.store_k_from_gpu(ik);
181+
177182
// output iteration information and reset avg_iter
178183
if (skip_charge)
179184
{

source/source_psi/psi_prepare.cpp

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -176,7 +176,16 @@ void PSIPrepare<T, Device>::initialize_psi(Psi<std::complex<double>>* psi,
176176
this->psi_initer->init_psig(psi_cpu->get_pointer(), ik);
177177
if (psi_device->get_pointer() != psi_cpu->get_pointer())
178178
{
179-
syncmem_h2d_op()(psi_device->get_pointer(), psi_cpu->get_pointer(), nbands_start * nbasis);
179+
if (kspw_psi->get_storage_mode() == psi::PsiStorageMode::PAGED_GPU)
180+
{
181+
T* dst_cpu = kspw_psi->get_cpu_pointer(ik);
182+
std::memcpy(dst_cpu, psi_cpu->get_pointer(), sizeof(T) * nbands_l * nbasis);
183+
kspw_psi->load_k_to_gpu(ik);
184+
}
185+
else
186+
{
187+
syncmem_h2d_op()(psi_device->get_pointer(), psi_cpu->get_pointer(), nbands_start * nbasis);
188+
}
180189
}
181190

182191

@@ -208,7 +217,13 @@ void PSIPrepare<T, Device>::initialize_psi(Psi<std::complex<double>>* psi,
208217
{
209218
if (psi_device->get_pointer() != kspw_psi->get_pointer())
210219
{
211-
syncmem_complex_op()(kspw_psi->get_pointer(), psi_device->get_pointer(), nbands_l * nbasis);
220+
if (kspw_psi->get_storage_mode() == psi::PsiStorageMode::PAGED_GPU)
221+
{
222+
}
223+
else
224+
{
225+
syncmem_complex_op()(kspw_psi->get_pointer(), psi_device->get_pointer(), nbands_l * nbasis);
226+
}
212227
}
213228
}
214229
}

source/source_psi/setup_psi_pw.cpp

Lines changed: 39 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,33 @@ void Setup_Psi_pw::before_runner_impl(
4040
}
4141

4242
if (inp.device == "gpu" || inp.precision == "single") {
43+
const int nks = kv.get_nks();
44+
psi::PsiStorageMode target_mode = psi::PsiStorageMode::ALL_GPU;
45+
if (inp.device_memory_mode == "paged")
46+
{
47+
target_mode = psi::PsiStorageMode::PAGED_GPU;
48+
}
49+
else if (inp.device_memory_mode == "full_gpu")
50+
{
51+
target_mode = psi::PsiStorageMode::ALL_GPU;
52+
}
53+
else if (nks > 10)
54+
{
55+
target_mode = psi::PsiStorageMode::PAGED_GPU;
56+
}
57+
if (target_mode != psi::PsiStorageMode::ALL_GPU)
58+
{
59+
this->psi_cpu->set_storage_mode(target_mode);
60+
}
61+
const char* mode_str = (target_mode == psi::PsiStorageMode::PAGED_GPU) ? "PAGED_GPU" : "ALL_GPU";
62+
GlobalV::ofs_running << " GPU memory mode for Psi: " << mode_str
63+
<< " (nks=" << nks << ", device_memory_mode=\""
64+
<< inp.device_memory_mode << "\")" << std::endl;
4365
this->psi_t = static_cast<void*>(new psi::Psi<T, Device>(this->psi_cpu[0]));
66+
if (target_mode != psi::PsiStorageMode::ALL_GPU)
67+
{
68+
this->psi_cpu->set_storage_mode(psi::PsiStorageMode::ALL_GPU);
69+
}
4470
} else {
4571
this->psi_t = static_cast<void*>(reinterpret_cast<psi::Psi<T, Device>*>(this->psi_cpu));
4672
}
@@ -178,9 +204,19 @@ template <typename T, typename Device>
178204
void Setup_Psi_pw::copy_d2h_impl()
179205
{
180206
auto* psi_t = this->get_psi_t<T, Device>();
181-
this->castmem_d2h_impl<T, Device>(this->psi_cpu[0].get_pointer() - this->psi_cpu[0].get_psi_bias(),
182-
psi_t->get_pointer() - psi_t->get_psi_bias(),
183-
this->psi_cpu[0].size());
207+
if (psi_t->get_storage_mode() == psi::PsiStorageMode::PAGED_GPU)
208+
{
209+
// PAGED_GPU: full data is on CPU in psi_cpu_, just memcpy
210+
const size_t total_size = sizeof(T) * psi_t->get_nk() * psi_t->get_nbands() * psi_t->get_nbasis();
211+
std::memcpy(this->psi_cpu[0].get_pointer() - this->psi_cpu[0].get_psi_bias(),
212+
psi_t->get_cpu_pointer(), total_size);
213+
}
214+
else
215+
{
216+
this->castmem_d2h_impl<T, Device>(this->psi_cpu[0].get_pointer() - this->psi_cpu[0].get_psi_bias(),
217+
psi_t->get_pointer() - psi_t->get_psi_bias(),
218+
this->psi_cpu[0].size());
219+
}
184220
}
185221

186222
void Setup_Psi_pw::copy_d2h()

source/source_pw/module_pwdft/forces_nl.cpp

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,9 @@ void Forces<FPTYPE, Device>::cal_force_nl(ModuleBase::matrix& forcenl,
3737
const int max_nbands = wg.nc;
3838
for (int ik = 0; ik < nks; ik++) // loop k points
3939
{
40+
// Ensure k-point is loaded to GPU for PAGED_GPU mode
41+
const_cast<psi::Psi<std::complex<FPTYPE>, Device>*>(psi_in)->ensure_k_on_gpu(ik);
42+
4043
// skip zero weights to speed up
4144
int nbands_occ = wg.nc;
4245
while (wg(ik, nbands_occ - 1) == 0.0)

source/source_pw/module_pwdft/forces_onsite.cpp

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,11 @@ void Forces<FPTYPE, Device>::cal_force_onsite(ModuleBase::matrix& force_onsite,
3232
const int nks = wfc_basis->nks;
3333
for (int ik = 0; ik < nks; ik++)
3434
{
35+
// In PAGED_GPU mode only one k-point resides on the GPU; make sure the
36+
// wavefunction of this k-point is loaded before the onsite projector
37+
// reads psi_ (otherwise psi_(ik,0,0) returns a host pointer that the
38+
// GPU gemm would dereference as a device pointer -> illegal access).
39+
const_cast<psi::Psi<std::complex<FPTYPE>, Device>*>(psi_in)->ensure_k_on_gpu(ik);
3540
int nbands_occ = wg.nc;
3641
while (wg(ik, nbands_occ - 1) == 0.0)
3742
{

source/source_pw/module_pwdft/onsite_proj.cpp

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -539,6 +539,7 @@ void projectors::OnsiteProjector<T, Device>::cal_occupations(
539539
const int nbands = psi_in->get_nbands();
540540
for(int ik = 0; ik < psi_in->get_nk(); ik++)
541541
{
542+
const_cast<psi::Psi<std::complex<T>, Device>*>(psi_in)->ensure_k_on_gpu(ik);
542543
psi_in->fix_k(ik);
543544
if(ik != 0)
544545
{

source/source_pw/module_pwdft/stress_kin.cpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,9 +18,12 @@ void Stress_Func<FPTYPE, Device>::stress_kin(ModuleBase::matrix& sigma,
1818

1919
this->ucell = &ucell_in;
2020

21-
hamilt::FS_Kin_tools<FPTYPE, Device> kin_tool(*this->ucell, p_kv, wfc_basis, wg);
21+
hamilt::FS_Kin_tools<FPTYPE, Device> kin_tool(*this->ucell, p_kv, wfc_basis, wg);
2222
for (int ik = 0; ik < wfc_basis->nks; ++ik)
2323
{
24+
// Ensure k-point is loaded to GPU for PAGED_GPU mode
25+
const_cast<psi::Psi<std::complex<FPTYPE>, Device>*>(psi_in)->ensure_k_on_gpu(ik);
26+
2427
int nbands_occ = wg.nc;
2528
while (wg(ik, nbands_occ - 1) == 0.0)
2629
{

source/source_pw/module_pwdft/stress_nl.cpp

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,9 @@ void Stress_Func<FPTYPE, Device>::stress_nl(ModuleBase::matrix& sigma,
3939
const int max_nbands = wg.nc;
4040
for (int ik = 0; ik < nks; ik++) // loop k points
4141
{
42+
// Ensure k-point is loaded to GPU for PAGED_GPU mode
43+
const_cast<psi::Psi<std::complex<FPTYPE>, Device>*>(psi_in)->ensure_k_on_gpu(ik);
44+
4245
// skip zero weights to speed up
4346
int nbands_occ = wg.nc;
4447
while (wg(ik, nbands_occ - 1) == 0.0)

source/source_pw/module_pwdft/stress_onsite.cpp

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,12 @@ void Stress_Func<FPTYPE, Device>::stress_onsite(
6060
// Loop over all k-points
6161
for (int ik = 0; ik < nks; ik++)
6262
{
63+
// In PAGED_GPU mode only one k-point resides on the GPU; load this
64+
// k-point's wavefunction before the onsite projector reads psi_
65+
// (otherwise psi_(ik,0,0) returns a host pointer used as a device
66+
// pointer by the GPU gemm -> illegal memory access).
67+
const_cast<psi::Psi<std::complex<FPTYPE>, Device>*>(
68+
static_cast<const psi::Psi<std::complex<FPTYPE>, Device>*>(psi_in))->ensure_k_on_gpu(ik);
6369
// Determine number of occupied bands (skip zero weights)
6470
int nbands_occ = wg.nc;
6571

0 commit comments

Comments
 (0)