Skip to content

Commit 4ba66ae

Browse files
authored
Merge branch 'develop' into fix/7681-cuda-fft-grid
2 parents a6bf6c6 + acf177d commit 4ba66ae

3 files changed

Lines changed: 232 additions & 9 deletions

File tree

source/source_pw/module_pwdft/kernels/cuda/stress_op.cu

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -350,10 +350,10 @@ __global__ void cal_stress_nl(
350350
{
351351
ps_qq = thrust::complex<FPTYPE>(- ekb_now * qq_nt[it * deeq_3 * deeq_4 + ip1 * deeq_4 + ip2], 0.0);
352352
}
353-
const thrust::complex<FPTYPE> ps0 = deeq_nc[((iat + ia) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
354-
const thrust::complex<FPTYPE> ps1 = deeq_nc[((1 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2];
355-
const thrust::complex<FPTYPE> ps2 = deeq_nc[((2 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2];
356-
const thrust::complex<FPTYPE> ps3 = deeq_nc[((3 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
353+
const thrust::complex<FPTYPE> ps0 = deeq_nc[((iat) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
354+
const thrust::complex<FPTYPE> ps1 = deeq_nc[((1 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2];
355+
const thrust::complex<FPTYPE> ps2 = deeq_nc[((2 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2];
356+
const thrust::complex<FPTYPE> ps3 = deeq_nc[((3 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
357357
const int inkb1 = sum + ip1;
358358
const int inkb2 = sum + ip2;
359359
//out<<"\n ps = "<<ps;

source/source_pw/module_pwdft/kernels/rocm/stress_op.hip.cu

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -311,10 +311,10 @@ __global__ void cal_stress_nl(
311311
{
312312
ps_qq = thrust::complex<FPTYPE>(- ekb_now * qq_nt[it * deeq_3 * deeq_4 + ip1 * deeq_4 + ip2], 0.0);
313313
}
314-
const thrust::complex<FPTYPE> ps0 = deeq_nc[((iat + ia) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
315-
const thrust::complex<FPTYPE> ps1 = deeq_nc[((1 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2];
316-
const thrust::complex<FPTYPE> ps2 = deeq_nc[((2 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2];
317-
const thrust::complex<FPTYPE> ps3 = deeq_nc[((3 * deeq_2 + iat + ia) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
314+
const thrust::complex<FPTYPE> ps0 = deeq_nc[((iat) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
315+
const thrust::complex<FPTYPE> ps1 = deeq_nc[((1 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2];
316+
const thrust::complex<FPTYPE> ps2 = deeq_nc[((2 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2];
317+
const thrust::complex<FPTYPE> ps3 = deeq_nc[((3 * deeq_2 + iat) * deeq_3 + ip1) * deeq_4 + ip2] + ps_qq;
318318
const int inkb1 = sum + ip1;
319319
const int inkb2 = sum + ip2;
320320
//out<<"\n ps = "<<ps;

source/source_pw/module_pwdft/kernels/test/stress_op_test.cpp

Lines changed: 224 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -302,4 +302,227 @@ TEST(TestSrcPWStressMultiDevice, cal_stress_nl_op_gpu)
302302
delmem_int_op()(d_atom_nh);
303303
delmem_int_op()(d_atom_na);
304304
}
305-
#endif // __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM
305+
#endif // __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM
306+
307+
// Non-collinear (nspin=4) cal_stress_nl test with deeq_nc.
308+
// Uses ntype = 1 and atom_na = {2} so that the second atom of the same
309+
// element type is indexed: this catches the GPU bug where iat was both
310+
// incremented inside the atom loop and added to ia (double counting).
311+
namespace
312+
{
313+
// Naive reference implementation of the non-collinear non-local stress
314+
// for a single element type starting at atom index 0.
315+
double ref_stress_nl_nc(const int nbands_occ,
316+
const int nkb,
317+
const int natom,
318+
const int nproj,
319+
const int deeq_2,
320+
const int deeq_3,
321+
const int deeq_4,
322+
const std::vector<double>& d_wg,
323+
const std::vector<std::complex<double>>& deeq_nc,
324+
const std::vector<std::complex<double>>& becp,
325+
const std::vector<std::complex<double>>& dbecp)
326+
{
327+
double ref = 0.0;
328+
for (int ib = 0; ib < nbands_occ; ib++)
329+
{
330+
const double fac = d_wg[ib];
331+
const int ib2 = ib * 2;
332+
for (int ia = 0; ia < natom; ia++)
333+
{
334+
for (int ip1 = 0; ip1 < nproj; ip1++)
335+
{
336+
for (int ip2 = 0; ip2 < nproj; ip2++)
337+
{
338+
const std::complex<double> ps0 = deeq_nc[((0 * deeq_2 + ia) * deeq_3 + ip1) * deeq_4 + ip2];
339+
const std::complex<double> ps1 = deeq_nc[((1 * deeq_2 + ia) * deeq_3 + ip1) * deeq_4 + ip2];
340+
const std::complex<double> ps2 = deeq_nc[((2 * deeq_2 + ia) * deeq_3 + ip1) * deeq_4 + ip2];
341+
const std::complex<double> ps3 = deeq_nc[((3 * deeq_2 + ia) * deeq_3 + ip1) * deeq_4 + ip2];
342+
const int inkb1 = ia * nproj + ip1;
343+
const int inkb2 = ia * nproj + ip2;
344+
const std::complex<double> dbb0 = std::conj(dbecp[ib2 * nkb + inkb1]) * becp[ib2 * nkb + inkb2];
345+
const std::complex<double> dbb1 = std::conj(dbecp[ib2 * nkb + inkb1]) * becp[(ib2 + 1) * nkb + inkb2];
346+
const std::complex<double> dbb2 = std::conj(dbecp[(ib2 + 1) * nkb + inkb1]) * becp[ib2 * nkb + inkb2];
347+
const std::complex<double> dbb3
348+
= std::conj(dbecp[(ib2 + 1) * nkb + inkb1]) * becp[(ib2 + 1) * nkb + inkb2];
349+
ref -= fac * (ps0 * dbb0 + ps1 * dbb1 + ps2 * dbb2 + ps3 * dbb3).real();
350+
}
351+
}
352+
}
353+
}
354+
return ref;
355+
}
356+
357+
// Deterministic, non-symmetric input data so that any wrong atom or
358+
// spin-block index gives a different result.
359+
void init_nc_inputs(std::vector<std::complex<double>>& deeq_nc,
360+
std::vector<std::complex<double>>& becp,
361+
std::vector<std::complex<double>>& dbecp)
362+
{
363+
for (size_t i = 0; i < deeq_nc.size(); i++)
364+
{
365+
deeq_nc[i] = std::complex<double>(0.11 * i + 0.03, -0.07 * i + 0.02);
366+
}
367+
for (size_t i = 0; i < becp.size(); i++)
368+
{
369+
becp[i] = std::complex<double>(0.05 * i - 0.31, 0.11 * i + 0.13);
370+
}
371+
for (size_t i = 0; i < dbecp.size(); i++)
372+
{
373+
dbecp[i] = std::complex<double>(-0.06 * i + 0.21, 0.04 * i - 0.52);
374+
}
375+
}
376+
} // namespace
377+
378+
TEST(TestSrcPWStressMultiDevice, cal_stress_nl_nc_op_cpu)
379+
{
380+
const int ipol = 0, jpol = 1;
381+
const int nkb = 4, nbands_occ = 2, ntype = 1;
382+
const int natom = 2, nproj = 2;
383+
const int deeq_2 = natom, deeq_3 = nproj, deeq_4 = nproj;
384+
385+
std::vector<int> atom_na{natom};
386+
std::vector<int> atom_nh{nproj};
387+
388+
std::vector<double> d_wg{0.71, 1.33};
389+
std::vector<double> qq_nt(1, 0.0); // unused: d_ekb is nullptr
390+
391+
std::vector<std::complex<double>> deeq_nc(4 * deeq_2 * deeq_3 * deeq_4);
392+
std::vector<std::complex<double>> becp(nbands_occ * 2 * nkb);
393+
std::vector<std::complex<double>> dbecp(nbands_occ * 2 * nkb);
394+
init_nc_inputs(deeq_nc, becp, dbecp);
395+
396+
const double expected = ref_stress_nl_nc(nbands_occ,
397+
nkb,
398+
natom,
399+
nproj,
400+
deeq_2,
401+
deeq_3,
402+
deeq_4,
403+
d_wg,
404+
deeq_nc,
405+
becp,
406+
dbecp);
407+
408+
std::vector<double> stress(9, 0.0);
409+
hamilt::cal_stress_nl_op<double, base_device::DEVICE_CPU>()(cpu_ctx,
410+
ipol,
411+
jpol,
412+
nkb,
413+
nbands_occ,
414+
ntype,
415+
deeq_2,
416+
deeq_3,
417+
deeq_4,
418+
atom_nh.data(),
419+
atom_na.data(),
420+
d_wg.data(),
421+
true,
422+
nullptr,
423+
qq_nt.data(),
424+
deeq_nc.data(),
425+
becp.data(),
426+
dbecp.data(),
427+
stress.data());
428+
429+
EXPECT_LT(fabs(stress[ipol * 3 + jpol] - expected), 1e-12);
430+
}
431+
432+
#if __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM
433+
TEST(TestSrcPWStressMultiDevice, cal_stress_nl_nc_op_gpu)
434+
{
435+
const int ipol = 0, jpol = 1;
436+
const int nkb = 4, nbands_occ = 2, ntype = 1;
437+
const int natom = 2, nproj = 2;
438+
const int deeq_2 = natom, deeq_3 = nproj, deeq_4 = nproj;
439+
440+
std::vector<int> atom_na{natom};
441+
std::vector<int> atom_nh{nproj};
442+
443+
std::vector<double> d_wg{0.71, 1.33};
444+
std::vector<double> qq_nt(1, 0.0); // unused: d_ekb is nullptr
445+
446+
std::vector<std::complex<double>> deeq_nc(4 * deeq_2 * deeq_3 * deeq_4);
447+
std::vector<std::complex<double>> becp(nbands_occ * 2 * nkb);
448+
std::vector<std::complex<double>> dbecp(nbands_occ * 2 * nkb);
449+
init_nc_inputs(deeq_nc, becp, dbecp);
450+
451+
const double expected = ref_stress_nl_nc(nbands_occ,
452+
nkb,
453+
natom,
454+
nproj,
455+
deeq_2,
456+
deeq_3,
457+
deeq_4,
458+
d_wg,
459+
deeq_nc,
460+
becp,
461+
dbecp);
462+
463+
std::vector<double> stress(9, 0.0);
464+
465+
using delmem_int_op = base_device::memory::delete_memory_op<int, base_device::DEVICE_GPU>;
466+
using resmem_int_op = base_device::memory::resize_memory_op<int, base_device::DEVICE_GPU>;
467+
using syncmem_int_h2d_op
468+
= base_device::memory::synchronize_memory_op<int, base_device::DEVICE_GPU, base_device::DEVICE_CPU>;
469+
470+
std::complex<double> *d_deeq_nc = nullptr, *d_becp = nullptr, *d_dbecp = nullptr;
471+
double *dev_wg = nullptr, *d_qq_nt = nullptr, *d_stress = nullptr;
472+
int *d_atom_nh = nullptr, *d_atom_na = nullptr;
473+
474+
resmem_zd_op()(d_deeq_nc, deeq_nc.size());
475+
resmem_zd_op()(d_becp, becp.size());
476+
resmem_zd_op()(d_dbecp, dbecp.size());
477+
syncmem_z2z_h2d_op()(d_deeq_nc, deeq_nc.data(), deeq_nc.size());
478+
syncmem_z2z_h2d_op()(d_becp, becp.data(), becp.size());
479+
syncmem_z2z_h2d_op()(d_dbecp, dbecp.data(), dbecp.size());
480+
481+
resmem_dd_op()(dev_wg, d_wg.size());
482+
resmem_dd_op()(d_qq_nt, qq_nt.size());
483+
resmem_dd_op()(d_stress, stress.size());
484+
syncmem_d2d_h2d_op()(dev_wg, d_wg.data(), d_wg.size());
485+
syncmem_d2d_h2d_op()(d_qq_nt, qq_nt.data(), qq_nt.size());
486+
syncmem_d2d_h2d_op()(d_stress, stress.data(), stress.size());
487+
488+
resmem_int_op()(d_atom_nh, atom_nh.size());
489+
resmem_int_op()(d_atom_na, atom_na.size());
490+
syncmem_int_h2d_op()(d_atom_nh, atom_nh.data(), atom_nh.size());
491+
syncmem_int_h2d_op()(d_atom_na, atom_na.data(), atom_na.size());
492+
493+
hamilt::cal_stress_nl_op<double, base_device::DEVICE_GPU>()(gpu_ctx,
494+
ipol,
495+
jpol,
496+
nkb,
497+
nbands_occ,
498+
ntype,
499+
deeq_2,
500+
deeq_3,
501+
deeq_4,
502+
d_atom_nh,
503+
d_atom_na,
504+
dev_wg,
505+
true,
506+
nullptr,
507+
d_qq_nt,
508+
d_deeq_nc,
509+
d_becp,
510+
d_dbecp,
511+
d_stress);
512+
513+
syncmem_d2d_d2h_op()(stress.data(), d_stress, stress.size());
514+
515+
EXPECT_LT(fabs(stress[ipol * 3 + jpol] - expected), 1e-12);
516+
517+
delmem_zd_op()(d_deeq_nc);
518+
delmem_zd_op()(d_becp);
519+
delmem_zd_op()(d_dbecp);
520+
521+
delmem_dd_op()(dev_wg);
522+
delmem_dd_op()(d_qq_nt);
523+
delmem_dd_op()(d_stress);
524+
525+
delmem_int_op()(d_atom_nh);
526+
delmem_int_op()(d_atom_na);
527+
}
528+
#endif // __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM

0 commit comments

Comments
 (0)