3131####################################################################################################
3232def train_continue ( wandb_id , epoch , Trainer , epoch_continue = - 1 ) :
3333
34- num_accs_per_task = int ( 4 / int ( os .environ .get ('SLURM_TASKS_PER_NODE' , '1' )[0 ] ))
35- device = init_torch ( num_accs_per_task )
34+ devices = init_torch ()
3635 with_ddp = True
3736 par_rank , par_size = setup_ddp ( with_ddp )
3837
3938 cf = Config ().load_json ( wandb_id )
4039
40+ cf .num_accs_per_task = len (devices ) # number of GPUs / accelerators per task
4141 cf .with_ddp = with_ddp
4242 cf .par_rank = par_rank
4343 cf .par_size = par_size
@@ -56,13 +56,13 @@ def train_continue( wandb_id, epoch, Trainer, epoch_continue = -1) :
5656 cf .with_mixed_precision = True
5757 if not hasattr (cf , 'years_val' ):
5858 cf .years_val = cf .years_test
59-
59+
6060 # any parameter in cf can be overwritten when training is continued, e.g. we can increase the
6161 # masking rate
6262 # cf.fields = [ [ 'specific_humidity', [ 1, 2048, [ ], 0 ],
6363 # [ 96, 105, 114, 123, 137 ],
6464 # [12, 6, 12], [3, 9, 9], [0.5, 0.9, 0.1, 0.05] ] ]
65-
65+
6666 setup_wandb ( cf .with_wandb , cf , par_rank , project_name = 'train' , mode = 'offline' )
6767 # resuming a run requires online mode, which is not available everywhere
6868 #setup_wandb( cf.with_wandb, cf, par_rank, wandb_id = wandb_id)
@@ -75,15 +75,14 @@ def train_continue( wandb_id, epoch, Trainer, epoch_continue = -1) :
7575 epoch_continue = epoch
7676
7777 # run
78- trainer = Trainer .load ( cf , wandb_id , epoch , device )
78+ trainer = Trainer .load ( cf , wandb_id , epoch , devices )
7979 print ( 'Loaded run \' {}\' at epoch {}.' .format ( wandb_id , epoch ))
8080 trainer .run ( epoch_continue )
8181
8282####################################################################################################
8383def train () :
8484
85- num_accs_per_task = int ( 4 / int ( os .environ .get ('SLURM_TASKS_PER_NODE' , '1' )[0 ] ))
86- device = init_torch ( num_accs_per_task )
85+ devices = init_torch ()
8786 with_ddp = True
8887 par_rank , par_size = setup_ddp ( with_ddp )
8988
@@ -93,7 +92,7 @@ def train() :
9392 cf = Config ()
9493 # parallelization
9594 cf .with_ddp = with_ddp
96- cf .num_accs_per_task = num_accs_per_task # number of GPUs / accelerators per task
95+ cf .num_accs_per_task = len ( devices ) # number of GPUs / accelerators per task
9796 cf .par_rank = par_rank
9897 cf .par_size = par_size
9998
@@ -108,32 +107,31 @@ def train() :
108107
109108 # cf.fields = [ [ 'temperature', [ 1, 1024, [ ], 0 ],
110109 # [ 96, 105, 114, 123, 137 ],
111- # [12, 6, 12 ], [3, 9, 9 ], [0.25 , 0.9, 0.1 , 0.05], 'local' ] ]
110+ # [12, 2, 4 ], [3, 27, 27 ], [0.5 , 0.9, 0.2 , 0.05], 'local' ] ]
112111 # cf.fields_prediction = [ [cf.fields[0][0], 1.] ]
113-
112+
114113 cf .fields = [ [ 'velocity_u' , [ 1 , 1024 , [ ], 0 ],
115114 [ 96 , 105 , 114 , 123 , 137 ],
116115 [12 , 3 , 6 ], [3 , 18 , 18 ], [0.5 , 0.9 , 0.2 , 0.05 ] ] ]
117116
118117 cf .fields_prediction = [ [cf .fields [0 ][0 ], 1. ] ]
119118
120-
119+
121120 # cf.fields = [ [ 'velocity_v', [ 1, 1024, [ ], 0 ],
122121 # [ 96, 105, 114, 123, 137 ],
123122 # [12, 3, 6], [3, 18, 18], [0.25, 0.9, 0.1, 0.05] ] ]
124123
125- # cf.fields = [ [ 'velocity_z', [ 1, 512 , [ ], 0 ],
124+ # cf.fields = [ [ 'velocity_z', [ 1, 1024 , [ ], 0 ],
126125 # [ 96, 105, 114, 123, 137 ],
127- # [12, 6, 12 ], [3, 9, 9 ], [0.25, 0.9, 0.1, 0.05] ] ]
126+ # [12, 3, 6 ], [3, 18, 18 ], [0.25, 0.9, 0.1, 0.05] ] ]
128127
129128 # cf.fields = [ [ 'specific_humidity', [ 1, 1024, [ ], 0 ],
130129 # [ 96, 105, 114, 123, 137 ],
131- # [12, 6, 12], [3, 9, 9], [0.25, 0.9, 0.1, 0.05], 'local' ] ]
132- # [12, 2, 4], [3, 27, 27], [0.5, 0.9, 0.1, 0.05], 'local' ] ]
133-
130+ # [12, 3, 6], [3, 18, 18], [0.25, 0.9, 0.1, 0.05] ] ]
131+
134132 cf .fields_targets = []
135133
136- cf .years_train = list ( range ( 2010 , 2021 ))
134+ cf .years_train = list ( range ( 1979 , 2021 ))
137135 cf .years_val = [2021 ] #[2018]
138136 cf .month = None
139137 cf .geo_range_sampling = [[ - 90. , 90. ], [ 0. , 360. ]]
@@ -143,7 +141,7 @@ def train() :
143141 # training params
144142 cf .batch_size_validation = 1 #64
145143 cf .batch_size = 96
146- cf .num_epochs = 128
144+ cf .num_epochs = 400 # 128
147145 cf .num_samples_per_epoch = 4096 * 12
148146 cf .num_samples_validate = 128 * 12
149147 cf .num_loader_workers = 8
@@ -176,7 +174,7 @@ def train() :
176174 cf .net_tail_num_nets = 16
177175 cf .net_tail_num_layers = 0
178176 # loss
179- cf .losses = ['mse_ensemble' , 'stats' ] # mse, mse_ensemble, stats, crps, weighted_mse
177+ cf .losses = ['mse_ensemble' , 'stats' ] # mse, mse_ensemble, stats, crps, weighted_mse
180178 # training
181179 cf .optimizer_zero = False
182180 cf .lr_start = 5. * 10e-7
@@ -185,6 +183,7 @@ def train() :
185183 cf .weight_decay = 0.05 #0.1
186184 cf .lr_decay_rate = 1.025
187185 cf .lr_start_epochs = 3
186+ cf .model_log_frequency = 256 #save checkpoint every X batches
188187 # BERT
189188 # strategies: 'BERT', 'forecast', 'temporal_interpolation'
190189 cf .BERT_strategy = 'BERT'
@@ -218,7 +217,7 @@ def train() :
218217 # # # cf.file_path = '/p/scratch/atmo-rep/data/era5_1deg/months/era5_y2021_res025_chunk8.zarr'
219218 # # cf.file_path = '/ec/res4/scratch/nacl/atmorep/era5_y2021_res025_chunk8_lat180_lon180.zarr'
220219 # # # cf.file_path = '/ec/res4/scratch/nacl/atmorep/era5_y2021_res025_chunk16.zarr'
221- cf .file_path = '/gpfs/scratch/ehpc03/era5_y2010_2021_res025_chunk8 .zarr/'
220+ cf .file_path = '/gpfs/scratch/ehpc03/era5_y1979_2021_res025_chunk8 .zarr/'
222221 # # # in steps x lat_degrees x lon_degrees
223222 cf .n_size = [36 , 0.25 * 9 * 6 , 0.25 * 9 * 12 ]
224223
@@ -230,7 +229,7 @@ def train() :
230229 cf .write_json ( wandb )
231230 cf .print ()
232231
233- trainer = Trainer_BERT ( cf , device ).create ()
232+ trainer = Trainer_BERT ( cf , devices ).create ()
234233 trainer .run ()
235234
236235####################################################################################################
@@ -239,8 +238,8 @@ def train() :
239238 try :
240239
241240 train ()
242-
243- # wandb_id, epoch, epoch_continue = '1jh2qvrx ', 392, 392
241+
242+ # wandb_id, epoch, epoch_continue = 'gxfywjzl ', 127, 127
244243 # Trainer = Trainer_BERT
245244 # train_continue( wandb_id, epoch, Trainer, epoch_continue)
246245
0 commit comments