import torch
import torch.nn as nn

print("="*60)
print("一、常见损失函数")
print("="*60)

# --- 回归: MSE / MAE / SmoothL1 ---
pred = torch.tensor([2.5, 0.0, 2.0])
target = torch.tensor([3.0, 0.0, 1.0])
print("MSE   :", nn.MSELoss()(pred, target).item(), " = mean((pred-target)^2)")
print("  手算:", ((pred-target)**2).mean().item())
print("MAE   :", nn.L1Loss()(pred, target).item(),  " = mean(|pred-target|)")
print("SmoothL1:", nn.SmoothL1Loss()(pred, target).item(), " 小误差像MSE,大误差像MAE(抗离群点)")

# --- 分类: CrossEntropy ---
print("\n--- 多分类 CrossEntropyLoss(含softmax) ---")
logits = torch.tensor([[2.0, 0.5, 0.1]])   # 3类的原始分数
label  = torch.tensor([0])                  # 正确类是第0类
ce = nn.CrossEntropyLoss()(logits, label)
# 手算: softmax -> -log(正确类概率)
sm = torch.softmax(logits, dim=1)
print("softmax概率:", sm.numpy().round(3))
print("CrossEntropy:", ce.item(), " 手算 -log(p0):", (-torch.log(sm[0,0])).item())

print("\n" + "="*60)
print("二、优化器对比：同一个问题，看谁收敛快")
print("="*60)

def train_with(opt_name, lr):
    torch.manual_seed(0)
    # 简单线性回归 y = 3x + 2
    x = torch.linspace(-1,1,50).unsqueeze(1)
    y = 3*x + 2
    model = nn.Linear(1,1)
    crit = nn.MSELoss()
    if opt_name == "SGD":
        opt = torch.optim.SGD(model.parameters(), lr=lr)
    elif opt_name == "Momentum":
        opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9)
    elif opt_name == "Adam":
        opt = torch.optim.Adam(model.parameters(), lr=lr)
    for _ in range(100):
        opt.zero_grad()
        loss = crit(model(x), y)
        loss.backward()
        opt.step()
    return loss.item()

for name, lr in [("SGD",0.1),("Momentum",0.1),("Adam",0.1)]:
    print(f"{name:10s}(lr={lr}) 100步后loss = {train_with(name,lr):.6f}")

print("\n" + "="*60)
print("三、学习率的影响（SGD）")
print("="*60)
for lr in [0.001, 0.05, 0.5, 1.5]:
    final = train_with_lr = None
    torch.manual_seed(0)
    x = torch.linspace(-1,1,50).unsqueeze(1); y = 3*x+2
    m = nn.Linear(1,1); c = nn.MSELoss()
    o = torch.optim.SGD(m.parameters(), lr=lr)
    for _ in range(100):
        o.zero_grad(); l = c(m(x),y); l.backward(); o.step()
    tag = "太小(学得慢)" if lr<=0.001 else ("合适" if lr<=0.5 else "太大(可能发散)")
    print(f"lr={lr:5} -> loss={l.item():.4f}  [{tag}]")
