-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathjelly_ai_future.py
More file actions
80 lines (59 loc) · 3.14 KB
/
Copy pathjelly_ai_future.py
File metadata and controls
80 lines (59 loc) · 3.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
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
75
76
77
78
79
80
# import tensorflow as tf
import numpy as np
import joblib
import tensorflow as tf
from sklearn.preprocessing import MinMaxScaler
from tensorflow.keras.models import load_model
class JellyAIFutureModel:
models = {}
def loadModels(self):
self.models['노무라입깃해파리'] = {
'appear2': [load_model('./models/N_best_model_iw60_fh2.keras'),60], #모델이 요구하는 데이터 사이즈도 같이 저장
'appear7': [load_model('./models/N_best_model_iw21_fh7.keras'),21]
}
self.models['보름달물해파리'] = {
'appear2': [load_model('./models/B_best_model_iw50_fh2.keras'),50],
'appear7': [load_model('./models/B_best_model_iw28_fh7.keras'),28]
}
self.scaler = joblib.load('./models/scaler.save')
def predictJelly(self, jelly_type, datas):#데이터를 가져올때 61일 데이터 넘파이 어레이 형태로 가져오기
if jelly_type not in self.models:
raise ValueError(f"해당 해파리 모델({jelly_type})이 존재하지 않습니다.")
nparr = datas.to_numpy()
datas=self.scaler.transform(nparr)
model_set = self.models[jelly_type]
appear_pred=[]
datas=datas[-model_set['appear2'][1]-1:-1]
temp=model_set['appear2'][0].predict(datas.reshape(1, datas.shape[0], datas.shape[1]))
#2일
datas=datas[-model_set['appear2'][1]:]
res=model_set['appear2'][0].predict(datas.reshape(1, datas.shape[0], datas.shape[1]))
appear_pred.append(res[0][0])
appear_pred[0] = np.where(appear_pred[0] >= 0.5, 1, 0).item()
appear_pred.append(res[0][1])
appear_pred[1] = np.where(appear_pred[1] >= 0.5, 1, 0).item()
#7일
datas=datas[-model_set['appear7'][1]:]
res=model_set['appear7'][0].predict(datas.reshape(1, datas.shape[0], datas.shape[1]))
appear_pred.append(res[0][2])
appear_pred[2] = np.where(appear_pred[2] >= 0.5, 1, 0).item()
appear_pred.append(res[0][3])
appear_pred[3] = np.where(appear_pred[3] >= 0.5, 1, 0).item()
appear_pred.append(res[0][4])
appear_pred[4] = np.where(appear_pred[4] >= 0.5, 1, 0).item()
appear_pred.append(res[0][5])
appear_pred[5] = np.where(appear_pred[5] >= 0.5, 1, 0).item()
appear_pred.append(res[0][6])
appear_pred[6] = np.where(appear_pred[6] >= 0.5, 1, 0).item()
return appear_pred #데이터를 내보낼때 1,2,3,4,5,6,7일 예측 리스트 형태로 내보내기
if __name__ == '__main__':
model = JellyAIFutureModel()
model.loadModels()
data = np.array([[0.42678829, 0.90817016, 0.06811594, 0.68564412, 0.46262909, 0.46666667, 0.09313725, 0.81476323, 0.02418605]])
jelly_types = ['노무라입깃해파리', '보름달물해파리']
for jelly_type in jelly_types:
print(f"====={jelly_type}=====")
appear_pred, density_pred, percent_loc = model.predictJelly(jelly_type, data)
print(f"appear_pred: {appear_pred}")
print(f"density_pred: {density_pred}")
print(f"percent_loc: {percent_loc}")