Skip to content

Commit ee607b1

Browse files
committed
Reuse PPCG Rayleigh-Ritz workspaces
1 parent d33ddb0 commit ee607b1

3 files changed

Lines changed: 32 additions & 22 deletions

File tree

source/source_hsolver/diago_ppcg.h

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,12 @@ class DiagoPPCG
9898
std::vector<T> p_; // previous search direction (for block subspace)
9999
std::vector<T> sp_; // S * p
100100
std::vector<T> hp_; // H * p
101+
std::vector<T> rr_psi_; // Rayleigh-Ritz rotation workspace
102+
std::vector<T> rr_spsi_;
103+
std::vector<T> rr_hpsi_;
104+
std::vector<T> rr_hsub_;
105+
std::vector<T> rr_ssub_;
106+
std::vector<Real> rr_eval_;
101107

102108
// Polak-Ribiere state (CONJUGATE_GRADIENT strategy)
103109
std::vector<T> grad_old_; // previous gradient

source/source_hsolver/ppcg/diago_ppcg_diag.hpp

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,12 @@ double DiagoPPCG<T, Device>::diag(const HPsiFunc& hpsi_func,
3333
p_.assign(sz, T(0));
3434
sp_.assign(sz, T(0));
3535
hp_.assign(sz, T(0));
36+
rr_psi_.resize(sz);
37+
rr_spsi_.resize(sz);
38+
rr_hpsi_.resize(sz);
39+
rr_hsub_.resize(ncol * ncol);
40+
rr_ssub_.resize(ncol * ncol);
41+
rr_eval_.resize(ncol);
3642

3743
std::vector<int> all_cols(ncol);
3844
std::iota(all_cols.begin(), all_cols.end(), 0);

source/source_hsolver/ppcg/diago_ppcg_orth.hpp

Lines changed: 20 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -141,37 +141,35 @@ void DiagoPPCG<T, Device>::rayleigh_ritz(
141141
std::vector<int>& active_cols,
142142
const std::vector<double>& ethr_band)
143143
{
144-
std::vector<T> hsub(n_band_ * n_band_, T(0));
145-
std::vector<T> ssub(n_band_ * n_band_, T(0));
146-
gram(psi, hpsi_.data(), n_band_, n_band_, hsub, n_band_);
147-
gram(psi, spsi_.data(), n_band_, n_band_, ssub, n_band_);
144+
gram(psi, hpsi_.data(), n_band_, n_band_, rr_hsub_, n_band_);
145+
gram(psi, spsi_.data(), n_band_, n_band_, rr_ssub_, n_band_);
148146

149-
std::vector<Real> eval(n_band_, static_cast<Real>(0));
150147
bool sygvd_ok = false;
151148
try
152149
{
153-
HermitianLapack<T>::sygvd(n_band_, hsub.data(), ssub.data(),
154-
eval.data());
150+
HermitianLapack<T>::sygvd(n_band_, rr_hsub_.data(), rr_ssub_.data(),
151+
rr_eval_.data());
155152
sygvd_ok = true;
156153
}
157154
catch (const std::runtime_error&)
158155
{
159156
// Fallback: diagonal Rayleigh quotients.
160157
// hsub and ssub may be corrupted by sygvd; re-form them.
161-
gram(psi, hpsi_.data(), n_band_, n_band_, hsub, n_band_);
162-
gram(psi, spsi_.data(), n_band_, n_band_, ssub, n_band_);
158+
gram(psi, hpsi_.data(), n_band_, n_band_, rr_hsub_, n_band_);
159+
gram(psi, spsi_.data(), n_band_, n_band_, rr_ssub_, n_band_);
163160
for (int ii = 0; ii < n_band_; ++ii)
164-
eval[ii] = static_cast<Real>(std::real(hsub[ii + ii * n_band_]))
161+
rr_eval_[ii] = static_cast<Real>(std::real(rr_hsub_[ii + ii * n_band_]))
165162
/ std::max(static_cast<Real>(
166-
std::real(ssub[ii + ii * n_band_])),
163+
std::real(rr_ssub_[ii + ii * n_band_])),
167164
static_cast<Real>(1e-30));
168165
}
169166

170167
if (sygvd_ok)
171168
{
172-
std::vector<T> psi_old(psi, psi + ld_psi_ * n_band_);
173-
std::vector<T> spsi_old = spsi_;
174-
std::vector<T> hpsi_old = hpsi_;
169+
const int sz = ld_psi_ * n_band_;
170+
std::copy(psi, psi + sz, rr_psi_.begin());
171+
std::copy(spsi_.begin(), spsi_.end(), rr_spsi_.begin());
172+
std::copy(hpsi_.begin(), hpsi_.end(), rr_hpsi_.begin());
175173

176174
std::fill(psi, psi + ld_psi_ * n_band_, T(0));
177175
set_zero(spsi_);
@@ -185,9 +183,9 @@ void DiagoPPCG<T, Device>::rayleigh_ritz(
185183
n_band_,
186184
n_band_,
187185
&one,
188-
psi_old.data(),
186+
rr_psi_.data(),
189187
ld_psi_,
190-
hsub.data(),
188+
rr_hsub_.data(),
191189
n_band_,
192190
&zero,
193191
psi,
@@ -198,9 +196,9 @@ void DiagoPPCG<T, Device>::rayleigh_ritz(
198196
n_band_,
199197
n_band_,
200198
&one,
201-
spsi_old.data(),
199+
rr_spsi_.data(),
202200
ld_psi_,
203-
hsub.data(),
201+
rr_hsub_.data(),
204202
n_band_,
205203
&zero,
206204
spsi_.data(),
@@ -211,24 +209,24 @@ void DiagoPPCG<T, Device>::rayleigh_ritz(
211209
n_band_,
212210
n_band_,
213211
&one,
214-
hpsi_old.data(),
212+
rr_hpsi_.data(),
215213
ld_psi_,
216-
hsub.data(),
214+
rr_hsub_.data(),
217215
n_band_,
218216
&zero,
219217
hpsi_.data(),
220218
ld_psi_);
221219

222220
for (int j = 0; j < n_band_; ++j)
223221
{
224-
eigenvalue[j] = eval[j];
222+
eigenvalue[j] = rr_eval_[j];
225223
}
226224
}
227225
else
228226
{
229227
// No rotation: just update eigenvalues with Rayleigh quotients.
230228
for (int j = 0; j < n_band_; ++j)
231-
eigenvalue[j] = eval[j];
229+
eigenvalue[j] = rr_eval_[j];
232230
}
233231

234232
// Compute residual: w_i = H|psi_i> - eps_i * S|psi_i>

0 commit comments

Comments
 (0)