Skip to content

Commit 56cbe2d

Browse files
author
Han Wang
committed
simplify three autograd to one by vmap, which was made inpossible by jit
1 parent a920ef6 commit 56cbe2d

1 file changed

Lines changed: 22 additions & 36 deletions

File tree

deepmd/pt_expt/model/transform_output.py

Lines changed: 22 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -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

6450
def task_deriv_one(

0 commit comments

Comments
 (0)