Skip to content

Commit b75d9d8

Browse files
committed
Complete PPCG pool reductions
1 parent 7f566b5 commit b75d9d8

5 files changed

Lines changed: 36 additions & 38 deletions

File tree

source/source_hsolver/ppcg/diago_ppcg_cg.hpp

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -39,11 +39,9 @@ void DiagoPPCG<T, Device>::orth_gradient(
3939
for (int i = 0; i < n_band_; ++i)
4040
{
4141
// Full complex inner product <psi_i | grad_j>
42-
T coeff = 0;
4342
const T* pi = psi + i * ld_psi_;
4443
const T* gj = grad.data() + j * ld_psi_;
45-
for (int ig = 0; ig < n_dim_; ++ig)
46-
coeff += std::conj(pi[ig]) * gj[ig];
44+
const T coeff = complex_dot(pi, gj);
4745
if (std::abs(coeff) <= std::numeric_limits<Real>::epsilon())
4846
continue;
4947
// grad_j -= S|psi_i> * coeff
@@ -102,6 +100,8 @@ void DiagoPPCG<T, Device>::update_polak_ribiere(
102100
beta_num_zr += static_cast<Real>(std::real(z * std::conj(g[ig])));
103101
beta_num_zo += static_cast<Real>(std::real(z * std::conj(r_old)));
104102
}
103+
reduce_pool_if_mpi_ready(beta_num_zr);
104+
reduce_pool_if_mpi_ready(beta_num_zo);
105105

106106
Real beta = 0;
107107
const Real denom = beta_denom[j];

source/source_hsolver/ppcg/diago_ppcg_diag.hpp

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,7 @@ double DiagoPPCG<T, Device>::diag(const HPsiFunc& hpsi_func,
288288
for (int ig = 0; ig < n_dim_; ++ig)
289289
nrm2 += static_cast<Real>(
290290
std::norm(grad[idx(ig, i, ld_psi_)]));
291+
reduce_pool_if_mpi_ready(nrm2);
291292
if (std::sqrt(nrm2) > std::max(static_cast<Real>(ethr_band[i]),
292293
diag_thr_))
293294
{

source/source_hsolver/ppcg/diago_ppcg_lapack.hpp

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
#include <ATen/kernels/lapack.h>
22

3+
#include "source_base/parallel_reduce.h"
4+
35
#include <cstdlib>
46
#include <fstream>
57

@@ -10,6 +12,32 @@ namespace hsolver {
1012
// =============================================================================
1113
namespace {
1214

15+
template <typename Value>
16+
void reduce_pool_if_mpi_ready(Value& value)
17+
{
18+
#ifdef __MPI
19+
int initialized = 0;
20+
int finalized = 0;
21+
MPI_Initialized(&initialized);
22+
MPI_Finalized(&finalized);
23+
if (initialized && !finalized)
24+
Parallel_Reduce::reduce_pool(value);
25+
#endif
26+
}
27+
28+
template <typename Value>
29+
void reduce_pool_if_mpi_ready(Value* value, const int n)
30+
{
31+
#ifdef __MPI
32+
int initialized = 0;
33+
int finalized = 0;
34+
MPI_Initialized(&initialized);
35+
MPI_Finalized(&finalized);
36+
if (initialized && !finalized)
37+
Parallel_Reduce::reduce_pool(value, n);
38+
#endif
39+
}
40+
1341
template <typename T, typename Real>
1442
Real max_generalized_residual(
1543
const T* hpsi,
@@ -28,6 +56,7 @@ Real max_generalized_residual(
2856
const T r = hpsi[ig + j * ld] - T(eigenvalue[j]) * spsi[ig + j * ld];
2957
nrm2 += static_cast<Real>(std::norm(r));
3058
}
59+
reduce_pool_if_mpi_ready(nrm2);
3160
max_res = std::max(max_res, std::sqrt(nrm2));
3261
}
3362
return max_res;

source/source_hsolver/ppcg/diago_ppcg_ops.hpp

Lines changed: 1 addition & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,38 +1,6 @@
11
#include "source_base/kernels/math_kernel_op.h"
2-
#include "source_base/parallel_reduce.h"
3-
42
namespace hsolver {
53

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-
364
// =============================================================================
375
// Constructor
386
// =============================================================================
@@ -220,11 +188,9 @@ void DiagoPPCG<T, Device>::project_against(
220188
for (const int bc : basis_cols)
221189
{
222190
// Full complex inner product <basis_bc | sx_c>
223-
T coeff = 0;
224191
const T* bb = basis + bc * ld_psi_;
225192
const T* sc = sx.data() + c * ld_psi_;
226-
for (int ig = 0; ig < n_dim_; ++ig)
227-
coeff += std::conj(bb[ig]) * sc[ig];
193+
const T coeff = complex_dot(bb, sc);
228194
if (std::abs(coeff) <= std::numeric_limits<Real>::epsilon())
229195
continue;
230196
const T* sb = sbasis + bc * ld_psi_;

source/source_hsolver/ppcg/diago_ppcg_subspace.hpp

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ void DiagoPPCG<T, Device>::lock_epairs(
1919
Real nrm2 = 0;
2020
for (int ig = 0; ig < n_dim_; ++ig)
2121
nrm2 += static_cast<Real>(std::norm(residual[idx(ig, j, ld_psi_)]));
22+
reduce_pool_if_mpi_ready(nrm2);
2223
const Real rnrm = std::sqrt(std::max(nrm2, static_cast<Real>(0)));
2324
const Real thr = std::max(static_cast<Real>(ethr_band[j]), diag_thr_);
2425
if (rnrm > thr)
@@ -80,6 +81,7 @@ void DiagoPPCG<T, Device>::build_small_subspace(
8081
for (int ig = 0; ig < n_dim_; ++ig)
8182
sn2 += std::real(std::conj(x[idx(ig, j, ld_psi_)])
8283
* sx[idx(ig, j, ld_psi_)]);
84+
reduce_pool_if_mpi_ready(sn2);
8385
Real sn = std::sqrt(std::max(sn2, static_cast<Real>(1e-30)));
8486
// Only scale if the norm is non-negligible; a near-zero
8587
// column is a converged band whose contribution is harmless.

0 commit comments

Comments
 (0)