Skip to content

Commit 0e210f8

Browse files
committed
Fix: preserve exact-origin spherical harmonics in GPU Gint
1 parent c1790ea commit 0e210f8

3 files changed

Lines changed: 60 additions & 30 deletions

File tree

‎source/source_base/kernels/cuda/sph_harm_gpu.cuh‎

Lines changed: 39 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -4,36 +4,15 @@
44

55
namespace ModuleBase {
66

7-
/// Spherical harmonics computation (table lookup method)
8-
/// Directly uses constexpr ylmcoef, compiler auto-inlines
9-
/// @param nwl Maximum angular momentum L (0 <= nwl <= 5)
10-
/// @param x,y,z Direction vector (need not be normalized, normalization is done internally)
11-
/// @param ylma Output array, size (nwl+1)^2
12-
__device__ static void sph_harm(
7+
/// Evaluate the existing spherical-harmonic recurrence directly.
8+
/// This helper performs no input normalization and no zero-vector fallback.
9+
__device__ static void sph_harm_direct(
1310
const int nwl,
14-
const double x_in,
15-
const double y_in,
16-
const double z_in,
11+
const double x,
12+
const double y,
13+
const double z,
1714
double* __restrict__ ylma)
1815
{
19-
// Normalize the input direction vector
20-
double r = sqrt(x_in * x_in + y_in * y_in + z_in * z_in);
21-
double x, y, z;
22-
if (r < 1e-10)
23-
{
24-
// At origin, default to z-axis direction
25-
x = 0.0;
26-
y = 0.0;
27-
z = 1.0;
28-
}
29-
else
30-
{
31-
const double inv_r = 1.0 / r;
32-
x = x_in * inv_r;
33-
y = y_in * inv_r;
34-
z = z_in * inv_r;
35-
}
36-
3716
/***************************
3817
L = 0
3918
***************************/
@@ -147,6 +126,39 @@ __device__ static void sph_harm(
147126
return;
148127
}
149128

129+
/// Spherical harmonics computation (table lookup method)
130+
/// Directly uses constexpr ylmcoef, compiler auto-inlines
131+
/// @param nwl Maximum angular momentum L (0 <= nwl <= 5)
132+
/// @param x,y,z Direction vector (need not be normalized, normalization is done internally)
133+
/// @param ylma Output array, size (nwl+1)^2
134+
__device__ static void sph_harm(
135+
const int nwl,
136+
const double x_in,
137+
const double y_in,
138+
const double z_in,
139+
double* __restrict__ ylma)
140+
{
141+
// Normalize the input direction vector
142+
double r = sqrt(x_in * x_in + y_in * y_in + z_in * z_in);
143+
double x, y, z;
144+
if (r < 1e-10)
145+
{
146+
// At origin, default to z-axis direction
147+
x = 0.0;
148+
y = 0.0;
149+
z = 1.0;
150+
}
151+
else
152+
{
153+
const double inv_r = 1.0 / r;
154+
x = x_in * inv_r;
155+
y = y_in * inv_r;
156+
z = z_in * inv_r;
157+
}
158+
159+
sph_harm_direct(nwl, x, y, z, ylma);
160+
}
161+
150162
/// Spherical harmonics and gradient computation
151163
__device__ static void grad_rl_sph_harm(
152164
const int nwl,

‎source/source_hamilt/module_gint/kernel/phi_operator_kernel.cuh‎

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,11 @@ __global__ void set_phi_kernel(
5353
const double3 coord = make_double3(mgrid_pos.x-rcoord.x, // coord is the relative coordinate of an atom and a meshgrid
5454
mgrid_pos.y-rcoord.y,
5555
mgrid_pos.z-rcoord.z);
56+
// Preserve the existing near-origin behavior. Only the exact
57+
// atomic grid point follows the CPU direct-recurrence semantics.
58+
const bool exact_origin
59+
= (coord.x == 0.0 && coord.y == 0.0 && coord.z == 0.0);
60+
5661
double dist = norm3d(coord.x, coord.y, coord.z);
5762
if (dist < rcut[atom_type])
5863
{
@@ -61,7 +66,20 @@ __global__ void set_phi_kernel(
6166
// since nwl is less or equal than 5, the size of ylma is (5+1)^2
6267
double ylma[36];
6368
const int nwl = ucell_atom_nwl[atom_type];
64-
sph_harm(nwl, coord.x/dist, coord.y/dist, coord.z/dist, ylma);
69+
if (exact_origin)
70+
{
71+
ModuleBase::sph_harm_direct(
72+
nwl, 0.0, 0.0, 0.0, ylma);
73+
}
74+
else
75+
{
76+
sph_harm(
77+
nwl,
78+
coord.x/dist,
79+
coord.y/dist,
80+
coord.z/dist,
81+
ylma);
82+
}
6583

6684
const double pos = dist / dr_uniform;
6785
const int ip = static_cast<int>(pos);

‎tests/03_NAO_multik/CASES_GPU.txt‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ scf_out_hsk
3333
scf_out_hsk_binary
3434
scf_out_hsr
3535
scf_out_hsr_binary_spin2
36-
#scf_out_hsr_spin4
36+
scf_out_hsr_spin4
3737
scf_out_dh_t
3838
scf_out_dos_spin4
3939
scf_out_mul
@@ -46,7 +46,7 @@ nscf_out_dos
4646
nscf_out_band_pband
4747
nscf_out_pot1
4848
nscf_out_mul
49-
#nscf_out_hsr_tr_rr
49+
nscf_out_hsr_tr_rr
5050
relax_bfgs2
5151
relax_old_cg
5252
relax_cell

0 commit comments

Comments
 (0)