Skip to content

Commit 3efcb12

Browse files
authored
Merge pull request #37 from clessig/iluise/head
move to Dask array + first working version of the multiformer
2 parents e944cb8 + 09f3b96 commit 3efcb12

11 files changed

Lines changed: 286 additions & 211 deletions

File tree

‎atmorep/core/atmorep_model.py‎

Lines changed: 23 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
import torch
1818
import numpy as np
1919
import code
20+
import os
2021
# code.interact(local=locals())
2122

2223
# import horovod.torch as hvd
@@ -238,8 +239,18 @@ def create( self, devices, load_pretrained=True) :
238239
if len(field_info[1]) > 4 and load_pretrained :
239240
# TODO: inconsistent with embeds_token_info -> version that can handle both
240241
# we could imply use the file name: embed_token_info vs embeds_token_info
241-
name = 'AtmoRep' + '_embed_token_info'
242+
name = 'AtmoRep' + '_embeds_token_info'
243+
if not os.path.exists(get_model_filename( name, field_info[1][4][0], field_info[1][4][1])):
244+
name = 'AtmoRep' + '_embed_token_info'
245+
242246
mloaded = torch.load( get_model_filename( name, field_info[1][4][0], field_info[1][4][1]))
247+
248+
if "weight" not in mloaded.keys(): #TODO: get rid of this
249+
mloaded["weight"] = mloaded["0.weight"]
250+
mloaded["bias"] = mloaded["0.bias"]
251+
del mloaded["0.weight"]
252+
del mloaded["0.bias"]
253+
243254
self.embeds_token_info[-1].load_state_dict( mloaded)
244255
print( 'Loaded embed_token_info from id = {}.'.format( field_info[1][4][0] ) )
245256
else :
@@ -260,7 +271,7 @@ def create( self, devices, load_pretrained=True) :
260271
if len(field_info[1]) > 4 and load_pretrained :
261272
self.load_block( field_info, 'encoder', self.encoders[-1])
262273
self.embeds.append( self.encoders[-1].embed)
263-
274+
264275
# indices of coupled fields for efficient access in forward
265276
self.fields_coupling_idx.append( [field_idx])
266277
for field_coupled in field_info[1][2] :
@@ -353,9 +364,17 @@ def load_block( self, field_info, block_name, block ) :
353364
param[ : , b_loaded[k].shape[1] : ] = 0.01 * torch.rand( param.shape[0],
354365
param.shape[1] - b_loaded[k].shape[1])
355366
keys_del += [ k ]
367+
368+
#for backward compatibility. solved in new runs
369+
if 'proj_heads_other' in name:
370+
for k in b_loaded.keys() :
371+
if name == k :
372+
if b_loaded[name].shape[0] == 0:
373+
keys_del += [ name ]
374+
356375
for k in keys_del :
357376
del b_loaded[k]
358-
377+
359378
# use strict=False so that differing blocks, e.g. through coupling, are ignored
360379
mkeys, _ = block.load_state_dict( b_loaded, False)
361380

@@ -394,11 +413,7 @@ def translate_weights(self, mloaded, mkeys, ukeys):
394413
del mloaded[f'encoders.0.heads.{layer}.heads_other.{head}.proj_ks.weight']
395414
del mloaded[f'encoders.0.heads.{layer}.heads_other.{head}.proj_vs.weight']
396415

397-
else:
398-
dim_mw = self.encoders[0].heads[0].proj_heads_other[0].weight.shape
399-
mw = torch.tensor(np.zeros(dim_mw))
400-
401-
mloaded[f'encoders.0.heads.{layer}.proj_heads_other.0.weight'] = mw
416+
mloaded[f'encoders.0.heads.{layer}.proj_heads_other.0.weight'] = mw
402417

403418
#decoder
404419
for iblock in range(0, 19, 2) :

‎atmorep/core/evaluate.py‎

Lines changed: 25 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -19,24 +19,24 @@
1919

2020
if __name__ == '__main__':
2121

22-
# models for individual fields
22+
# arXiv 2023: models for individual fields
2323
#model_id = '4nvwbetz' # vorticity
2424
#model_id = 'oxpycr7w' # divergence
2525
#model_id = '1565pb1f' # specific_humidity
2626
#model_id = '3kdutwqb' # total precip
27-
model_id = 'dys79lgw' # velocity_u
27+
#model_id = 'dys79lgw' # velocity_u
2828
#model_id = '22j6gysw' # velocity_v
29-
# model_id = '15oisw8d' # velocity_z
30-
#model_id = '3qou60es' # temperature (also 2147fkco)
29+
#model_id = '15oisw8d' # velocity_z
30+
#model_id = '3qou60es' # temperature
3131
#model_id = '2147fkco' # temperature (also 2147fkco)
32-
33-
# multi-field configurations with either velocity or voritcity+divergence
34-
#model_id = '1jh2qvrx' # multiformer, velocity
35-
# model_id = 'wqqy94oa' # multiformer, vorticity
36-
#model_id = '3cizyl1q' # 3 field config: u,v,T
37-
# model_id = '1v4qk0qx' # pre-trained, 3h forecasting
38-
# model_id = '1m79039j' # pre-trained, 6h forecasting
39-
#model_id='34niv2nu'
32+
33+
# new runs 2024
34+
#model_id='j8dwr5qj' #velocity_u
35+
#model_id='0tlnm5up' #velocity_v
36+
#model_id='v63l01zu' #specific humidity
37+
#model_id='9l1errbo' #velocity_z
38+
model_id='7ojls62c' #temperature 1024
39+
4040
# supported modes: test, forecast, fixed_location, temporal_interpolation, global_forecast,
4141
# global_forecast_range
4242
# options can be used to over-write parameters in config; some modes also have specific options,
@@ -45,7 +45,7 @@
4545
#Add 'attention' : True to options to store the attention maps. NB. supported only for single field runs.
4646

4747
# BERT masked token model
48-
mode, options = 'BERT', {'years_test' : [2021], 'num_samples_validate' : 128, 'with_pytest' : True }
48+
#mode, options = 'BERT', {'years_val' : [2021], 'num_samples_validate' : 128, 'with_pytest' : True}
4949

5050
# BERT forecast mode
5151
#mode, options = 'forecast', {'forecast_num_tokens' : 2, 'num_samples_validate' : 128, 'with_pytest' : True }
@@ -55,19 +55,19 @@
5555
#mode, options = 'temporal_interpolation', {'idx_time_mask': [5,6,7], 'num_samples_validate' : 128, 'with_pytest' : True}
5656

5757
# BERT forecast with patching to obtain global forecast
58-
# mode, options = 'global_forecast', {
59-
# 'dates' : [[2021, 1, 10, 18]],
60-
# # # 'dates' : [ #[2021, 1, 10, 18]
61-
# # # [2021, 1, 10, 12] , [2021, 1, 11, 0], [2021, 1, 11, 12], [2021, 1, 12, 0], [2021, 1, 12, 12], [2021, 1, 13, 0],
62-
# # # [2021, 4, 10, 12], [2021, 4, 11, 0], [2021, 4, 11, 12], [2021, 4, 12, 0], [2021, 4, 12, 12], [2021, 4, 13, 0],
63-
# # # [2021, 7, 10, 12], [2021, 7, 11, 0], [2021, 7, 11, 12], [2021, 7, 12, 0], [2021, 7, 12, 12], [2021, 7, 13, 0],
64-
# # # [2021, 10, 10, 12], [2021, 10, 11, 0], [2021, 10, 11, 12], [2021, 10, 12, 0], [2021, 10, 12, 12], [2021, 10, 13, 0]
65-
# # # ],
66-
# 'token_overlap' : [0, 0],
67-
# 'forecast_num_tokens' : 2,
68-
# 'with_pytest' : True }
58+
mode, options = 'global_forecast', {
59+
#'dates' : [[2021, 2, 10, 12]]
60+
'dates' : [
61+
[2021, 1, 10, 12] , [2021, 1, 11, 0], [2021, 1, 11, 12], [2021, 1, 12, 0], #[2021, 1, 12, 12], [2021, 1, 13, 0],
62+
[2021, 4, 10, 12], [2021, 4, 11, 0], [2021, 4, 11, 12], [2021, 4, 12, 0], #[2021, 4, 12, 12], [2021, 4, 13, 0],
63+
[2021, 7, 10, 12], [2021, 7, 11, 0], [2021, 7, 11, 12], [2021, 7, 12, 0], #[2021, 7, 12, 12], [2021, 7, 13, 0],
64+
[2021, 10, 10, 12], [2021, 10, 11, 0], [2021, 10, 11, 12], #[2021, 10, 12, 0], [2021, 10, 12, 12], [2021, 10, 13, 0]
65+
],
66+
'token_overlap' : [0, 0],
67+
'forecast_num_tokens' : 2,
68+
'with_pytest' : True }
6969

70-
file_path = '/gpfs/scratch/ehpc03/era5_y2010_2021_res025_chunk8.zarr'
70+
file_path = '/gpfs/scratch/ehpc03/era5_y1979_2021_res025_chunk8.zarr'
7171

7272
now = time.time()
7373
Evaluator.evaluate( mode, model_id, file_path, options)

‎atmorep/core/evaluator.py‎

Lines changed: 10 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -69,18 +69,13 @@ def run( cf, model_id, model_epoch, devices) :
6969
@staticmethod
7070
def evaluate( mode, model_id, file_path, args = {}, model_epoch=-2) :
7171

72-
# SLURM_TASKS_PER_NODE is controlled by #SBATCH --ntasks-per-node=1; should be 1 for multiformer
73-
with_ddp = True
74-
if '-1' == os.environ.get('MASTER_ADDR', '-1') :
75-
with_ddp = False
76-
num_accs_per_task = 1
77-
else :
78-
num_accs_per_task = int( 4 / int( os.environ.get('SLURM_TASKS_PER_NODE', '1')[0] ))
79-
devices = init_torch( num_accs_per_task)
80-
#devices = ['cuda:1']
81-
72+
devices = init_torch()
73+
with_ddp = True
8274
par_rank, par_size = setup_ddp( with_ddp)
75+
8376
cf = Config().load_json( model_id)
77+
78+
cf.num_accs_per_task = len(devices)
8479
cf.file_path = file_path
8580
cf.with_wandb = True
8681
cf.with_ddp = with_ddp
@@ -120,7 +115,7 @@ def evaluate( mode, model_id, file_path, args = {}, model_epoch=-2) :
120115
if cf.with_pytest:
121116
fields = [field[0] for field in cf.fields_prediction]
122117
for field in fields:
123-
pytest.main(["-x", "./atmorep/tests/validation_test.py", "--field", field, "--model_id", cf.wandb_id, "--strategy", cf.BERT_strategy])
118+
pytest.main(["-x", "-s", "./atmorep/tests/validation_test.py", "--field", field, "--model_id", cf.wandb_id, "--strategy", cf.BERT_strategy])
124119

125120
##############################################
126121
@staticmethod
@@ -155,11 +150,11 @@ def global_forecast( cf, model_id, model_epoch, devices, args = {}) :
155150
cf.batch_size_test = 24
156151
cf.num_loader_workers = 12 #1
157152
cf.log_test_num_ranks = 1
153+
154+
#TODO: temporary solution. Add support for batch_size > 1
155+
cf.batch_size_validation = 1 #64
156+
cf.batch_size = 1
158157

159-
if not hasattr(cf, 'batch_size'):
160-
cf.batch_size = 196 #14
161-
if not hasattr(cf, 'batch_size_validation'):
162-
cf.batch_size_validation = 1 #64
163158
if not hasattr(cf, 'num_samples_validate'):
164159
cf.num_samples_validate = 196
165160
#if not hasattr(cf,'with_mixed_precision'):

‎atmorep/core/train.py‎

Lines changed: 22 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -31,13 +31,13 @@
3131
####################################################################################################
3232
def 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
####################################################################################################
8383
def 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

Comments
 (0)