(sample)
| 114 | return torch.dot(sample, weight) |
| 115 | |
| 116 | def grad_sample(sample): |
| 117 | return torch.autograd.functional.vjp(lambda weight: model(sample), weight)[1] |
| 118 | |
| 119 | # The following doesn't actually work in the vmap prototype. But it |
| 120 | # could be an API for computing per-sample-gradients. |