import torch def net(w1, b1, w2, b2, x=torch.tensor(1.0), y=torch.tensor(2.0)): h1 = w1 * x + b1 a1 = torch.relu(h1) h2 = w2 * a1 + b2 return (h2 - y) ** 2, h1 p =[torch.tensor(v, requires_grad=True) for v in (2.0, 0.0, 3.0, 1.0)] opt = torch.optim.SGD(p, lr=0.1) loss, h1 = net(*p) opt.zero_grad() loss.backward() opt.step() print([round(t.item(), 4) for t in p]) # [-1.0, -3.0, 1.0, 0.0] print(net(*p)[0].item(), net(*p)[1].item()) # 4.0 -4.0