@@ -18,47 +18,33 @@ def atomic_virial_corr(
1818 atom_energy : torch .Tensor ,
1919) -> torch .Tensor :
2020 nall = extended_coord .shape [1 ]
21+ nf = extended_coord .shape [0 ]
2122 nloc = atom_energy .shape [1 ]
2223 coord , _ = torch .split (extended_coord , [nloc , nall - nloc ], dim = 1 )
2324 # no derivative with respect to the loc coord.
2425 coord = coord .detach ()
2526 ce = coord * atom_energy
26- sumce0 , sumce1 , sumce2 = torch .split (torch .sum (ce , dim = 1 ), [1 , 1 , 1 ], dim = - 1 )
27- faked_grad = torch .ones_like (sumce0 )
28- lst : list [torch .Tensor | None ] = [faked_grad ]
29- extended_virial_corr0 = torch .autograd .grad (
30- [sumce0 ],
31- [extended_coord ],
32- grad_outputs = lst ,
33- create_graph = False ,
34- retain_graph = True ,
35- )[0 ]
36- assert extended_virial_corr0 is not None
37- extended_virial_corr1 = torch .autograd .grad (
38- [sumce1 ],
39- [extended_coord ],
40- grad_outputs = lst ,
41- create_graph = False ,
42- retain_graph = True ,
43- )[0 ]
44- assert extended_virial_corr1 is not None
45- extended_virial_corr2 = torch .autograd .grad (
46- [sumce2 ],
47- [extended_coord ],
48- grad_outputs = lst ,
49- create_graph = False ,
50- retain_graph = True ,
51- )[0 ]
52- assert extended_virial_corr2 is not None
53- extended_virial_corr = torch .concat (
54- [
55- extended_virial_corr0 .unsqueeze (- 1 ),
56- extended_virial_corr1 .unsqueeze (- 1 ),
57- extended_virial_corr2 .unsqueeze (- 1 ),
58- ],
59- dim = - 1 ,
60- )
61- return extended_virial_corr
27+ sumce = torch .sum (ce , dim = 1 ) # [nf, 3]
28+
29+ # Use vmap to batch the 3 backward passes (one per spatial component)
30+ basis = torch .eye (3 , dtype = sumce .dtype , device = sumce .device ) # [3, 3]
31+ basis = basis .unsqueeze (1 ).expand (3 , nf , 3 ) # [3, nf, 3]
32+
33+ def grad_fn (grad_output : torch .Tensor ) -> torch .Tensor :
34+ result = torch .autograd .grad (
35+ [sumce ],
36+ [extended_coord ],
37+ grad_outputs = [grad_output ],
38+ create_graph = False ,
39+ retain_graph = True ,
40+ )[0 ]
41+ assert result is not None
42+ return result
43+
44+ # [3, nf, nall, 3] — batched over the 3 spatial components
45+ extended_virial_corr = torch .vmap (grad_fn )(basis )
46+ # [3, nf, nall, 3] -> [nf, nall, 3, 3]
47+ return extended_virial_corr .permute (1 , 2 , 3 , 0 )
6248
6349
6450def task_deriv_one (
0 commit comments