summaryrefslogtreecommitdiff
path: root/src/zbp_scaling/zbp/wp.py
blob: ff2511d7878c9b1ba8124fcae7e933eeb5c37d8f (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
"""Weight perturbation (parameter-space zeroth order) baseline: MeZO / CD-RGE / weight-space forward
gradient.  All parameters are perturbed by eps*u (Rademacher), the loss is measured at +/- and the
gradient estimate is (1/n) sum_i u_i (L(theta+eps u_i) - L(theta-eps u_i)) / (2 eps).
2n forward passes of the whole network per step, no backward pass, variance ~ P/n with P = #params."""
import torch


def wp_step(model, loss_fn, x, y, n, eps, gen):
    params = [p for p in model.parameters() if p.requires_grad]
    grads = [torch.zeros_like(p) for p in params]
    with torch.no_grad():
        for i in range(n):
            us = [(torch.randint(0, 2, p.shape, generator=gen, device=p.device).to(p.dtype) * 2 - 1) for p in params]
            for p, u in zip(params, us):
                p.add_(eps * u)
            lp = loss_fn(model(x), y).item()
            for p, u in zip(params, us):
                p.sub_(2 * eps * u)
            lm = loss_fn(model(x), y).item()
            for p, u in zip(params, us):
                p.add_(eps * u)
            D = (lp - lm) / (2 * eps)
            for g, u in zip(grads, us):
                g.add_(u, alpha=D / n)
        for p, g in zip(params, grads):
            p.grad = g
        loss = loss_fn(model(x), y)
    return loss