import torch
import torch.nn as nn

# ===== 三大激活函数的雅可比都是对角矩阵（逐元素函数）=====
x = torch.tensor([-2.0, -0.5, 0.0, 0.5, 2.0], requires_grad=True)

def show_grad(name, y):
    g = torch.ones_like(y)
    x.grad = None
    y.backward(g, retain_graph=True)
    print(f"{name:8s} | y={y.detach().numpy().round(3)} | dy/dx={x.grad.numpy().round(3)}")

print("输入 x =", x.detach().numpy())
print("-"*70)
show_grad("ReLU",    torch.relu(x))       # 导数: x>0->1, x<=0->0
show_grad("Sigmoid", torch.sigmoid(x))    # 导数: σ(1-σ), 最大0.25
show_grad("Tanh",    torch.tanh(x))       # 导数: 1-tanh^2, 最大1.0

print("-"*70)
print("观察:")
print("  ReLU 导数是 0/1 开关 -> 正区间梯度恒为1, 不衰减")
print("  Sigmoid 导数最大才0.25 -> 多层连乘 0.25^N 迅速趋0 -> 梯度消失")
print("  Tanh 导数最大1.0(在0处), 比sigmoid好但两端仍饱和")

# ===== 演示梯度消失: 10层sigmoid vs 10层relu =====
print("\n=== 10层堆叠后, 输入端梯度对比 ===")
for act_name, act in [("Sigmoid", torch.sigmoid), ("ReLU", torch.relu)]:
    z = torch.tensor([1.0], requires_grad=True)
    h = z
    for _ in range(10):
        h = act(h * 1.0)   # 简化: 每层过一次激活
    z.grad = None
    h.backward()
    print(f"  {act_name:8s}: 10层后输入端梯度 = {z.grad.item():.6e}")
