Skip to content

Commit 7f566b5

Browse files
committed
Use BLAS and pool reductions in PPCG projections
1 parent 2c8bf32 commit 7f566b5

3 files changed

Lines changed: 155 additions & 39 deletions

File tree

source/source_hsolver/ppcg/diago_ppcg_ops.hpp

Lines changed: 51 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,38 @@
11
#include "source_base/kernels/math_kernel_op.h"
2+
#include "source_base/parallel_reduce.h"
23

34
namespace hsolver {
45

6+
namespace {
7+
8+
template <typename Value>
9+
void reduce_pool_if_mpi_ready(Value& value)
10+
{
11+
#ifdef __MPI
12+
int initialized = 0;
13+
int finalized = 0;
14+
MPI_Initialized(&initialized);
15+
MPI_Finalized(&finalized);
16+
if (initialized && !finalized)
17+
Parallel_Reduce::reduce_pool(value);
18+
#endif
19+
}
20+
21+
template <typename Value>
22+
void reduce_pool_if_mpi_ready(Value* value, const int n)
23+
{
24+
#ifdef __MPI
25+
int initialized = 0;
26+
int finalized = 0;
27+
MPI_Initialized(&initialized);
28+
MPI_Finalized(&finalized);
29+
if (initialized && !finalized)
30+
Parallel_Reduce::reduce_pool(value, n);
31+
#endif
32+
}
33+
34+
} // anonymous namespace
35+
536
// =============================================================================
637
// Constructor
738
// =============================================================================
@@ -93,7 +124,9 @@ template <typename T, typename Device>
93124
typename DiagoPPCG<T, Device>::Real
94125
DiagoPPCG<T, Device>::gamma_dot(const T* x, const T* y) const
95126
{
96-
return ModuleBase::dot_real_op<T, Device>()(n_dim_, x, y, false);
127+
Real result = ModuleBase::dot_real_op<T, Device>()(n_dim_, x, y, false);
128+
reduce_pool_if_mpi_ready(result);
129+
return result;
97130
}
98131

99132
template <typename T, typename Device>
@@ -102,6 +135,7 @@ T DiagoPPCG<T, Device>::complex_dot(const T* x, const T* y) const
102135
T acc = T(0);
103136
for (int i = 0; i < n_dim_; ++i)
104137
acc += std::conj(x[i]) * y[i];
138+
reduce_pool_if_mpi_ready(&acc, 1);
105139
return acc;
106140
}
107141

@@ -115,10 +149,22 @@ void DiagoPPCG<T, Device>::gram(const T* a, const T* b,
115149
int ld_out) const
116150
{
117151
out.assign(ld_out * ncol_b, T(0));
118-
for (int jb = 0; jb < ncol_b; ++jb)
119-
for (int ia = 0; ia < ncol_a; ++ia)
120-
out[ia + jb * ld_out] = complex_dot(a + ia * ld_psi_,
121-
b + jb * ld_psi_);
152+
const T one = T(1);
153+
const T zero = T(0);
154+
ModuleBase::gemm_op<T, Device>()('C',
155+
'N',
156+
ncol_a,
157+
ncol_b,
158+
n_dim_,
159+
&one,
160+
a,
161+
ld_psi_,
162+
b,
163+
ld_psi_,
164+
&zero,
165+
out.data(),
166+
ld_out);
167+
reduce_pool_if_mpi_ready(out.data(), ld_out * ncol_b);
122168
}
123169

124170
// =============================================================================

source/source_hsolver/ppcg/diago_ppcg_orth.hpp

Lines changed: 42 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -167,18 +167,50 @@ void DiagoPPCG<T, Device>::rayleigh_ritz(
167167
set_zero(spsi_);
168168
set_zero(hpsi_);
169169

170+
const T one = T(1);
171+
const T zero = T(0);
172+
ModuleBase::gemm_op<T, Device>()('N',
173+
'N',
174+
n_dim_,
175+
n_band_,
176+
n_band_,
177+
&one,
178+
psi_old.data(),
179+
ld_psi_,
180+
hsub.data(),
181+
n_band_,
182+
&zero,
183+
psi,
184+
ld_psi_);
185+
ModuleBase::gemm_op<T, Device>()('N',
186+
'N',
187+
n_dim_,
188+
n_band_,
189+
n_band_,
190+
&one,
191+
spsi_old.data(),
192+
ld_psi_,
193+
hsub.data(),
194+
n_band_,
195+
&zero,
196+
spsi_.data(),
197+
ld_psi_);
198+
ModuleBase::gemm_op<T, Device>()('N',
199+
'N',
200+
n_dim_,
201+
n_band_,
202+
n_band_,
203+
&one,
204+
hpsi_old.data(),
205+
ld_psi_,
206+
hsub.data(),
207+
n_band_,
208+
&zero,
209+
hpsi_.data(),
210+
ld_psi_);
211+
170212
for (int j = 0; j < n_band_; ++j)
171213
{
172-
for (int i = 0; i < n_band_; ++i)
173-
{
174-
const T c = hsub[i + j * n_band_];
175-
for (int ig = 0; ig < n_dim_; ++ig)
176-
{
177-
psi[ idx(ig, j, ld_psi_)] += psi_old[ idx(ig, i, ld_psi_)] * c;
178-
spsi_[idx(ig, j, ld_psi_)] += spsi_old[idx(ig, i, ld_psi_)] * c;
179-
hpsi_[idx(ig, j, ld_psi_)] += hpsi_old[idx(ig, i, ld_psi_)] * c;
180-
}
181-
}
182214
eigenvalue[j] = eval[j];
183215
}
184216
}

source/source_hsolver/ppcg/diago_ppcg_subspace.hpp

Lines changed: 62 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -213,41 +213,79 @@ void DiagoPPCG<T, Device>::update_one_block(
213213
std::vector<T> sp_new(ld_psi_ * l, T(0));
214214
std::vector<T> hp_new(ld_psi_ * l, T(0));
215215

216+
std::vector<T> coeff_state(dim * l, T(0));
217+
std::vector<T> coeff_dir(dim * l, T(0));
216218
for (int j = 0; j < l; ++j)
217219
{
218220
for (int i = 0; i < l; ++i)
219221
{
220-
const T cpsi = eigvec[i + j * dim];
221-
const T cw = eigvec[(l + i) + j * dim] * subspace.w_scale[i];
222-
223-
for (int ig = 0; ig < n_dim_; ++ig)
222+
coeff_state[i + j * dim] = eigvec[i + j * dim];
223+
const T cw = eigvec[(l + i) + j * dim] * subspace.w_scale[i];
224+
coeff_state[(l + i) + j * dim] = cw;
225+
coeff_dir[(l + i) + j * dim] = cw;
226+
if (use_p)
224227
{
225-
psi_new[idx(ig, j, ld_psi_)] += psi_l[idx(ig, i, ld_psi_)] * cpsi
226-
+ w_l[ idx(ig, i, ld_psi_)] * cw;
227-
spsi_new[idx(ig, j, ld_psi_)] += spsi_l[idx(ig, i, ld_psi_)] * cpsi
228-
+ sw_l[ idx(ig, i, ld_psi_)] * cw;
229-
hpsi_new[idx(ig, j, ld_psi_)] += hpsi_l[idx(ig, i, ld_psi_)] * cpsi
230-
+ hw_l[ idx(ig, i, ld_psi_)] * cw;
231-
p_new[idx(ig, j, ld_psi_)] += w_l[ idx(ig, i, ld_psi_)] * cw;
232-
sp_new[idx(ig, j, ld_psi_)] += sw_l[ idx(ig, i, ld_psi_)] * cw;
233-
hp_new[idx(ig, j, ld_psi_)] += hw_l[ idx(ig, i, ld_psi_)] * cw;
228+
const T cp = eigvec[(2*l + i) + j * dim] * subspace.p_scale[i];
229+
coeff_state[(2*l + i) + j * dim] = cp;
230+
coeff_dir[(2*l + i) + j * dim] = cp;
234231
}
232+
}
233+
}
235234

235+
auto fill_basis = [&](const std::vector<T>& a,
236+
const std::vector<T>& b,
237+
const std::vector<T>& c,
238+
std::vector<T>& basis)
239+
{
240+
basis.assign(ld_psi_ * dim, T(0));
241+
for (int j = 0; j < l; ++j)
242+
{
243+
std::copy(a.begin() + j * ld_psi_,
244+
a.begin() + (j + 1) * ld_psi_,
245+
basis.begin() + j * ld_psi_);
246+
std::copy(b.begin() + j * ld_psi_,
247+
b.begin() + (j + 1) * ld_psi_,
248+
basis.begin() + (l + j) * ld_psi_);
236249
if (use_p)
237250
{
238-
const T cp = eigvec[(2*l + i) + j * dim] * subspace.p_scale[i];
239-
for (int ig = 0; ig < n_dim_; ++ig)
240-
{
241-
psi_new[idx(ig, j, ld_psi_)] += p_l[ idx(ig, i, ld_psi_)] * cp;
242-
spsi_new[idx(ig, j, ld_psi_)] += sp_l[idx(ig, i, ld_psi_)] * cp;
243-
hpsi_new[idx(ig, j, ld_psi_)] += hp_l[idx(ig, i, ld_psi_)] * cp;
244-
p_new[idx(ig, j, ld_psi_)] += p_l[ idx(ig, i, ld_psi_)] * cp;
245-
sp_new[idx(ig, j, ld_psi_)] += sp_l[idx(ig, i, ld_psi_)] * cp;
246-
hp_new[idx(ig, j, ld_psi_)] += hp_l[idx(ig, i, ld_psi_)] * cp;
247-
}
251+
std::copy(c.begin() + j * ld_psi_,
252+
c.begin() + (j + 1) * ld_psi_,
253+
basis.begin() + (2 * l + j) * ld_psi_);
248254
}
249255
}
250-
}
256+
};
257+
258+
auto combine = [&](const std::vector<T>& a,
259+
const std::vector<T>& b,
260+
const std::vector<T>& c,
261+
const std::vector<T>& coeff,
262+
std::vector<T>& out)
263+
{
264+
std::vector<T> basis;
265+
fill_basis(a, b, c, basis);
266+
const T one = T(1);
267+
const T zero = T(0);
268+
ModuleBase::gemm_op<T, Device>()('N',
269+
'N',
270+
n_dim_,
271+
l,
272+
dim,
273+
&one,
274+
basis.data(),
275+
ld_psi_,
276+
coeff.data(),
277+
dim,
278+
&zero,
279+
out.data(),
280+
ld_psi_);
281+
};
282+
283+
combine(psi_l, w_l, p_l, coeff_state, psi_new);
284+
combine(spsi_l, sw_l, sp_l, coeff_state, spsi_new);
285+
combine(hpsi_l, hw_l, hp_l, coeff_state, hpsi_new);
286+
combine(psi_l, w_l, p_l, coeff_dir, p_new);
287+
combine(spsi_l, sw_l, sp_l, coeff_dir, sp_new);
288+
combine(hpsi_l, hw_l, hp_l, coeff_dir, hp_new);
251289

252290
scatter_cols(psi, cols, psi_new);
253291
scatter_cols(spsi_.data(), cols, spsi_new);

0 commit comments

Comments
 (0)