import torch
import torch.nn as nn

# ===== 0. 造数据：拟合 y = sin(x) 这个非线性函数 =====
torch.manual_seed(0)
x = torch.linspace(-3.14, 3.14, 200).unsqueeze(1)   # [200, 1]
y = torch.sin(x)                                     # [200, 1]

# ===== 1. 定义模型（万能骨架）=====
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(1, 32)    # 输入1维 -> 隐藏32维
        self.fc2 = nn.Linear(32, 32)   # 隐藏 -> 隐藏
        self.fc3 = nn.Linear(32, 1)    # 隐藏 -> 输出1维
    def forward(self, x):
        x = torch.relu(self.fc1(x))    # 非线性激活
        x = torch.relu(self.fc2(x))
        return self.fc3(x)             # 输出层不加激活(回归)

model = MLP()

# ===== 2. 损失函数 + 优化器 =====
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# ===== 3. 训练循环（五步雷打不动）=====
for epoch in range(1000):
    optimizer.zero_grad()       # 1. 清空梯度
    pred = model(x)             # 2. 前向(自动调forward)
    loss = criterion(pred, y)   # 3. 算loss
    loss.backward()             # 4. 反向(自动算梯度)
    optimizer.step()            # 5. 更新参数
    if (epoch+1) % 200 == 0:
        print(f"Epoch {epoch+1:4d} | Loss: {loss.item():.6f}")

# ===== 4. 看看学得怎么样 =====
with torch.no_grad():
    test_x = torch.tensor([[0.0], [1.5708], [-1.5708]])  # 0, π/2, -π/2
    test_pred = model(test_x)
    print("\n预测 sin(0)   =", round(test_pred[0].item(), 4), " 真实 =", 0.0)
    print("预测 sin(π/2) =", round(test_pred[1].item(), 4), " 真实 =", 1.0)
    print("预测 sin(-π/2)=", round(test_pred[2].item(), 4), " 真实 =", -1.0)

# ===== 5. 看看模型里的参数(就是W和b)=====
print("\n=== 模型参数形状 ===")
for name, p in model.named_parameters():
    print(f"{name}: {tuple(p.shape)}")
