Skip to content

Commit 0c8b6dc

Browse files
dyzhengdyzheng
andauthored
Feature: add Hessian operator <\phi|\nabla_x\nabla_y|\phi> (deepmodeling#6888)
* Feature: add Hessian operator <\phi|\nabla_x\nabla_y|\phi> * fix: UT of twocenterintegral --------- Co-authored-by: dyzheng <zhengdy@bjaisi.com>
1 parent 70e5f7d commit 0c8b6dc

8 files changed

Lines changed: 895 additions & 40 deletions

File tree

source/source_base/test/ylm_test.cpp

Lines changed: 436 additions & 1 deletion
Large diffs are not rendered by default.

source/source_base/ylm.cpp

Lines changed: 217 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1315,9 +1315,224 @@ void Ylm::hes_rl_sph_harm
13151315
if (Lmax == 4) return;
13161316

13171317
/***************************
1318-
L > 4
1318+
L = 5
1319+
***************************/
1320+
//m=0 : (63z^5 - 70z^3*r^2 + 15z*r^4)
1321+
coeff = sqrt(11.0 / ModuleBase::PI) / 16.0;
1322+
hrly[25][0] = (180*x*x*z + 60*y*y*z - 80*z*z*z) * coeff;
1323+
hrly[25][1] = (120*x*y*z) * coeff;
1324+
hrly[25][2] = (60*x*x*x + 60*x*y*y - 240*x*z*z) * coeff;
1325+
hrly[25][3] = (60*x*x*z + 180*y*y*z - 80*z*z*z) * coeff;
1326+
hrly[25][4] = (60*x*x*y + 60*y*y*y - 240*y*z*z) * coeff;
1327+
hrly[25][5] = (-240*x*x*z - 240*y*y*z + 160*z*z*z) * coeff;
1328+
1329+
//m=1 : x(21z^4 - 14z^2*r^2 + r^4)
1330+
coeff = sqrt(165.0 / 2.0 / ModuleBase::PI) / 16.0;
1331+
hrly[26][0] = (20*x*x*x + 12*x*y*y - 72*x*z*z) * coeff;
1332+
hrly[26][1] = (12*x*x*y + 4*y*y*y - 24*y*z*z) * coeff;
1333+
hrly[26][2] = (-72*x*x*z - 24*y*y*z + 32*z*z*z) * coeff;
1334+
hrly[26][3] = (4*x*x*x + 12*x*y*y - 24*x*z*z) * coeff;
1335+
hrly[26][4] = (-48*x*y*z) * coeff;
1336+
hrly[26][5] = (-24*x*x*x - 24*x*y*y + 96*x*z*z) * coeff;
1337+
1338+
//m=-1 : y(21z^4 - 14z^2*r^2 + r^4)
1339+
hrly[27][0] = (12*x*x*y + 4*y*y*y - 24*y*z*z) * coeff;
1340+
hrly[27][1] = (4*x*x*x + 12*x*y*y - 24*x*z*z) * coeff;
1341+
hrly[27][2] = (-48*x*y*z) * coeff;
1342+
hrly[27][3] = (12*x*x*y + 20*y*y*y - 72*y*z*z) * coeff;
1343+
hrly[27][4] = (-24*x*x*z - 72*y*y*z + 32*z*z*z) * coeff;
1344+
hrly[27][5] = (-24*x*x*y - 24*y*y*y + 96*y*z*z) * coeff;
1345+
1346+
//m=2 : (x^2 - y^2)(3z^3 - z*r^2)
1347+
coeff = sqrt(1155.0 / ModuleBase::PI) / 8.0;
1348+
hrly[28][0] = (-12*x*x*z + 4*z*z*z) * coeff;
1349+
hrly[28][1] = 0.0;
1350+
hrly[28][2] = (-4*x*x*x + 12*x*z*z) * coeff;
1351+
hrly[28][3] = (12*y*y*z - 4*z*z*z) * coeff;
1352+
hrly[28][4] = (4*y*y*y - 12*y*z*z) * coeff;
1353+
hrly[28][5] = (12*x*x*z - 12*y*y*z) * coeff;
1354+
1355+
//m=-2 : xy(3z^3 - z*r^2)
1356+
hrly[29][0] = (-6*x*y*z) * coeff;
1357+
hrly[29][1] = (-3*x*x*z - 3*y*y*z + 2*z*z*z) * coeff;
1358+
hrly[29][2] = (-3*x*x*y - y*y*y + 6*y*z*z) * coeff;
1359+
hrly[29][3] = (-6*x*y*z) * coeff;
1360+
hrly[29][4] = (-x*x*x - 3*x*y*y + 6*x*z*z) * coeff;
1361+
hrly[29][5] = (12*x*y*z) * coeff;
1362+
1363+
//m=3 : x(x^2 - 3y^2)(9z^2 - r^2)
1364+
coeff = sqrt(385.0 / 2.0 / ModuleBase::PI) / 16.0;
1365+
hrly[30][0] = (-20*x*x*x + 12*x*y*y + 48*x*z*z) * coeff;
1366+
hrly[30][1] = (12*x*x*y + 12*y*y*y - 48*y*z*z) * coeff;
1367+
hrly[30][2] = (48*x*x*z - 48*y*y*z) * coeff;
1368+
hrly[30][3] = (4*x*x*x + 36*x*y*y - 48*x*z*z) * coeff;
1369+
hrly[30][4] = (-96*x*y*z) * coeff;
1370+
hrly[30][5] = (16*x*x*x - 48*x*y*y) * coeff;
1371+
1372+
//m=-3 : y(3x^2 - y^2)(9z^2 - r^2)
1373+
hrly[31][0] = (-36*x*x*y - 4*y*y*y + 48*y*z*z) * coeff;
1374+
hrly[31][1] = (-12*x*x*x - 12*x*y*y + 48*x*z*z) * coeff;
1375+
hrly[31][2] = (96*x*y*z) * coeff;
1376+
hrly[31][3] = (-12*x*x*y + 20*y*y*y - 48*y*z*z) * coeff;
1377+
hrly[31][4] = (48*x*x*z - 48*y*y*z) * coeff;
1378+
hrly[31][5] = (48*x*x*y - 16*y*y*y) * coeff;
1379+
1380+
//m=4 : (x^4 - 6x^2*y^2 + y^4) * z
1381+
coeff = sqrt(385.0 / ModuleBase::PI) / 16.0;
1382+
hrly[32][0] = (12*x*x*z - 12*y*y*z) * coeff;
1383+
hrly[32][1] = (-24*x*y*z) * coeff;
1384+
hrly[32][2] = (4*x*x*x - 12*x*y*y) * coeff;
1385+
hrly[32][3] = (-12*x*x*z + 12*y*y*z) * coeff;
1386+
hrly[32][4] = (-12*x*x*y + 4*y*y*y) * coeff;
1387+
hrly[32][5] = 0.0;
1388+
1389+
//m=-4 : xy(x^2 - y^2) * z
1390+
hrly[33][0] = (6*x*y*z) * coeff;
1391+
hrly[33][1] = (3*x*x*z - 3*y*y*z) * coeff;
1392+
hrly[33][2] = (3*x*x*y - y*y*y) * coeff;
1393+
hrly[33][3] = (-6*x*y*z) * coeff;
1394+
hrly[33][4] = (x*x*x - 3*x*y*y) * coeff;
1395+
hrly[33][5] = 0.0;
1396+
1397+
//m=5 : x(x^4 - 10x^2*y^2 + 5y^4)
1398+
coeff = sqrt(77.0 / 2.0 / ModuleBase::PI) / 16.0;
1399+
hrly[34][0] = (20.0 * x*x*x - 60.0 * x * y*y) * coeff;
1400+
hrly[34][1] = (-60.0 * x*x * y + 20.0 * y*y*y) * coeff;
1401+
hrly[34][2] = 0.0;
1402+
hrly[34][3] = (-20.0 * x*x*x + 60.0 * x * y*y) * coeff;
1403+
hrly[34][4] = 0.0;
1404+
hrly[34][5] = 0.0;
1405+
1406+
//m=-5 : y(5x^4 - 10x^2*y^2 + y^4)
1407+
hrly[35][0] = (60.0 * x*x * y - 20.0 * y*y*y) * coeff;
1408+
hrly[35][1] = (20.0 * x*x*x - 60.0 * x * y*y) * coeff;
1409+
hrly[35][2] = 0.0;
1410+
hrly[35][3] = (-60.0 * x*x * y + 20.0 * y*y*y) * coeff;
1411+
hrly[35][4] = 0.0;
1412+
hrly[35][5] = 0.0;
1413+
1414+
if (Lmax == 5) return;
1415+
1416+
/***************************
1417+
L = 6
1418+
***************************/
1419+
//m=0 : (231z^6 - 315z^4*r^2 + 105z^2*r^4 - 5r^6)
1420+
coeff = sqrt(13.0 / ModuleBase::PI) / 32.0;
1421+
hrly[36][0] = (-150*x*x*x*x - 180*x*x*y*y + 1080*x*x*z*z - 30*y*y*y*y + 360*y*y*z*z - 240*z*z*z*z) * coeff;
1422+
hrly[36][1] = (-120*x*x*x*y - 120*x*y*y*y + 720*x*y*z*z) * coeff;
1423+
hrly[36][2] = (720*x*x*x*z + 720*x*y*y*z - 960*x*z*z*z) * coeff;
1424+
hrly[36][3] = (-30*x*x*x*x - 180*x*x*y*y + 360*x*x*z*z - 150*y*y*y*y + 1080*y*y*z*z - 240*z*z*z*z) * coeff;
1425+
hrly[36][4] = (720*x*x*y*z + 720*y*y*y*z - 960*y*z*z*z) * coeff;
1426+
hrly[36][5] = (180*x*x*x*x + 360*x*x*y*y - 1440*x*x*z*z + 180*y*y*y*y - 1440*y*y*z*z + 480*z*z*z*z) * coeff;
1427+
1428+
//m=1 : x(33z^5 - 30z^3*r^2 + 5z*r^4)
1429+
coeff = sqrt(273.0 / 2.0 / ModuleBase::PI) / 16.0;
1430+
hrly[37][0] = (100*x*x*x*z + 60*x*y*y*z - 120*x*z*z*z) * coeff;
1431+
hrly[37][1] = (60*x*x*y*z + 20*y*y*y*z - 40*y*z*z*z) * coeff;
1432+
hrly[37][2] = (25*x*x*x*x + 30*x*x*y*y - 180*x*x*z*z + 5*y*y*y*y - 60*y*y*z*z + 40*z*z*z*z) * coeff;
1433+
hrly[37][3] = (20*x*x*x*z + 60*x*y*y*z - 40*x*z*z*z) * coeff;
1434+
hrly[37][4] = (20*x*x*x*y + 20*x*y*y*y - 120*x*y*z*z) * coeff;
1435+
hrly[37][5] = (-120*x*x*x*z - 120*x*y*y*z + 160*x*z*z*z) * coeff;
1436+
1437+
//m=-1 : y(33z^5 - 30z^3*r^2 + 5z*r^4)
1438+
hrly[38][0] = (60*x*x*y*z + 20*y*y*y*z - 40*y*z*z*z) * coeff;
1439+
hrly[38][1] = (20*x*x*x*z + 60*x*y*y*z - 40*x*z*z*z) * coeff;
1440+
hrly[38][2] = (20*x*x*x*y + 20*x*y*y*y - 120*x*y*z*z) * coeff;
1441+
hrly[38][3] = (60*x*x*y*z + 100*y*y*y*z - 120*y*z*z*z) * coeff;
1442+
hrly[38][4] = (5*x*x*x*x + 30*x*x*y*y - 60*x*x*z*z + 25*y*y*y*y - 180*y*y*z*z + 40*z*z*z*z) * coeff;
1443+
hrly[38][5] = (-120*x*x*y*z - 120*y*y*y*z + 160*y*z*z*z) * coeff;
1444+
1445+
//m=2 : (x^2 - y^2)(33z^4 - 18z^2*r^2 + r^4)
1446+
coeff = sqrt(1365.0 / ModuleBase::PI) / 32.0;
1447+
hrly[39][0] = (30*x*x*x*x + 12*x*x*y*y - 192*x*x*z*z - 2*y*y*y*y + 32*z*z*z*z) * coeff;
1448+
hrly[39][1] = (8*x*x*x*y - 8*x*y*y*y) * coeff;
1449+
hrly[39][2] = (-128*x*x*x*z + 128*x*z*z*z) * coeff;
1450+
hrly[39][3] = (2*x*x*x*x - 12*x*x*y*y - 30*y*y*y*y + 192*y*y*z*z - 32*z*z*z*z) * coeff;
1451+
hrly[39][4] = (128*y*y*y*z - 128*y*z*z*z) * coeff;
1452+
hrly[39][5] = (-32*x*x*x*x + 192*x*x*z*z + 32*y*y*y*y - 192*y*y*z*z) * coeff;
1453+
1454+
//m=-2 : xy(33z^4 - 18z^2*r^2 + r^4)
1455+
hrly[40][0] = (20*x*x*x*y + 12*x*y*y*y - 96*x*y*z*z) * coeff;
1456+
hrly[40][1] = (20*x*x*x*x + 36*x*x*y*y - 96*x*x*z*z + 20*y*y*y*y - 96*y*y*z*z + 32*z*z*z*z) * coeff;
1457+
hrly[40][2] = (-96*x*x*y*z - 32*y*y*y*z + 64*y*z*z*z) * coeff;
1458+
hrly[40][3] = (12*x*x*x*y + 20*x*y*y*y - 96*x*y*z*z) * coeff;
1459+
hrly[40][4] = (-32*x*x*x*z - 96*x*y*y*z + 64*x*z*z*z) * coeff;
1460+
hrly[40][5] = (-32*x*x*x*y - 32*x*y*y*y + 192*x*y*z*z) * coeff;
1461+
1462+
//m=3 : x(x^2 - 3y^2)(11z^3 - 3z*r^2)
1463+
coeff = sqrt(1365.0 / ModuleBase::PI) / 16.0;
1464+
hrly[41][0] = (-60*x*x*x*z + 36*x*y*y*z + 48*x*z*z*z) * coeff;
1465+
hrly[41][1] = (36*x*x*y*z + 36*y*y*y*z - 48*y*z*z*z) * coeff;
1466+
hrly[41][2] = (-30*x*x*x*x + 36*x*x*y*y + 72*x*x*z*z + 18*y*y*y*y - 72*y*y*z*z) * coeff;
1467+
hrly[41][3] = (12*x*x*x*z + 108*x*y*y*z - 48*x*z*z*z) * coeff;
1468+
hrly[41][4] = (12*x*x*x*y + 36*x*y*y*y - 144*x*y*z*z) * coeff;
1469+
hrly[41][5] = (48*x*x*x*z - 144*x*y*y*z) * coeff;
1470+
1471+
//m=-3 : y(3x^2 - y^2)(11z^3 - 3z*r^2)
1472+
hrly[42][0] = (-108*x*x*y*z - 12*y*y*y*z + 48*y*z*z*z) * coeff;
1473+
hrly[42][1] = (-36*x*x*x*z - 36*x*y*y*z + 48*x*z*z*z) * coeff;
1474+
hrly[42][2] = (-36*x*x*x*y - 12*x*y*y*y + 144*x*y*z*z) * coeff;
1475+
hrly[42][3] = (-36*x*x*y*z + 60*y*y*y*z - 48*y*z*z*z) * coeff;
1476+
hrly[42][4] = (-18*x*x*x*x - 36*x*x*y*y + 72*x*x*z*z + 30*y*y*y*y - 72*y*y*z*z) * coeff;
1477+
hrly[42][5] = (144*x*x*y*z - 48*y*y*y*z) * coeff;
1478+
1479+
//m=4 : (x^4 - 6x^2*y^2 + y^4)(11z^2 - r^2)
1480+
coeff = sqrt(91.0 / ModuleBase::PI) / 32.0;
1481+
hrly[43][0] = (-30*x*x*x*x + 60*x*x*y*y + 120*x*x*z*z + 10*y*y*y*y - 120*y*y*z*z) * coeff;
1482+
hrly[43][1] = (40*x*x*x*y + 40*x*y*y*y - 240*x*y*z*z) * coeff;
1483+
hrly[43][2] = (80*x*x*x*z - 240*x*y*y*z) * coeff;
1484+
hrly[43][3] = (10*x*x*x*x + 60*x*x*y*y - 120*x*x*z*z - 30*y*y*y*y + 120*y*y*z*z) * coeff;
1485+
hrly[43][4] = (-240*x*x*y*z + 80*y*y*y*z) * coeff;
1486+
hrly[43][5] = (20*x*x*x*x - 120*x*x*y*y + 20*y*y*y*y) * coeff;
1487+
1488+
//m=-4 : xy(x^2 - y^2)(11z^2 - r^2)
1489+
hrly[44][0] = (-20*x*x*x*y + 60*x*y*z*z) * coeff;
1490+
hrly[44][1] = (-5*x*x*x*x + 30*x*x*z*z + 5*y*y*y*y - 30*y*y*z*z) * coeff;
1491+
hrly[44][2] = (60*x*x*y*z - 20*y*y*y*z) * coeff;
1492+
hrly[44][3] = (20*x*y*y*y - 60*x*y*z*z) * coeff;
1493+
hrly[44][4] = (20*x*x*x*z - 60*x*y*y*z) * coeff;
1494+
hrly[44][5] = (20*x*x*x*y - 20*x*y*y*y) * coeff;
1495+
1496+
//m=5 : x(x^4 - 10x^2*y^2 + 5y^4) * z
1497+
coeff = sqrt(1001.0 / 2.0 / ModuleBase::PI) / 16.0;
1498+
hrly[45][0] = (20*x*x*x*z - 60*x*y*y*z) * coeff;
1499+
hrly[45][1] = (-60*x*x*y*z + 20*y*y*y*z) * coeff;
1500+
hrly[45][2] = (5*x*x*x*x - 30*x*x*y*y + 5*y*y*y*y) * coeff;
1501+
hrly[45][3] = (-20*x*x*x*z + 60*x*y*y*z) * coeff;
1502+
hrly[45][4] = (-20*x*x*x*y + 20*x*y*y*y) * coeff;
1503+
hrly[45][5] = 0.0;
1504+
1505+
//m=-5 : y(5x^4 - 10x^2*y^2 + y^4) * z
1506+
hrly[46][0] = (60*x*x*y*z - 20*y*y*y*z) * coeff;
1507+
hrly[46][1] = (20*x*x*x*z - 60*x*y*y*z) * coeff;
1508+
hrly[46][2] = (20*x*x*x*y - 20*x*y*y*y) * coeff;
1509+
hrly[46][3] = (-60*x*x*y*z + 20*y*y*y*z) * coeff;
1510+
hrly[46][4] = (5*x*x*x*x - 30*x*x*y*y + 5*y*y*y*y) * coeff;
1511+
hrly[46][5] = 0.0;
1512+
1513+
//m=6 : (x^6 - 15x^4*y^2 + 15x^2*y^4 - y^6)
1514+
coeff = sqrt(3003.0 / ModuleBase::PI) / 32.0;
1515+
hrly[47][0] = (30*x*x*x*x - 180*x*x*y*y + 30*y*y*y*y) * coeff;
1516+
hrly[47][1] = (-120*x*x*x*y + 120*x*y*y*y) * coeff;
1517+
hrly[47][2] = 0.0;
1518+
hrly[47][3] = (-30*x*x*x*x + 180*x*x*y*y - 30*y*y*y*y) * coeff;
1519+
hrly[47][4] = 0.0;
1520+
hrly[47][5] = 0.0;
1521+
1522+
//m=-6 : xy(3x^4 - 10x^2*y^2 + 3y^4)
1523+
hrly[48][0] = (60*x*x*x*y - 60*x*y*y*y) * coeff;
1524+
hrly[48][1] = (15*x*x*x*x - 90*x*x*y*y + 15*y*y*y*y) * coeff;
1525+
hrly[48][2] = 0.0;
1526+
hrly[48][3] = (-60*x*x*x*y + 60*x*y*y*y) * coeff;
1527+
hrly[48][4] = 0.0;
1528+
hrly[48][5] = 0.0;
1529+
1530+
if (Lmax == 6) return;
1531+
1532+
/***************************
1533+
L > 6
13191534
***************************/
1320-
ModuleBase::WARNING_QUIT("hes_rl_sph_harm","l>4 not implemented!");
1535+
ModuleBase::WARNING_QUIT("hes_rl_sph_harm","l>6 not implemented!");
13211536

13221537

13231538
return;

source/source_basis/module_nao/test/two_center_integrator_test.cpp

Lines changed: 147 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -218,6 +218,153 @@ TEST_F(TwoCenterIntegratorTest, SphericalBessel)
218218
delete[] zeros;
219219
}
220220

221+
TEST_F(TwoCenterIntegratorTest, HessianSymmetry)
222+
{
223+
nfile = 3;
224+
orb.build(nfile, file, 'o');
225+
226+
ModuleBase::SphericalBesselTransformer sbt;
227+
orb.set_transformer(sbt);
228+
229+
double rmax = orb.rcut_max() * 2.0;
230+
double dr = 0.01;
231+
int nr = static_cast<int>(rmax / dr) + 1;
232+
233+
orb.set_uniform_grid(true, nr, rmax, 'i', true);
234+
235+
S_intor.tabulate(orb, orb, 'S', nr, rmax);
236+
T_intor.tabulate(orb, orb, 'T', nr, rmax);
237+
238+
ModuleBase::Vector3<double> R(1.5, 2.0, 1.0);
239+
double hess[9];
240+
241+
// Test S operator
242+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, nullptr, hess);
243+
244+
EXPECT_NEAR(hess[1], hess[3], 1e-10); // H_xy == H_yx
245+
EXPECT_NEAR(hess[2], hess[6], 1e-10); // H_xz == H_zx
246+
EXPECT_NEAR(hess[5], hess[7], 1e-10); // H_yz == H_zy
247+
248+
// Test T operator
249+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, nullptr, hess);
250+
251+
EXPECT_NEAR(hess[1], hess[3], 1e-10); // H_xy == H_yx
252+
EXPECT_NEAR(hess[2], hess[6], 1e-10); // H_xz == H_zx
253+
EXPECT_NEAR(hess[5], hess[7], 1e-10); // H_yz == H_zy
254+
}
255+
256+
TEST_F(TwoCenterIntegratorTest, HessianFiniteDifference)
257+
{
258+
nfile = 3;
259+
orb.build(nfile, file, 'o');
260+
261+
ModuleBase::SphericalBesselTransformer sbt;
262+
orb.set_transformer(sbt);
263+
264+
double rmax = orb.rcut_max() * 2.0;
265+
double dr = 0.01;
266+
int nr = static_cast<int>(rmax / dr) + 1;
267+
268+
orb.set_uniform_grid(true, nr, rmax, 'i', true);
269+
270+
S_intor.tabulate(orb, orb, 'S', nr, rmax);
271+
T_intor.tabulate(orb, orb, 'T', nr, rmax);
272+
273+
ModuleBase::Vector3<double> R(1.5, 2.0, 1.0);
274+
double hess_analytical[9];
275+
double hess_numerical[9];
276+
double eps = 1e-5;
277+
278+
// Test S operator
279+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, nullptr, hess_analytical);
280+
281+
// Compute numerical Hessian via finite differences
282+
for (int alpha = 0; alpha < 3; ++alpha)
283+
{
284+
for (int beta = 0; beta < 3; ++beta)
285+
{
286+
ModuleBase::Vector3<double> R_plus = R, R_minus = R;
287+
R_plus[beta] += eps;
288+
R_minus[beta] -= eps;
289+
290+
double grad_plus[3], grad_minus[3];
291+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R_plus, nullptr, grad_plus, nullptr);
292+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R_minus, nullptr, grad_minus, nullptr);
293+
294+
hess_numerical[alpha * 3 + beta] = (grad_plus[alpha] - grad_minus[alpha]) / (2.0 * eps);
295+
}
296+
}
297+
298+
// Compare with tolerance appropriate for finite differences
299+
for (int i = 0; i < 9; ++i)
300+
{
301+
EXPECT_NEAR(hess_analytical[i], hess_numerical[i], 1e-5);
302+
}
303+
304+
// Test T operator
305+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, nullptr, hess_analytical);
306+
307+
for (int alpha = 0; alpha < 3; ++alpha)
308+
{
309+
for (int beta = 0; beta < 3; ++beta)
310+
{
311+
ModuleBase::Vector3<double> R_plus = R, R_minus = R;
312+
R_plus[beta] += eps;
313+
R_minus[beta] -= eps;
314+
315+
double grad_plus[3], grad_minus[3];
316+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R_plus, nullptr, grad_plus, nullptr);
317+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R_minus, nullptr, grad_minus, nullptr);
318+
319+
hess_numerical[alpha * 3 + beta] = (grad_plus[alpha] - grad_minus[alpha]) / (2.0 * eps);
320+
}
321+
}
322+
323+
for (int i = 0; i < 9; ++i)
324+
{
325+
EXPECT_NEAR(hess_analytical[i], hess_numerical[i], 1e-5);
326+
}
327+
}
328+
329+
TEST_F(TwoCenterIntegratorTest, HessianDoesNotBreakGradient)
330+
{
331+
nfile = 3;
332+
orb.build(nfile, file, 'o');
333+
334+
ModuleBase::SphericalBesselTransformer sbt;
335+
orb.set_transformer(sbt);
336+
337+
double rmax = orb.rcut_max() * 2.0;
338+
double dr = 0.01;
339+
int nr = static_cast<int>(rmax / dr) + 1;
340+
341+
orb.set_uniform_grid(true, nr, rmax, 'i', true);
342+
343+
S_intor.tabulate(orb, orb, 'S', nr, rmax);
344+
T_intor.tabulate(orb, orb, 'T', nr, rmax);
345+
346+
ModuleBase::Vector3<double> R(1.5, 2.0, 1.0);
347+
double grad_only[3], grad_with_hess[3], hess[9];
348+
349+
// Test S operator
350+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, grad_only, nullptr);
351+
S_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, grad_with_hess, hess);
352+
353+
for (int i = 0; i < 3; ++i)
354+
{
355+
EXPECT_NEAR(grad_only[i], grad_with_hess[i], 1e-12);
356+
}
357+
358+
// Test T operator
359+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, grad_only, nullptr);
360+
T_intor.calculate(0, 1, 0, 0, 1, 1, 0, 0, R, nullptr, grad_with_hess, hess);
361+
362+
for (int i = 0; i < 3; ++i)
363+
{
364+
EXPECT_NEAR(grad_only[i], grad_with_hess[i], 1e-12);
365+
}
366+
}
367+
221368
int main(int argc, char** argv)
222369
{
223370

0 commit comments

Comments
 (0)