-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathDataTraining.py
More file actions
74 lines (55 loc) · 2.48 KB
/
Copy pathDataTraining.py
File metadata and controls
74 lines (55 loc) · 2.48 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
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
import pickle
from Network import Network
from XGBoost import XGBobj
import torch
import numpy as np
def getData():
sickFileTrain = open('Data/network_sick_train.data', 'rb')
sickTrain = pickle.load(sickFileTrain)
sickFileVal = open('Data/network_sick_val.data', 'rb')
sickVal = pickle.load(sickFileVal)
sickFileTest = open('Data/network_sick_test.data', 'rb')
sickTest = pickle.load(sickFileTest)
healthyFileTrain = open('Data/network_healthy_train.data', 'rb')
healthyTrain = pickle.load(healthyFileTrain)
healthyFileVal = open('Data/network_healthy_val.data', 'rb')
healthyVal = pickle.load(healthyFileVal)
healthyFileTest = open('Data/network_healthy_test.data', 'rb')
healthyTest = pickle.load(healthyFileTest)
print(sickTrain.shape)
print(sickVal.shape)
print(sickTest.shape)
print(healthyTrain.shape)
print(healthyVal.shape)
print(healthyTest.shape)
return (sickTrain,sickVal,sickTest,healthyTrain,healthyVal,healthyTest)
def duplicateRandom(data, numDuplicates):
originalData = data
for _ in range(numDuplicates):
data = torch.cat((data, originalData+(torch.rand(originalData.shape)-0.5)/1.5),0)
return data
def augmentData(sickTrain, healthyTrain, numDuplicates):
scaleAmount = int(healthyTrain.shape[0]/sickTrain.shape[0])
sickTrain = duplicateRandom(sickTrain, numDuplicates*scaleAmount)
healthyTrain = duplicateRandom(healthyTrain, numDuplicates)
return (sickTrain,healthyTrain)
def NNTrain():
sickTrain,sickVal,sickTest,healthyTrain,healthyVal,healthyTest = getData()
n = Network()
n.train(healthyTrain,sickTrain,healthyVal,sickVal,8,0.0001,0.9,1000)
def XGBTrain():
sickTrain,sickVal,sickTest,healthyTrain,healthyVal,healthyTest = getData()
sickTrain,healthyTrain = augmentData(sickTrain,healthyTrain, 10)
x = XGBobj()
train_X = torch.cat((sickTrain,healthyTrain),dim = 0)
val_X = torch.cat((sickVal,healthyVal),dim = 0)
test_X = torch.cat((sickTest,healthyTest),dim = 0)
train_Y = torch.cat((torch.ones([sickTrain.shape[0]]),torch.zeros([healthyTrain.shape[0]])), dim=0)
val_Y = torch.cat((torch.ones([sickVal.shape[0]]),torch.zeros([healthyVal.shape[0]])), dim=0)
test_Y = torch.cat((torch.ones([sickTest.shape[0]]),torch.zeros([healthyTest.shape[0]])), dim=0)
print(train_Y)
print(val_Y)
print(test_Y)
x.train(train_X.numpy(),train_Y.numpy(),val_X.numpy(),val_Y.numpy(),10000)
x.test(val_X.numpy(),val_Y.numpy())
XGBTrain()