"""第 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):
    raise NotImplementedError("请补全：执行一次训练更新；参考答案见 solution.py")

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("本课补全练习通过。")
