-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
41 lines (36 loc) · 1.14 KB
/
Copy pathmain.py
File metadata and controls
41 lines (36 loc) · 1.14 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
import torch
from timeit import default_timer as timer
from model import Modl
from utils import how_long_did_it_take, evaluate_the_model
from data import create_dataset, get_the_dataloader
from config import *
from train import train_the_model
def launch_model():
#1.
training_dataset, testing_dataset = create_dataset()
print(training_dataset)
print(testing_dataset)
#2.
[training_dataloader, testing_dataloader]=get_the_dataloader([training_dataset, testing_dataset])
#3.
Model=Modl(Input_features, Hidden_units, Output_features)
#4.
start=timer()
#5.
train_the_model(Model, EPOCHS, training_dataloader, testing_dataloader)
#6.
end=timer()
how_long_did_it_take(start, end)
#7.
Model.eval()
all_preds=[]
all_targets=[]
with torch.inference_mode():
for batch in testing_dataloader:
images = batch['img'].to(DEVICE)
preds = torch.argmax(Model(images), dim=1)
all_preds.extend(preds.cpu().numpy())
all_targets.extend(batch['label'].numpy())
evaluate_the_model(torch.tensor(all_preds), torch.tensor(all_targets))
if '__main__'==__name__:
launch_model()