Skip to content

Commit 2c6ce39

Browse files
authored
Validate explicit diago_proc bounds (#7739)
1 parent 3a7ce9f commit 2c6ce39

6 files changed

Lines changed: 58 additions & 3 deletions

File tree

source/source_io/module_parameter/read_input_item_system.cpp

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -594,11 +594,21 @@ Available options are:
594594
item.availability = "Used only for plane wave basis set.";
595595
read_sync_int(input.diago_proc);
596596
item.reset_value = [](const Input_Item& item, Parameter& para) {
597-
if (para.input.diago_proc > GlobalV::NPROC || para.input.diago_proc <= 0)
597+
if (para.input.diago_proc == 0)
598598
{
599599
para.input.diago_proc = GlobalV::NPROC;
600600
}
601601
};
602+
item.check_value = [](const Input_Item& item, const Parameter& para) {
603+
if (para.input.diago_proc < 0)
604+
{
605+
ModuleBase::WARNING_QUIT("ReadInput", "diago_proc must not be negative");
606+
}
607+
if (para.input.diago_proc > GlobalV::NPROC)
608+
{
609+
ModuleBase::WARNING_QUIT("ReadInput", "diago_proc cannot exceed the number of MPI processes");
610+
}
611+
};
602612
this->add_item(item);
603613
}
604614
{

source/source_io/test/read_input_ptest.cpp

Lines changed: 24 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,7 @@ TEST_F(InputParaTest, ParaRead)
147147
EXPECT_EQ(param.inp.ndx, 0);
148148
EXPECT_EQ(param.inp.ndy, 0);
149149
EXPECT_EQ(param.inp.ndz, 0);
150-
EXPECT_EQ(param.inp.diago_proc, std::min(GlobalV::NPROC, 4));
150+
EXPECT_EQ(param.inp.diago_proc, GlobalV::NPROC);
151151
EXPECT_EQ(param.inp.pw_diag_nmax, 50);
152152
EXPECT_EQ(param.inp.diago_cg_prec, 1);
153153
EXPECT_EQ(param.inp.pw_diag_ndim, 4);
@@ -456,6 +456,29 @@ TEST_F(InputParaTest, ParaRead)
456456
EXPECT_DOUBLE_EQ(param.inp.rdmft_power_alpha, 0.656);
457457
}
458458

459+
TEST_F(InputParaTest, DiagoProc)
460+
{
461+
int rank = 0;
462+
int nproc = 0;
463+
MPI_Comm_rank(MPI_COMM_WORLD, &rank);
464+
MPI_Comm_size(MPI_COMM_WORLD, &nproc);
465+
466+
ModuleIO::ReadInput readinput(rank);
467+
readinput.check_ntype_flag = false;
468+
Parameter full_param;
469+
readinput.read_parameters(full_param, "./support/INPUT.diago_proc_full");
470+
EXPECT_EQ(full_param.inp.diago_proc, nproc);
471+
472+
if (nproc == 4)
473+
{
474+
ModuleIO::ReadInput subset_readinput(rank);
475+
subset_readinput.check_ntype_flag = false;
476+
Parameter subset_param;
477+
subset_readinput.read_parameters(subset_param, "./support/INPUT.diago_proc_subset");
478+
EXPECT_EQ(subset_param.inp.diago_proc, 2);
479+
}
480+
}
481+
459482
// comment out this part of tests, since Parameter is in another directory now, mohan 2025-05-18
460483
// besides, the following tests will cause strange error in MPI_Finalize()
461484
// I tried the following modification, it worked well in my own environment, but not in the Github test, Xinyuan 2025-05-25

source/source_io/test/support/INPUT

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ out_freq_elec 0 #the frequency ( >= 0) of electronic iter to ou
3333
dft_plus_dmft 0 #true:DFT+DMFT; false: standard DFT calcullation(default)
3434
rpa 0 #true:generate output files used in rpa calculation; false:(default)
3535
mem_saver 0 #Only for nscf calculations. if set to 1, then a memory saving technique will be used for many k point calculations.
36-
diago_proc 4 #the number of procs used to do diagonalization
36+
diago_proc 0 #the number of procs used to do diagonalization
3737
nbspline -1 #the order of B-spline basis
3838
soc_lambda 1 #The fraction of averaged SOC pseudopotential is given by (1-soc_lambda)
3939
cal_force 0 #if calculate the force at the end of the electronic iteration
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
INPUT_PARAMETERS
2+
diago_proc 0
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
INPUT_PARAMETERS
2+
diago_proc 2

source/source_io/test_serial/read_input_test.cpp

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -336,6 +336,24 @@ TEST_F(InputTest, ValidateDeepksOutputFrequency)
336336
EXPECT_EQ(enabled_param.inp.deepks_out_freq_elec, 2);
337337
}
338338

339+
TEST_F(InputTest, ValidateDiagoProc)
340+
{
341+
set_nproc(4);
342+
343+
Parameter full_param;
344+
EXPECT_NO_THROW(read_parameters("diago_proc_full_INPUT", "diago_proc 0\n", full_param));
345+
EXPECT_EQ(full_param.inp.diago_proc, 4);
346+
347+
Parameter subset_param;
348+
EXPECT_NO_THROW(read_parameters("diago_proc_subset_INPUT", "diago_proc 2\n", subset_param));
349+
EXPECT_EQ(subset_param.inp.diago_proc, 2);
350+
351+
expect_invalid_input("diago_proc_negative_INPUT", "diago_proc -1\n", "diago_proc must not be negative");
352+
expect_invalid_input("diago_proc_oversized_INPUT",
353+
"diago_proc 5\n",
354+
"diago_proc cannot exceed the number of MPI processes");
355+
}
356+
339357
TEST_F(InputTest, Check)
340358
{
341359
ModuleIO::ReadInput readinput(0);

0 commit comments

Comments
 (0)