"""第 16 课补全练习：执行一次训练更新。修改 solve，保持下方检查不变。"""
import math, io, argparse
import torch
from torch import nn
from torch.nn import functional as F
torch.set_num_threads(2)
torch.manual_seed(26)

def solve(model, optimizer, x, y):
    optimizer.zero_grad(set_to_none=True)
    prediction=model(x)
    loss=F.mse_loss(prediction,y)
    loss.backward()
    optimizer.step()
    return loss.item()

model=nn.Linear(1,1);optimizer=torch.optim.SGD(model.parameters(),lr=.1)
x=torch.ones(4,1);y=torch.ones(4,1)*3
before=F.mse_loss(model(x),y).item()
solve(model,optimizer,x,y)
assert F.mse_loss(model(x),y).item()<before
print("本课补全练习通过。")
