import torch

torch.manual_seed(0)

# 一个全连接层 y = W x + b
W = torch.randn(3, 4, requires_grad=True)   # 输出3维, 输入4维
b = torch.randn(3, requires_grad=True)
x = torch.randn(4, requires_grad=True)

y = W @ x + b                # 前向: y = Wx + b
g = torch.tensor([1.0, 2.0, 3.0])   # 假设上游传来的梯度向量 v (dL/dy)

# ===== autograd 自动反向 =====
y.backward(g)   # 传入上游梯度向量 -> 触发 VJP

print("=== autograd 算出来的梯度 ===")
print("dL/dx (autograd):", x.grad)
print("dL/dW (autograd):\n", W.grad)
print("dL/db (autograd):", b.grad)

# ===== 手推雅可比公式 =====
print("\n=== 手推雅可比公式验证 ===")
# dL/dx = W^T @ g   (VJP: 向量 × 雅可比W)
print("dL/dx 手推 = W^T @ g:", W.t() @ g)
# dL/dW = g (外积) x^T
print("dL/dW 手推 = g ⊗ x:\n", torch.outer(g, x))
# dL/db = g  (因为 dy/db = I 单位阵, VJP 就是 g 本身)
print("dL/db 手推 = g:", g)

# ===== 对比 =====
print("\n=== 是否完全一致 ===")
print("dx 一致:", torch.allclose(x.grad, W.t() @ g))
print("dW 一致:", torch.allclose(W.grad, torch.outer(g, x)))
print("db 一致:", torch.allclose(b.grad, g))
