Skip to content

Commit feaa838

Browse files
authored
Merge pull request #6045 from ye-luo/fix-LMY
Protect DescentEngine object exposure
2 parents ee36d7d + cfbc892 commit feaa838

10 files changed

Lines changed: 45 additions & 48 deletions

src/QMCDrivers/WFOpt/QMCCostFunction.cpp

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -346,14 +346,14 @@ void QMCCostFunction::checkConfigurations(EngineHandle& handle)
346346
*In future, both the LM and descent engines should be children of some parent engine base class.
347347
* */
348348
void QMCCostFunction::engine_checkConfigurations(cqmc::engine::LMYEngine<Return_t>& EngineObj,
349-
DescentEngine& descentEngineObj,
350-
const std::string& MinMethod)
349+
OptionalRef<DescentEngine> descentEngineObj)
351350
{
352351
const auto num_opt_vars = opt_vars.size();
353-
if (MinMethod == "descent")
352+
if (descentEngineObj)
354353
{
354+
DescentEngine& descent_engine(*descentEngineObj);
355355
//Reset vectors and scalars from any previous iteration
356-
descentEngineObj.prepareStorage(omp_get_max_threads(), num_opt_vars);
356+
descent_engine.prepareStorage(omp_get_max_threads(), num_opt_vars);
357357
}
358358
RealType et_tot = 0.0;
359359
RealType e2_tot = 0.0;
@@ -426,20 +426,18 @@ void QMCCostFunction::engine_checkConfigurations(cqmc::engine::LMYEngine<Return_
426426
le_der_samp[i + 1] = HDsaved[i] + etmp * Dsaved[i];
427427

428428
#ifdef HAVE_LMY_ENGINE
429-
if (MinMethod == "adaptive")
430-
{
431-
// pass into engine
432-
EngineObj.take_sample(der_rat_samp, le_der_samp, le_der_samp, 1.0, saved[REWEIGHT]);
433-
}
434-
else if (MinMethod == "descent")
429+
if (descentEngineObj)
435430
{
431+
DescentEngine& descent_engine(*descentEngineObj);
436432
//Could remove this copying over if LM engine becomes compatible with complex numbers
437433
//so that der_rat_samp and le_der_samp are vectors of std::complex<double> when QMC_COMPLEX=1
438434
std::vector<FullPrecValueType> der_rat_samp_comp(der_rat_samp.begin(), der_rat_samp.end());
439435
std::vector<FullPrecValueType> le_der_samp_comp(le_der_samp.begin(), le_der_samp.end());
440436

441-
descentEngineObj.takeSample(ip, der_rat_samp_comp, le_der_samp_comp, le_der_samp_comp, 1.0, saved[REWEIGHT]);
437+
descent_engine.takeSample(ip, der_rat_samp_comp, le_der_samp_comp, le_der_samp_comp, 1.0, saved[REWEIGHT]);
442438
}
439+
else
440+
EngineObj.take_sample(der_rat_samp, le_der_samp, le_der_samp, 1.0, saved[REWEIGHT]);
443441
#endif
444442
}
445443
else
@@ -477,10 +475,13 @@ void QMCCostFunction::engine_checkConfigurations(cqmc::engine::LMYEngine<Return_
477475

478476
#ifdef HAVE_LMY_ENGINE
479477
// engine finish taking samples
480-
if (MinMethod == "adaptive")
478+
if (descentEngineObj)
479+
{
480+
DescentEngine& descent_engine(*descentEngineObj);
481+
descent_engine.sample_finish();
482+
}
483+
else
481484
EngineObj.sample_finish();
482-
else if (MinMethod == "descent")
483-
descentEngineObj.sample_finish();
484485
#endif
485486

486487
app_log().flush();

src/QMCDrivers/WFOpt/QMCCostFunction.h

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -41,8 +41,7 @@ class QMCCostFunction : public QMCCostFunctionBase, public CloneManager
4141
void checkConfigurations(EngineHandle& handle) override;
4242
#ifdef HAVE_LMY_ENGINE
4343
void engine_checkConfigurations(cqmc::engine::LMYEngine<Return_t>& EngineObj,
44-
DescentEngine& descentEngineObj,
45-
const std::string& MinMethod) override;
44+
OptionalRef<DescentEngine> descentEngineObj) override;
4645
#endif
4746

4847

src/QMCDrivers/WFOpt/QMCCostFunctionBase.h

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
#include "QMCWaveFunctions/TrialWaveFunction.h"
2323
#include "Message/MPIObjectBase.h"
2424
#include "libxml/xpath.h"
25+
#include "type_traits/OptionalRef.hpp"
2526

2627
#ifdef HAVE_LMY_ENGINE
2728
#include "formic/utils/matrix.h"
@@ -158,9 +159,11 @@ class QMCCostFunctionBase : public MPIObjectBase
158159
//for SR method
159160
virtual void checkConfigurationsSR(EngineHandle& handle);
160161
#ifdef HAVE_LMY_ENGINE
162+
/** similar to checkConfigurations. With additioal interaction with the LMY engine.
163+
* if descentEngineObj is not nullopt, collect results to the descentEngineObj.
164+
*/
161165
virtual void engine_checkConfigurations(cqmc::engine::LMYEngine<Return_t>& EngineObj,
162-
DescentEngine& descentEngineObj,
163-
const std::string& MinMethod) = 0;
166+
OptionalRef<DescentEngine> descentEngineObj) = 0;
164167

165168
#endif
166169

src/QMCDrivers/WFOpt/QMCCostFunctionBatched.cpp

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -627,8 +627,7 @@ void QMCCostFunctionBatched::checkConfigurationsSR(EngineHandle& handle)
627627

628628
#ifdef HAVE_LMY_ENGINE
629629
void QMCCostFunctionBatched::engine_checkConfigurations(cqmc::engine::LMYEngine<Return_t>& EngineObj,
630-
DescentEngine& descentEngineObj,
631-
const std::string& MinMethod)
630+
OptionalRef<DescentEngine> descentEngineObj)
632631
{ APP_ABORT("LMYEngine not implemented with batch optimization"); }
633632
#endif
634633

src/QMCDrivers/WFOpt/QMCCostFunctionBatched.h

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -55,8 +55,7 @@ class QMCCostFunctionBatched : public QMCCostFunctionBase, public QMCTraits
5555
void checkConfigurationsSR(EngineHandle& handle) override;
5656
#ifdef HAVE_LMY_ENGINE
5757
void engine_checkConfigurations(cqmc::engine::LMYEngine<Return_t>& EngineObj,
58-
DescentEngine& descentEngineObj,
59-
const std::string& MinMethod) override;
58+
OptionalRef<DescentEngine> descentEngineObj) override;
6059
#endif
6160

6261

src/QMCDrivers/WFOpt/QMCFixedSampleLinearOptimize.cpp

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -981,7 +981,7 @@ bool QMCFixedSampleLinearOptimize::adaptive_three_shift_run()
981981
EngineObj->reset();
982982

983983
// generate samples and compute weights, local energies, and derivative vectors
984-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
984+
engine_start();
985985

986986
// get dimension of the linear method matrices
987987
size_t N = numParams + 1;
@@ -1018,7 +1018,7 @@ bool QMCFixedSampleLinearOptimize::adaptive_three_shift_run()
10181018
finish();
10191019

10201020
// take sample
1021-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1021+
engine_start();
10221022
}
10231023

10241024
// say what we are doing
@@ -1396,7 +1396,7 @@ bool QMCFixedSampleLinearOptimize::descent_run()
13961396
optTarget->setneedGrads(true);
13971397

13981398
//Compute Lagrangian derivatives needed for parameter updates with engine_checkConfigurations, which is called inside engine_start
1399-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1399+
engine_start();
14001400

14011401
int descent_num = descentEngineObj->getDescentNum();
14021402

@@ -1440,7 +1440,7 @@ bool QMCFixedSampleLinearOptimize::descent_run()
14401440
#ifdef HAVE_LMY_ENGINE
14411441
bool QMCFixedSampleLinearOptimize::hybrid_run()
14421442
{
1443-
app_log() << "This is methodName: " << MinMethod << std::endl;
1443+
app_log() << "This method name is: " << MinMethod << std::endl;
14441444

14451445
//Either the adaptive BLM or descent optimization is run
14461446

@@ -1518,9 +1518,7 @@ void QMCFixedSampleLinearOptimize::start()
15181518
}
15191519

15201520
#ifdef HAVE_LMY_ENGINE
1521-
void QMCFixedSampleLinearOptimize::engine_start(cqmc::engine::LMYEngine<ValueType>& EngineObj,
1522-
DescentEngine& descentEngineObj,
1523-
std::string MinMethod)
1521+
void QMCFixedSampleLinearOptimize::engine_start()
15241522
{
15251523
app_log() << "entering engine_start function" << std::endl;
15261524

@@ -1545,8 +1543,10 @@ void QMCFixedSampleLinearOptimize::engine_start(cqmc::engine::LMYEngine<ValueTyp
15451543
Timer t2;
15461544
optTarget->getConfigurations(h5FileRoot);
15471545
optTarget->setRng(vmcEngine->getRngRefs());
1548-
optTarget->engine_checkConfigurations(EngineObj, descentEngineObj,
1549-
MinMethod); // computes derivative ratios and pass into engine
1546+
// computes derivative ratios and pass into engine
1547+
optTarget->engine_checkConfigurations(*EngineObj,
1548+
MinMethod == "descent" ? makeOptionalRef<DescentEngine>(*descentEngineObj)
1549+
: std::nullopt);
15501550
app_log() << " Execution time = " << std::setprecision(4) << t2.elapsed() << std::endl;
15511551
}
15521552
app_log() << " </log>" << std::endl;

src/QMCDrivers/WFOpt/QMCFixedSampleLinearOptimize.h

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -233,9 +233,7 @@ class QMCFixedSampleLinearOptimize : public QMCDriver, public LinearMethod, priv
233233
///common operation to start optimization, used by the derived classes
234234
void start();
235235
#ifdef HAVE_LMY_ENGINE
236-
void engine_start(cqmc::engine::LMYEngine<ValueType>& EngineObj,
237-
DescentEngine& descentEngineObj,
238-
std::string MinMethod);
236+
void engine_start();
239237
#endif
240238
///common operation to finish optimization, used by the derived classes
241239
void finish();

src/QMCDrivers/WFOpt/QMCFixedSampleLinearOptimizeBatched.cpp

Lines changed: 9 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -177,17 +177,15 @@ void QMCFixedSampleLinearOptimizeBatched::start()
177177
}
178178

179179
#ifdef HAVE_LMY_ENGINE
180-
void QMCFixedSampleLinearOptimizeBatched::engine_start(cqmc::engine::LMYEngine<ValueType>& EngineObj,
181-
DescentEngine& descentEngineObj,
182-
std::string MinMethod)
180+
void QMCFixedSampleLinearOptimizeBatched::engine_start()
183181
{
184182
app_log() << "entering engine_start function" << std::endl;
185183

186184
std::unique_ptr<EngineHandle> handle;
187185
if (MinMethod == "descent")
188-
handle = std::make_unique<DescentEngineHandle>(descentEngineObj);
186+
handle = std::make_unique<DescentEngineHandle>(*descentEngineObj);
189187
else if (MinMethod == "adaptive")
190-
handle = std::make_unique<LMYEngineHandle>(EngineObj);
188+
handle = std::make_unique<LMYEngineHandle>(*EngineObj);
191189
else
192190
handle = std::make_unique<NullEngineHandle>();
193191

@@ -1176,7 +1174,7 @@ bool QMCFixedSampleLinearOptimizeBatched::adaptive_three_shift_run()
11761174
EngineObj->reset();
11771175

11781176
// generate samples and compute weights, local energies, and derivative vectors
1179-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1177+
engine_start();
11801178

11811179
int new_num = 0;
11821180

@@ -1320,7 +1318,7 @@ bool QMCFixedSampleLinearOptimizeBatched::adaptive_three_shift_run()
13201318
finish();
13211319

13221320
// take sample
1323-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1321+
engine_start();
13241322
}
13251323
else
13261324
{
@@ -1331,12 +1329,12 @@ bool QMCFixedSampleLinearOptimizeBatched::adaptive_three_shift_run()
13311329

13321330
if (options_LMY_.filter_param)
13331331
{
1334-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1332+
engine_start();
13351333
EngineObj->buildMatricesFromDerivatives();
13361334
}
13371335
else
13381336
{
1339-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1337+
engine_start();
13401338
app_log() << "Should be building matrices from stored samples" << std::endl;
13411339
EngineObj->buildMatricesFromDerivatives();
13421340
}
@@ -1942,7 +1940,7 @@ bool QMCFixedSampleLinearOptimizeBatched::stochastic_reconfiguration_conjugate_g
19421940
bool QMCFixedSampleLinearOptimizeBatched::descent_run()
19431941
{
19441942
//Compute Lagrangian derivatives needed for parameter updates with engine_checkConfigurations, which is called inside engine_start
1945-
engine_start(*EngineObj, *descentEngineObj, MinMethod);
1943+
engine_start();
19461944

19471945
int descent_num = descentEngineObj->getDescentNum();
19481946

@@ -1984,7 +1982,7 @@ bool QMCFixedSampleLinearOptimizeBatched::descent_run()
19841982
#ifdef HAVE_LMY_ENGINE
19851983
bool QMCFixedSampleLinearOptimizeBatched::hybrid_run()
19861984
{
1987-
app_log() << "This is methodName: " << MinMethod << std::endl;
1985+
app_log() << "This method name is: " << MinMethod << std::endl;
19881986

19891987
//Either the adaptive BLM or descent optimization is run
19901988

src/QMCDrivers/WFOpt/QMCFixedSampleLinearOptimizeBatched.h

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -81,9 +81,7 @@ class QMCFixedSampleLinearOptimizeBatched : public QMCDriverNew, LinearMethod
8181

8282
#ifdef HAVE_LMY_ENGINE
8383
using ValueType = QMCTraits::ValueType;
84-
void engine_start(cqmc::engine::LMYEngine<ValueType>& EngineObj,
85-
DescentEngine& descentEngineObj,
86-
std::string MinMethod);
84+
void engine_start();
8785
#endif
8886

8987

src/formic/utils/lmyengine/engine.cpp

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,8 @@ cqmc::engine::LMYEngine<S>::LMYEngine(const formic::VarDeps* dep_ptr,
106106
_lm_ham_shift_i(lm_ham_shift_i),
107107
_lm_ham_shift_s(lm_ham_shift_s),
108108
_lm_max_update_abs(lm_max_update_abs),
109+
_vg(1),
110+
_weight(1),
109111
_shift_scale(shift_scale),
110112
_dep_ptr(dep_ptr),
111113
_mbuilder(_der_rat,

0 commit comments

Comments
 (0)