练习题代码:

import torch
import numpy as np
import matplotlib.pyplot as plt
import math

t_c = [0.5,14.0,15.0,28.0,11.0,8.0,3.0,-4.0,6.0,13.0,21.0]
t_u = [35.7,55.9,58.2,81.9,56.3,48.9,33.9,21.8,48.4,60.4,68.4]
t_c = torch.tensor(t_c, dtype=torch.float32)    #仅添加声明数据类型
t_u = torch.tensor(t_u, dtype=torch.float32)

def model(t_u, w2,w1, b):
    return w2*t_u**2+w1*t_u + b

def loss_fn(t_p, t_c):
    return ((t_p - t_c)**2).mean()

def training_loop(n_epochs,optimizer,learning_rate, params, t_u, t_c):
    for n in range(1,n_epochs+1):
        if params.grad is not None:
            params.grad.zero_()
        
        t_p =model(t_u,*params)
        loss = loss_fn(t_p, t_c)
        loss.backward()
        optimizer.step()
        with torch.no_grad():
            params -= learning_rate * params.grad
        if n % 10000 == 0 or n == 1:
            print(f"Epoch {n}, Loss {loss.item():.4f}, w2[{params[0].item():.4f}],w1[{params[1].item():.4f}], b[{params[2].item():.4f}]")
    return params


if __name__ == "__main__":
    #训练
    params = torch.tensor([0.000000001,1, 0.0],requires_grad=True,dtype=torch.float32)
    learning_rate = 1e-8   #学习率
    opertimizer = torch.optim.SGD([params], lr=learning_rate)
    n_epochs = 1_000_000  #轮数
    trained_params = training_loop(n_epochs, opertimizer,learning_rate, params, t_u, t_c)
    print(f"Trained parameters: w2{trained_params[0].item():.4f},w1{trained_params[1].item():.4f}, b {trained_params[2].item():.4f}")

    #画图
    t_u_np = t_u.numpy()
    t_c_np = t_c.numpy()
    w2 = trained_params[0].item()
    w1 = trained_params[1].item()
    b = trained_params[2].item()

    u_line = np.linspace(t_u_np.min(), t_u_np.max(), 100)
    t_p_line = w2 * u_line**2 + w1 * u_line + b
    plt.figure(figsize=(8, 6))
    plt.scatter(t_u_np, t_c_np, color='red', label='Data (t_u vs t_c)')
    plt.plot(u_line, t_p_line, color='blue',
             label=f'Predicted: t = {w2:.4f} * u**2 + {w1:.4f} * u + {b:.4f}')
    plt.xlabel('t_u')
    plt.ylabel('Temperature (°C)')
    plt.title('Data and Learned Quadratic Model')
    plt.grid(True)
    plt.legend()
    plt.show()

效果图:

更多推荐