| |
| |
| |
| |
| |
|
|
| """ |
| Example of training a DLWP model on the cubed sphere with the Keras functional API. |
| """ |
|
|
| import argparse |
| import os |
| import shutil |
| import time |
| import numpy as np |
| import pandas as pd |
| import xarray as xr |
| from datetime import datetime |
| from DLWP.model import DLWPFunctional, ArrayDataGenerator, tf_data_generator |
| from DLWP.model.preprocessing import get_constants, prepare_data_array |
| from DLWP.util import save_model |
| from tensorflow.keras.callbacks import TensorBoard |
| from azureml.core import Run |
|
|
| from tensorflow.keras.layers import Input, UpSampling3D, AveragePooling3D, concatenate, ReLU, Reshape, Concatenate, \ |
| Permute |
| from DLWP.custom import CubeSpherePadding2D, CubeSphereConv2D, RNNResetStates, EarlyStoppingMin, \ |
| RunHistory, SaveWeightsOnEpoch, GeneratorEpochEnd |
| from tensorflow.keras.models import Model |
| from tensorflow.keras.optimizers import Adam |
|
|
| import tensorflow as tf |
| |
| tf.compat.v1.logging.set_verbosity(tf.compat.v1.logging.ERROR) |
| |
| |
| |
| |
| |
|
|
|
|
| |
|
|
| parser = argparse.ArgumentParser() |
| parser.add_argument('--root-directory', type=str, dest='root_directory', default='.', |
| help='Destination root data directory on Azure Blob storage') |
| parser.add_argument('--predictor-file', type=str, dest='predictor_file', |
| help='Path and name of data file in root-directory') |
| parser.add_argument('--model-file', type=str, dest='model_file', |
| help='Path and name of model save file in root-directory') |
| parser.add_argument('--log-directory', type=str, dest='log_directory', default='./logs', |
| help='Destination for log files in root-directory') |
| parser.add_argument('--temp-dir', type=str, dest='temp_dir', default='None', |
| help='If specified, copies the predictor file here for use during training (e.g., fast SSD)') |
| parser.add_argument('--seed', type=int, dest='seed', default=-1, |
| help='Specify random number seed >= 0') |
|
|
| args = parser.parse_args() |
| if args.temp_dir != 'None': |
| os.makedirs(args.temp_dir, exist_ok=True) |
|
|
| if args.seed >= 0: |
| np.random.seed(args.seed) |
| tf.compat.v1.set_random_seed(args.seed) |
|
|
|
|
| |
|
|
| |
| root_directory = args.root_directory |
| predictor_file = os.path.join(root_directory, args.predictor_file) |
| model_file = os.path.join(root_directory, args.model_file) |
| log_directory = os.path.join(root_directory, args.log_directory) |
| reverse_lat = False |
|
|
| |
| constant_fields = [ |
| (os.path.join(root_directory, 'era5/era5_2deg_3h_CS2_land_sea_mask.nc'), 'lsm'), |
| (os.path.join(root_directory, 'era5/era5_2deg_3h_CS2_scaled_topo.nc'), 'z') |
| ] |
|
|
| |
| |
| cnn_model_name = 'unet2' |
| base_filter_number = 32 |
| min_epochs = 100 |
| max_epochs = 1000 |
| patience = 50 |
| batch_size = 64 |
| loss_by_step = None |
| shuffle = True |
| independent_north_pole = False |
|
|
| |
| |
| |
| |
| |
| io_selection = {'varlev': ['z/500', 'tau/300-700', 'z/1000', 't2m/0']} |
| io_time_steps = 2 |
| integration_steps = 2 |
| data_interval = 2 |
| |
| add_solar = True |
|
|
| |
| n_gpu = 1 |
|
|
| |
| use_mp_optimizer = True |
|
|
| |
| |
| |
| validation_set = list(pd.date_range(datetime(2013, 1, 1, 0), datetime(2013, 2, 28, 18), freq='3H')) |
| train_set = list(pd.date_range(datetime(1979, 1, 1, 0), datetime(1979, 12, 31, 18), freq='3H')) |
|
|
|
|
| |
|
|
| if args.temp_dir != 'None': |
| start_time = time.time() |
| new_predictor_file = os.path.join(args.temp_dir, args.predictor_file.split(os.sep)[-1]) |
| print('Copying predictor file to %s...' % new_predictor_file) |
| if os.path.isfile(new_predictor_file): |
| print('File already exists!') |
| else: |
| shutil.copy(predictor_file, new_predictor_file, follow_symlinks=True) |
| data = xr.open_dataset(new_predictor_file, chunks={'sample': batch_size}) |
| total_time = time.time() - start_time |
| print('Time to copy file: %d m %0.2f s' % (np.floor(total_time / 60), total_time % 60)) |
| else: |
| data = xr.open_dataset(predictor_file, chunks={'sample': 1}) |
|
|
| if reverse_lat: |
| data.lat.load() |
| data.lat[:] = -1. * data.lat.values |
|
|
| has_constants = not(not constant_fields) |
| constants = get_constants(constant_fields or None) |
|
|
|
|
| |
|
|
| dlwp = DLWPFunctional(is_convolutional=True, is_recurrent=False, time_dim=io_time_steps) |
|
|
| |
| if train_set is None: |
| train_set = np.isin(data.sample.values, np.array(validation_set, dtype='datetime64[ns]'), |
| assume_unique=True, invert=True) |
| validation_data = data.sel(sample=validation_set) |
| train_data = data.sel(sample=train_set) |
|
|
| |
| print('Loading data to memory...') |
| start_time = time.time() |
| train_array, input_ind, output_ind, sol = prepare_data_array(train_data, input_sel=io_selection, |
| output_sel=io_selection, add_insolation=add_solar) |
| generator = ArrayDataGenerator(dlwp, train_array, rank=3, input_slice=input_ind, output_slice=output_ind, |
| input_time_steps=io_time_steps, output_time_steps=io_time_steps, |
| sequence=integration_steps, interval=data_interval, insolation_array=sol, |
| batch_size=batch_size, shuffle=shuffle, constants=constants, channels_last=True, |
| drop_remainder=True) |
| input_names = ['main_input'] + ['solar_%d' % i for i in range(1, integration_steps)] + \ |
| (['constants'] if has_constants else []) |
| tf_train_data = tf_data_generator(generator, batch_size=batch_size, input_names=input_names) |
| if validation_data is not None: |
| print('Loading validation data to memory...') |
| val_array, input_ind, output_ind, sol = prepare_data_array(validation_data, input_sel=io_selection, |
| output_sel=io_selection, add_insolation=add_solar) |
| val_generator = ArrayDataGenerator(dlwp, val_array, rank=3, input_slice=input_ind, output_slice=output_ind, |
| input_time_steps=io_time_steps, output_time_steps=io_time_steps, |
| sequence=integration_steps, interval=data_interval, insolation_array=sol, |
| batch_size=batch_size, shuffle=False, constants=constants, channels_last=True) |
| tf_val_data = tf_data_generator(val_generator, input_names=input_names) |
| else: |
| tf_val_data = None |
|
|
| total_time = time.time() - start_time |
| print('Time to load data: %d m %0.2f s' % (np.floor(total_time / 60), total_time % 60)) |
|
|
|
|
| |
|
|
| |
| cs = generator.convolution_shape |
| cso = generator.output_convolution_shape |
| input_solar = (integration_steps > 1 and (isinstance(add_solar, str) or add_solar)) |
|
|
| |
| main_input = Input(shape=cs, name='main_input') |
| if input_solar: |
| solar_inputs = [Input(shape=generator.insolation_shape, name='solar_%d' % d) for d in range(1, integration_steps)] |
| if has_constants: |
| constant_input = Input(shape=(6, 48, 48, 2), name='constants') |
| cube_padding_1 = CubeSpherePadding2D(1, data_format='channels_last') |
| pooling_2 = AveragePooling3D((1, 2, 2), data_format='channels_last') |
| up_sampling_2 = UpSampling3D((1, 2, 2), data_format='channels_last') |
| relu = ReLU(negative_slope=0.1, max_value=10.) |
| conv_kwargs = { |
| 'dilation_rate': 1, |
| 'padding': 'valid', |
| 'activation': 'linear', |
| 'data_format': 'channels_last', |
| 'independent_north_pole': independent_north_pole, |
| 'flip_north_pole': not independent_north_pole |
| } |
| skip_connections = 'unet' in cnn_model_name.lower() |
| conv_2d_1 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_1_2 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_1_3 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_2 = CubeSphereConv2D(base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_2_2 = CubeSphereConv2D(base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_2_3 = CubeSphereConv2D(base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_3 = CubeSphereConv2D(base_filter_number * 4, 3, **conv_kwargs) |
| conv_2d_3_2 = CubeSphereConv2D(base_filter_number * 4, 3, **conv_kwargs) |
| conv_2d_4 = CubeSphereConv2D(base_filter_number * 4 if skip_connections else base_filter_number * 8, 3, **conv_kwargs) |
| conv_2d_4_2 = CubeSphereConv2D(base_filter_number * 8, 3, **conv_kwargs) |
| conv_2d_5 = CubeSphereConv2D(base_filter_number * 2 if skip_connections else base_filter_number * 4, 3, **conv_kwargs) |
| conv_2d_5_2 = CubeSphereConv2D(base_filter_number * 4, 3, **conv_kwargs) |
| conv_2d_5_3 = CubeSphereConv2D(base_filter_number * 4, 3, **conv_kwargs) |
| conv_2d_6 = CubeSphereConv2D(base_filter_number if skip_connections else base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_6_2 = CubeSphereConv2D(base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_6_3 = CubeSphereConv2D(base_filter_number * 2, 3, **conv_kwargs) |
| conv_2d_7 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_7_2 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_7_3 = CubeSphereConv2D(base_filter_number, 3, **conv_kwargs) |
| conv_2d_8 = CubeSphereConv2D(cso[-1], 1, name='output', **conv_kwargs) |
|
|
|
|
| |
|
|
| def basic(x): |
| x = cube_padding_1(x) |
| x = relu(conv_2d_1(x)) |
| x = pooling_2(x) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_2(x)) |
| x = pooling_2(x) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_3(x)) |
| x = up_sampling_2(x) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6(x)) |
| x = up_sampling_2(x) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_2(x)) |
| x = conv_2d_8(x) |
| return x |
|
|
|
|
| def unet(x): |
| x0 = cube_padding_1(x) |
| x0 = relu(conv_2d_1(x0)) |
| x1 = pooling_2(x0) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2(x1)) |
| x2 = pooling_2(x1) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_3(x2)) |
| x2 = up_sampling_2(x2) |
| x = concatenate([x2, x1], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6(x)) |
| x = up_sampling_2(x) |
| x = concatenate([x, x0], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_2(x)) |
| x = conv_2d_8(x) |
| return x |
|
|
|
|
| def unet2(x): |
| x0 = cube_padding_1(x) |
| x0 = relu(conv_2d_1(x0)) |
| x0 = cube_padding_1(x0) |
| x0 = relu(conv_2d_1_2(x0)) |
| x1 = pooling_2(x0) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2(x1)) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2_2(x1)) |
| x2 = pooling_2(x1) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_5_2(x2)) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_5(x2)) |
| x2 = up_sampling_2(x2) |
| x = concatenate([x2, x1], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6_2(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6(x)) |
| x = up_sampling_2(x) |
| x = concatenate([x, x0], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_2(x)) |
| x = conv_2d_8(x) |
| return x |
|
|
|
|
| def unet3(x): |
| x0 = cube_padding_1(x) |
| x0 = relu(conv_2d_1(x0)) |
| x0 = cube_padding_1(x0) |
| x0 = relu(conv_2d_1_2(x0)) |
| x0 = cube_padding_1(x0) |
| x0 = relu(conv_2d_1_3(x0)) |
| x1 = pooling_2(x0) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2(x1)) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2_2(x1)) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2_3(x1)) |
| x2 = pooling_2(x1) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_5_3(x2)) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_5_2(x2)) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_5(x2)) |
| x2 = up_sampling_2(x2) |
| x = concatenate([x2, x1], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6_3(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6_2(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6(x)) |
| x = up_sampling_2(x) |
| x = concatenate([x, x0], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_2(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_3(x)) |
| x = conv_2d_8(x) |
| return x |
|
|
|
|
| def unet4(x): |
| x0 = cube_padding_1(x) |
| x0 = relu(conv_2d_1(x0)) |
| x0 = cube_padding_1(x0) |
| x0 = relu(conv_2d_1_2(x0)) |
| x1 = pooling_2(x0) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2(x1)) |
| x1 = cube_padding_1(x1) |
| x1 = relu(conv_2d_2_2(x1)) |
| x2 = pooling_2(x1) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_3_2(x2)) |
| x2 = cube_padding_1(x2) |
| x2 = relu(conv_2d_3(x2)) |
| x3 = pooling_2(x2) |
| x3 = cube_padding_1(x3) |
| x3 = relu(conv_2d_4_2(x3)) |
| x3 = cube_padding_1(x3) |
| x3 = relu(conv_2d_4(x3)) |
| x3 = up_sampling_2(x3) |
| x = concatenate([x3, x2], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_5_2(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_5(x)) |
| x = up_sampling_2(x) |
| x = concatenate([x, x1], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6_2(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_6(x)) |
| x = up_sampling_2(x) |
| x = concatenate([x, x0], axis=-1) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7(x)) |
| x = cube_padding_1(x) |
| x = relu(conv_2d_7_2(x)) |
| x = conv_2d_8(x) |
| return x |
|
|
|
|
| def complete_model(x_in): |
| outputs = [] |
| model_function = globals()[cnn_model_name] |
| is_seq = isinstance(x_in, (list, tuple)) |
| xi = x_in[0] if is_seq else x_in |
| if is_seq and has_constants: |
| xi = Concatenate(axis=-1)([xi, x_in[-1]]) |
| outputs.append(model_function(xi)) |
| for step in range(1, integration_steps): |
| xo = outputs[step - 1] |
| if is_seq and input_solar: |
| xo = Reshape(cs[:-1] + (io_time_steps, -1))(xo) |
| xo = Concatenate(axis=-1)([xo, Permute((2, 3, 4, 1, 5))(x_in[step])]) |
| xo = Reshape(cs)(xo) |
| if is_seq and has_constants: |
| xo = Concatenate(axis=-1)([xo, x_in[-1]]) |
| outputs.append(model_function(xo)) |
|
|
| return outputs |
|
|
|
|
| |
| if not input_solar and not has_constants: |
| inputs = main_input |
| else: |
| inputs = [main_input] |
| if input_solar: |
| inputs = inputs + solar_inputs |
| if has_constants: |
| inputs = inputs + [constant_input] |
| model = Model(inputs=inputs, outputs=complete_model(inputs)) |
|
|
| |
| loss_function = 'mse' |
| if loss_by_step is None: |
| loss_by_step = [1./integration_steps] * integration_steps |
|
|
| |
| opt = tf.train.experimental.enable_mixed_precision_graph_rewrite(Adam()) if use_mp_optimizer else Adam() |
| dlwp.build_model(model, loss=loss_function, loss_weights=loss_by_step, optimizer=opt, metrics=['mae'], gpus=n_gpu) |
| print(dlwp.base_model.summary()) |
|
|
|
|
| |
|
|
| |
| start_time = time.time() |
| print('Begin training...') |
| run = Run.get_context() |
| history = RunHistory(run) |
| early = EarlyStoppingMin(monitor='val_loss' if validation_data is not None else 'loss', min_delta=0., |
| min_epochs=min_epochs, max_epochs=max_epochs, patience=patience, |
| restore_best_weights=True, verbose=1) |
| tensorboard = TensorBoard(log_dir=log_directory, update_freq='epoch') |
| save = SaveWeightsOnEpoch(weights_file=model_file + '.keras.tmp', interval=25) |
|
|
| try: |
| dlwp.model.load_weights('%s.keras.tmp' % model_file) |
| print('Loaded weights from existing model temporary file') |
| except: |
| pass |
|
|
| dlwp.fit_generator(tf_train_data, epochs=max_epochs + 1, |
| verbose=2, validation_data=tf_val_data, |
| callbacks=[history, RNNResetStates(), early, save, GeneratorEpochEnd(generator)]) |
| end_time = time.time() |
|
|
| |
| if model_file is not None: |
| save_model(dlwp, model_file, history=history) |
| print('Wrote model %s' % model_file) |
|
|
| |
| print("\nTrain time: %d m %s s" % (np.floor(total_time / 60), total_time % 60)) |
| try: |
| print('Train loss:', history.history['loss'][-patience - 1]) |
| run.log('TRAIN_LOSS', history.history['loss'][-patience - 1]) |
| print('Train mean absolute error:', history.history['mean_absolute_error'][-patience - 1]) |
| except (KeyError, IndexError): |
| pass |
| if validation_data is not None: |
| score = dlwp.evaluate(*val_generator.generate([]), verbose=0) |
| print('Validation loss:', score[0]) |
| try: |
| print('Validation mean absolute error:', score[1]) |
| except: |
| pass |
| run.log('VAL_LOSS', score[0]) |
|
|