# # Copyright (c) 2019 Jonathan Weyn # # See the file LICENSE for your rights. # """ 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 # Disable warning logging tf.compat.v1.logging.set_verbosity(tf.compat.v1.logging.ERROR) # # Set only GPU 0 # device = tf.config.list_physical_devices('GPU')[0] # tf.config.set_visible_devices([device], 'GPU') # # Allow memory growth # tf.config.experimental.set_memory_growth(device, True) #%% Parse user arguments 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) #%% Parameters # File paths and names 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 # Optional paths to files containing constant fields to add to the inputs 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') ] # NN parameters. Regularization is applied to LSTM layers by default. weight_loss indicates whether to weight the # loss function preferentially in the mid-latitudes. 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 # Data parameters. Specify the input/output variables/levels and input/output time steps. DLWPFunctional requires that # the inputs and outputs match exactly (for now). Ensure that the selections use LISTS of values (even for only 1) to # keep dimensions correct. The number of output iterations to train on is given by integration_steps. The actual number # of forecast steps (units of model delta t) is io_time_steps * integration_steps. The parameter data_interval # governs what the effective delta t is; it is a multiplier for the temporal resolution of the data file. io_selection = {'varlev': ['z/500', 'tau/300-700', 'z/1000', 't2m/0']} io_time_steps = 2 integration_steps = 2 data_interval = 2 # Add incoming solar radiation forcing add_solar = True # Use multiple GPUs, if available n_gpu = 1 # Optimize the optimizer for GPU tensor cores by using mixed precision use_mp_optimizer = True # Validation set to use. Either an integer (number of validation samples, taken from the end), or an iterable of # pandas datetime objects. The train set can be set to the first samples, an iterable of dates, or None to # simply use the remaining points. Match the type of validation_set. 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')) #%% Open data 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) #%% Create a model and the data generators dlwp = DLWPFunctional(is_convolutional=True, is_recurrent=False, time_dim=io_time_steps) # Find the validation set 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) # Build the data generators 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)) #%% Compile the model structure with some generator data information # Up-sampling convolutional network or U-net cs = generator.convolution_shape cso = generator.output_convolution_shape input_solar = (integration_steps > 1 and (isinstance(add_solar, str) or add_solar)) # Define layers. Must be defined outside of model function so we use the same weights at each integration step. 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) # Define the model functions. 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 # Build the model with inputs and 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)) # No weighted loss available for cube sphere at the moment, but we can weight each integration sequence loss_function = 'mse' if loss_by_step is None: loss_by_step = [1./integration_steps] * integration_steps # Build the DLWP model 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()) #%% Train, evaluate, and save the model # Train and evaluate the model 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() # Save the model if model_file is not None: save_model(dlwp, model_file, history=history) print('Wrote model %s' % model_file) # Evaluate the model 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])