# tensorplay.func whirlwind tour Source: https://www.tensorplay.cn/docs/func.whirlwind_tour.html A hands-on pass over the transforms, in the order you would usually meet them. Every snippet is self-contained and can be run as-is. ## grad() — gradient computation grad(func) returns a new function that computes the gradient of func. It assumes func returns a single-element tensor and by default differentiates with respect to the first argument. ``` import tensorplay as tp from tensorplay.func import grad x = tp.randn([]) cos_x = grad(lambda x: tp.sin(x))(x) assert tp.allclose(cos_x, x.cos()) # second-order gradients are just grad of grad neg_sin_x = grad(grad(lambda x: tp.sin(x)))(x) assert tp.allclose(neg_sin_x, -x.sin()) ``` ## vmap() — auto-vectorization vmap(func) returns a new function that maps func over a dimension (default: 0) of each tensor input. Write the function for a single sample; vmap handles the batch: ``` import tensorplay as tp from tensorplay.func import vmap batch_size, feature_size = 3, 5 weights = tp.randn(feature_size, requires_grad=True) def model(feature_vec): assert feature_vec.dim() == 1 return (feature_vec * weights).sum().relu() examples = tp.randn(batch_size, feature_size) result = vmap(model)(examples) assert result.shape == (batch_size,) ``` For a pure function, vmap(f)(x) is equivalent to stacking one call per sample: ``` import tensorplay as tp from tensorplay.func import vmap xs = tp.randn(4, 3) assert tp.allclose( vmap(lambda x: x * x)(xs), tp.stack([x * x for x in xs.unbind(0)]), ) ``` [vmap()](/docs/generated/tensorplay.func.vmap.html#tensorplay.func.vmap) imposes restrictions on the code it can map over — see [UX limitations](/docs/func.ux_limitations.html#ux-limitations). ### Per-sample gradients Composing vmap and grad computes per-sample gradients: each sample’s gradient is produced independently, with no averaging across the batch. ``` import tensorplay as tp from tensorplay.func import grad, vmap batch_size, feature_size = 3, 5 def compute_loss(weights, example, target): y = (example * weights).sum().relu() return (y - target) ** 2 weights = tp.randn(feature_size, requires_grad=True) examples = tp.randn(batch_size, feature_size) targets = tp.randn(batch_size) grad_per_example = vmap(grad(compute_loss), in_dims=(None, 0, 0))( weights, examples, targets ) assert grad_per_example.shape == (batch_size, feature_size) ``` The in_dims=(None, 0, 0) says the weights are shared across the batch while the examples and targets are mapped. ## vjp() — vector-Jacobian product vjp applies func to its inputs and returns a new function that computes the vector-Jacobian product for given cotangents: ``` import tensorplay as tp from tensorplay.func import vjp inputs = tp.randn(3) outputs, vjp_fn = vjp(tp.sin, inputs) vjps = vjp_fn(tp.randn(3)) ``` ## jvp() — Jacobian-vector product jvp computes forward-mode AD. Unlike most other transforms it is not a higher-order function: it returns the outputs of func(inputs) together with the Jacobian-vector products for the given tangents: ``` import tensorplay as tp from tensorplay.func import jvp x, y = tp.randn(5), tp.randn(5) _, out_tangent = jvp(lambda x, y: x * y, (x, y), (tp.ones(5), tp.ones(5))) assert tp.allclose(out_tangent, x + y) ``` ## jacrev(), jacfwd(), and hessian() jacrev returns a new function that takes in x and returns the Jacobian using reverse-mode AD; jacfwd is its forward-mode counterpart: ``` import tensorplay as tp from tensorplay.func import jacfwd, jacrev x = tp.randn(5) assert tp.allclose(jacrev(tp.sin)(x), tp.diag(tp.cos(x))) assert tp.allclose(jacfwd(tp.sin)(x), tp.diag(tp.cos(x))) ``` Composing the two directions produces Hessians, and [hessian()](/docs/generated/tensorplay.func.hessian.html#tensorplay.func.hessian) is the convenience wrapper: ``` import tensorplay as tp from tensorplay.func import hessian, jacfwd, jacrev def f(x): return x.sin().sum() x = tp.randn(5) h0 = jacrev(jacrev(f))(x) h1 = jacfwd(jacrev(f))(x) h2 = hessian(f)(x) assert tp.allclose(h0, h2) and tp.allclose(h1, h2) ``` ## linearize() linearize evaluates func once and returns the value together with a callable forward-mode linearization at that point: ``` import tensorplay as tp from tensorplay.func import linearize x = tp.randn(5) y, jvp_fn = linearize(tp.sin, x) assert tp.allclose(jvp_fn(tp.ones(5)), tp.cos(x)) ``` ## What composes today Not every composition of transforms is available yet. The verified compositions: - vmap(grad(f)) with a shared differentiable argument — the per-sample gradient pattern above. - vmap(vmap(f)) — mapping over two independent batch dimensions. - grad(vmap(f)) — differentiating a batched evaluation. vmap(jacrev(f)) and vmap(jacfwd(f)) — batched Jacobians — currently raise; see [UX limitations](/docs/func.ux_limitations.html#ux-limitations) for the details.