In [1]:
import torch
import torch.nn as nn
import torch.optim as optim
import pytorch_soom
import numpy as np
import matplotlib.pyplot as plt
from torch.utils.data import DataLoader, TensorDataset
from sklearn.datasets import load_diabetes
from sklearn.preprocessing import MinMaxScaler

In [2]:
device = 'cpu'

In [3]:
class Net(nn.Module):
    def __init__(self, input_size, device='cpu'):
        super().__init__()
        self.f1 = nn.Linear(input_size, 10, device=device)
        self.f2 = nn.Linear(10, 20, device=device)
        self.f3 = nn.Linear(20, 20, device=device)
        self.f4 = nn.Linear(20, 10, device=device)
        self.f5 = nn.Linear(10, 1, device=device)

        self.activation = nn.ReLU()
        # self.activation = nn.Sigmoid()

    def forward(self, x):
        x = self.activation(self.f1(x))
        x = self.activation(self.f2(x))
        x = self.activation(self.f3(x))
        x = self.activation(self.f4(x))
        x = self.f5(x)
        
        return x


In [4]:
X, y = load_diabetes(return_X_y = True, scaled=False)

X_scaler = MinMaxScaler()
X = X_scaler.fit_transform(X)

y_scaler = MinMaxScaler()
y = y_scaler.fit_transform(y.reshape((-1, 1)))

torch_data = TensorDataset(torch.Tensor(X).to(device), torch.Tensor(y).to(device))
data_loader = DataLoader(torch_data, batch_size=1000)

In [5]:
model = Net(input_size = X.shape[1], device=device)
loss_fn = nn.MSELoss()
opt = optim.SGD(model.parameters(), lr = 1e-2)

all_loss = {}
for epoch in range(100):
    print('epoch: ', epoch, end='')
    all_loss[epoch+1] = 0
    for batch_idx, (b_x, b_y) in enumerate(data_loader):
        pre = model(b_x)
        loss = loss_fn(pre, b_y)
        opt.zero_grad()
        loss.backward()

        # parameter update step based on optimizer
        opt.step()

        all_loss[epoch+1] += loss
    all_loss[epoch+1] /= len(data_loader)
    print(', loss: {}'.format(all_loss[epoch+1].detach().cpu().numpy().item()))

epoch:  0, loss: 0.24149444699287415
epoch:  1, loss: 0.23211970925331116
epoch:  2, loss: 0.22325003147125244
epoch:  3, loss: 0.21485504508018494
epoch:  4, loss: 0.20690634846687317
epoch:  5, loss: 0.1993778944015503
epoch:  6, loss: 0.1922454535961151
epoch:  7, loss: 0.1854863464832306
epoch:  8, loss: 0.17907896637916565
epoch:  9, loss: 0.1730034500360489
epoch:  10, loss: 0.16724146902561188
epoch:  11, loss: 0.16177572309970856
epoch:  12, loss: 0.1565900295972824
epoch:  13, loss: 0.15166905522346497
epoch:  14, loss: 0.14699888229370117
epoch:  15, loss: 0.14256924390792847
epoch:  16, loss: 0.13837528228759766
epoch:  17, loss: 0.1344127207994461
epoch:  18, loss: 0.1307019740343094
epoch:  19, loss: 0.1272537112236023
epoch:  20, loss: 0.12399578839540482
epoch:  21, loss: 0.12089867889881134
epoch:  22, loss: 0.11795064061880112
epoch:  23, loss: 0.11514674872159958
epoch:  24, loss: 0.11247842758893967
epoch:  25, loss: 0.10994468629360199
epoch:  26, loss: 0.1075386628

In [6]:
model = Net(input_size = X.shape[1], device=device)
loss_fn = nn.MSELoss()
opt = optim.Adam(model.parameters(), lr = 1e-2)

all_loss = {}
for epoch in range(100):
    print('epoch: ', epoch, end='')
    all_loss[epoch+1] = 0
    for batch_idx, (b_x, b_y) in enumerate(data_loader):
        pre = model(b_x)
        loss = loss_fn(pre, b_y)
        opt.zero_grad()
        loss.backward()

        # parameter update step based on optimizer
        opt.step()

        all_loss[epoch+1] += loss
    all_loss[epoch+1] /= len(data_loader)
    print(', loss: {}'.format(all_loss[epoch+1].detach().cpu().numpy().item()))

epoch:  0, loss: 0.0623062402009964
epoch:  1, loss: 0.05777193233370781
epoch:  2, loss: 0.05699910596013069
epoch:  3, loss: 0.057509712874889374
epoch:  4, loss: 0.057704027742147446
epoch:  5, loss: 0.05710136517882347
epoch:  6, loss: 0.05624129995703697
epoch:  7, loss: 0.05549498274922371
epoch:  8, loss: 0.054914798587560654
epoch:  9, loss: 0.05434662103652954
epoch:  10, loss: 0.05350053310394287
epoch:  11, loss: 0.052188653498888016
epoch:  12, loss: 0.050379347056150436
epoch:  13, loss: 0.0482637956738472
epoch:  14, loss: 0.046167828142642975
epoch:  15, loss: 0.04404192417860031
epoch:  16, loss: 0.0415063351392746
epoch:  17, loss: 0.038924235850572586
epoch:  18, loss: 0.03688272461295128
epoch:  19, loss: 0.035059381276369095
epoch:  20, loss: 0.033763617277145386
epoch:  21, loss: 0.03339529410004616
epoch:  22, loss: 0.033116746693849564
epoch:  23, loss: 0.03317643702030182
epoch:  24, loss: 0.03263682499527931
epoch:  25, loss: 0.0319666787981987
epoch:  26, loss