-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
78 lines (64 loc) · 2.68 KB
/
Copy pathmain.py
File metadata and controls
78 lines (64 loc) · 2.68 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
import numpy as np
import torch
from visdom import Visdom
from torch.nn import L1Loss
from torch.optim import Adam
from torch import device, cuda
from model import RUNet
from dataset import Data_Set
from train import train_model
class Task:
def __init__(self, init, variable, validation_year: list, test_year: list):
self.info = (init,variable)
self.validation_year = validation_year
self.model = RUNet(1, 1)
self.data = (Data_Set(validation_year=validation_year,test_year=test_year,info=self.info),
Data_Set(validation_year=validation_year,test_year=test_year,info=self.info
,train=False),
Data_Set(validation_year=validation_year,test_year=test_year,info=self.info,
train=False,test=True))
self.data_path = self.data[0].data_path()
self.loss_function = L1Loss()
self.optimizer = Adam
self.adjustable_parameters = {"lr": 0.0002, "epoch": 100, "batch_size": 30}
def __str__(self):
return "model: {}\nparameters: {}".format(self.model, self.adjustable_parameters)
if __name__ == '__main__':
# init_mon = 'dec'
# var = 't2m'
# init_mon = 'dec'
# var = 'tp'
# init_mon = 'jun'
# var = 't2m'
# init_mon = 'jun'
# var = 'tp'
torch.manual_seed(42)
torch.cuda.manual_seed(42)
torch.cuda.manual_seed_all(42)
np.random.seed(42)
validate_year_all = [[2018, 2019, 2020],
[1993, 1994, 1995],
[1998, 1999, 2000],
[2003, 2004, 2005],
[2008, 2009, 2010],
[2013, 2014, 2015]]
test_year_all = [[1991, 1992, 1993, 1994, 1995],
[1996, 1997, 1998, 1999, 2000],
[2001, 2002, 2003, 2004, 2005],
[2006, 2007, 2008, 2009, 2010],
[2011, 2012, 2013, 2014, 2015],
[2016, 2017, 2018, 2019, 2020]]
init_list = ['dec','jun']
var_list = ['t2m', 'tp']
for init_mon in init_list:
for var in var_list:
for i in range(6):
validation_year = validate_year_all[i]
test_year = test_year_all[i]
dev = device("cuda" if cuda.is_available() else "cpu")
device_ids = [0]
nwp_improving = Task(init_mon,var,validation_year, test_year=test_year)
print(i, dev, nwp_improving, sep='\n')
model_env = 'nwp_improve_{}_{}'.format(init_mon,var)
visualizer = Visdom(env=f'{model_env}')
train_model(nwp_improving, device_ids, visualizer, i)