import torch
import torch.nn as nn
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np

torch.manual_seed(0)

# ===== 数据 =====
x = torch.linspace(-3.14, 3.14, 200).unsqueeze(1)
y = torch.sin(x)

# ===== 模型 =====
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(1, 32)
        self.fc2 = nn.Linear(32, 32)
        self.fc3 = nn.Linear(32, 1)
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        return self.fc3(x)

model = MLP()
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# ===== 训练 + 记录中间过程 =====
loss_history = []
snapshots = {}     # 记录不同epoch的预测曲线
snap_epochs = [0, 20, 100, 1000]

for epoch in range(1001):
    optimizer.zero_grad()
    pred = model(x)
    loss = criterion(pred, y)
    loss.backward()
    optimizer.step()
    loss_history.append(loss.item())
    if epoch in snap_epochs:
        with torch.no_grad():
            snapshots[epoch] = model(x).squeeze().numpy()

# ===== 画图 =====
fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# 左图: 拟合过程（不同epoch的曲线）
ax = axes[0]
xn = x.squeeze().numpy()
ax.plot(xn, y.squeeze().numpy(), 'k-', lw=3, label='True sin(x)', alpha=0.6)
colors = ['#ffcccc', '#ff8888', '#ff3333', '#cc0000']
for c, ep in zip(colors, snap_epochs):
    ax.plot(xn, snapshots[ep], color=c, lw=1.8, label=f'epoch {ep}')
ax.set_title('MLP fitting sin(x): from random to perfect', fontsize=13)
ax.set_xlabel('x'); ax.set_ylabel('y')
ax.legend(); ax.grid(alpha=0.3)

# 右图: loss下降曲线（对数）
ax = axes[1]
ax.semilogy(loss_history, 'b-', lw=1.5)
ax.set_title('Training Loss (log scale)', fontsize=13)
ax.set_xlabel('epoch'); ax.set_ylabel('MSE Loss (log)')
ax.grid(alpha=0.3)
for ep in snap_epochs:
    ax.axvline(ep, color='r', ls='--', alpha=0.3)
    ax.text(ep, loss_history[ep], f'ep{ep}', fontsize=8, color='r')

plt.tight_layout()
plt.savefig('sin_fit_viz.png', dpi=130, bbox_inches='tight')
print("saved sin_fit_viz.png")
print(f"final loss = {loss_history[-1]:.6e}")
print(f"epoch0 loss = {loss_history[0]:.4f}  ->  epoch1000 loss = {loss_history[-1]:.6e}")