MCPcopy Create free account
hub / github.com/pytorch/tutorials / grad_sample

Function grad_sample

unstable_source/vmap_recipe.py:116–117  ·  view source on GitHub ↗
(sample)

Source from the content-addressed store, hash-verified

114 return torch.dot(sample, weight)
115
116def 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.

Callers

nothing calls this directly

Calls 1

modelFunction · 0.70

Tested by

no test coverage detected