| |
| |
| |
| |
| |
|
|
| """ |
| Example of training a DLWP model using a dataset of predictors generated with DLWP.model.Preprocessor. |
| |
| Uses Microsoft Azure resources. Launch this script as an Azure experiment using 'Train on Azure.ipynb' |
| """ |
|
|
| 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 DLWPNeuralNet, SeriesDataGenerator |
| from DLWP.util import save_model, train_test_split_ind |
| from DLWP.custom import RNNResetStates, EarlyStoppingMin, latitude_weighted_loss, RunHistory, anomaly_correlation_loss |
|
|
| from tensorflow.keras.losses import mean_squared_error |
| from tensorflow.keras.callbacks import TensorBoard |
| from azureml.core import Run |
| import tensorflow as tf |
|
|
|
|
| |
|
|
| 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) |
|
|
| |
| |
| model_is_convolutional = True |
| model_is_recurrent = False |
| min_epochs = 200 |
| max_epochs = 1000 |
| patience = 50 |
| batch_size = 64 |
| lambda_ = 1.e-4 |
| weight_loss = False |
| acc_loss = False |
| shuffle = True |
|
|
| |
| |
| |
| input_selection = {'varlev': ['HGT/500', 'THICK/300-700']} |
| output_selection = {'varlev': ['HGT/500', 'THICK/300-700']} |
| input_time_steps = 1 |
| output_time_steps = 1 |
| step_interval = 6 |
| |
| crop_north_pole = True |
| |
| add_solar = False |
|
|
| |
| |
| load_memory = True |
|
|
| |
| n_gpu = 1 |
|
|
| |
| |
| use_keras_fit = False |
|
|
| |
| |
| |
| validation_set = list(pd.date_range(datetime(2003, 1, 1, 0), datetime(2006, 12, 31, 18), freq='6H')) |
| train_set = list(pd.date_range(datetime(1979, 1, 1, 6), datetime(2002, 12, 31, 18), freq='6H')) |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| |
|
|
| if args.temp_dir != 'None': |
| new_predictor_file = os.path.join(args.temp_dir, args.predictor_file) |
| 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}) |
| else: |
| data = xr.open_dataset(predictor_file, chunks={'sample': batch_size}) |
|
|
| if 'time_step' in data.dims: |
| time_dim = data.dims['time_step'] |
| else: |
| time_dim = 1 |
| n_sample = data.dims['sample'] |
|
|
| if crop_north_pole: |
| data = data.isel(lat=(data.lat < 90.0)) |
|
|
|
|
| |
|
|
| dlwp = DLWPNeuralNet(is_convolutional=model_is_convolutional, is_recurrent=model_is_recurrent, time_dim=time_dim, |
| scaler_type=None, scale_targets=False) |
|
|
| |
| if isinstance(validation_set, int): |
| n_sample = data.dims['sample'] |
| ts, val_set = train_test_split_ind(n_sample, validation_set, method='last') |
| if train_set is None: |
| train_set = ts |
| elif isinstance(train_set, int): |
| train_set = list(range(train_set)) |
| validation_data = data.isel(sample=val_set) |
| train_data = data.isel(sample=train_set) |
| elif validation_set is None: |
| if train_set is None: |
| train_set = data.sample.values |
| validation_data = None |
| train_data = data.sel(sample=train_set) |
| else: |
| 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) |
|
|
| |
| batch_size = n_gpu * batch_size |
|
|
| |
| if load_memory or use_keras_fit: |
| print('Loading data to memory...') |
| generator = SeriesDataGenerator(dlwp, train_data, input_sel=input_selection, output_sel=output_selection, |
| input_time_steps=input_time_steps, output_time_steps=output_time_steps, |
| batch_size=batch_size, add_insolation=add_solar, load=load_memory, shuffle=shuffle, |
| interval=step_interval) |
| if use_keras_fit: |
| p_train, t_train = generator.generate([]) |
| if validation_data is not None: |
| val_generator = SeriesDataGenerator(dlwp, validation_data, input_sel=input_selection, output_sel=output_selection, |
| input_time_steps=input_time_steps, output_time_steps=output_time_steps, |
| batch_size=batch_size, add_insolation=add_solar, load=load_memory, |
| interval=step_interval) |
| if use_keras_fit: |
| val = val_generator.generate([]) |
| else: |
| val_generator = None |
| if use_keras_fit: |
| val = None |
|
|
|
|
| |
|
|
| |
| cs = generator.convolution_shape |
| cso = generator.output_convolution_shape |
| layers = ( |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| ('PeriodicPadding2D', ((0, 2),), { |
| 'data_format': 'channels_first', |
| 'input_shape': cs |
| }), |
| ('ZeroPadding2D', ((2, 0),), {'data_format': 'channels_first'}), |
| ('Conv2D', (32, 3), { |
| 'dilation_rate': 2, |
| 'padding': 'valid', |
| 'activation': 'tanh', |
| 'data_format': 'channels_first' |
| }), |
| |
| ('MaxPooling2D', (2,), {'data_format': 'channels_first'}), |
| ('PeriodicPadding2D', ((0, 1),), {'data_format': 'channels_first'}), |
| ('ZeroPadding2D', ((1, 0),), {'data_format': 'channels_first'}), |
| ('Conv2D', (64, 3), { |
| 'dilation_rate': 1, |
| 'padding': 'valid', |
| 'activation': 'tanh', |
| 'data_format': 'channels_first' |
| }), |
| |
| ('MaxPooling2D', (2,), {'data_format': 'channels_first'}), |
| ('PeriodicPadding2D', ((0, 1),), {'data_format': 'channels_first'}), |
| ('ZeroPadding2D', ((1, 0),), {'data_format': 'channels_first'}), |
| ('Conv2D', (128, 3), { |
| 'dilation_rate': 1, |
| 'padding': 'valid', |
| 'activation': 'tanh', |
| 'data_format': 'channels_first' |
| }), |
| |
| ('UpSampling2D', (2,), {'data_format': 'channels_first'}), |
| ('PeriodicPadding2D', ((0, 1),), {'data_format': 'channels_first'}), |
| ('ZeroPadding2D', ((1, 0),), {'data_format': 'channels_first'}), |
| ('Conv2D', (64, 3), { |
| 'dilation_rate': 1, |
| 'padding': 'valid', |
| 'activation': 'tanh', |
| 'data_format': 'channels_first' |
| }), |
| |
| ('UpSampling2D', (2,), {'data_format': 'channels_first'}), |
| ('PeriodicPadding2D', ((0, 2),), {'data_format': 'channels_first'}), |
| ('ZeroPadding2D', ((2, 0),), {'data_format': 'channels_first'}), |
| ('Conv2D', (32, 3), { |
| 'dilation_rate': 2, |
| 'padding': 'valid', |
| 'activation': 'tanh', |
| 'data_format': 'channels_first' |
| }), |
| |
| ('PeriodicPadding2D', ((0, 2),), {'data_format': 'channels_first'}), |
| ('ZeroPadding2D', ((2, 0),), {'data_format': 'channels_first'}), |
| |
| ('Conv2D', (cso[0], 5), { |
| 'padding': 'valid', |
| 'activation': 'linear', |
| 'data_format': 'channels_first' |
| }), |
| |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| if acc_loss: |
| |
| |
| print('Finding climatology for ACC loss...') |
| p_fit, t_fit = generator.generate([], scale_and_impute=False) |
| climo = t_fit.mean(axis=0, keepdims=True) |
| p_fit, t_fit = (None, None) |
| loss_function = anomaly_correlation_loss(climo, regularize_mean='mse', reverse=True) |
| else: |
| loss_function = mean_squared_error |
| if weight_loss: |
| loss_function = latitude_weighted_loss(loss_function, generator.ds.lat.values, generator.convolution_shape, |
| axis=-2, weighting='midlatitude') |
|
|
| |
| try: |
| dlwp.build_model(layers, loss=loss_function, optimizer='adam', metrics=['mae'], gpus=n_gpu) |
| except (ValueError, IndexError): |
| for layer in dlwp.base_model.layers: |
| print(layer.name, layer.output_shape) |
| raise |
| print(dlwp.base_model.summary()) |
|
|
|
|
| |
|
|
| |
| start_time = time.time() |
| print('Begin training...') |
| run = Run.get_context() |
| history = RunHistory(run) |
| early = EarlyStoppingMin(min_epochs=min_epochs, monitor='val_loss' if val_generator is not None else 'loss', |
| min_delta=0., patience=patience, restore_best_weights=True, verbose=1) |
| tensorboard = TensorBoard(log_dir=log_directory, batch_size=batch_size, update_freq='epoch') |
|
|
| if use_keras_fit: |
| dlwp.fit(p_train, t_train, batch_size=batch_size, epochs=max_epochs, verbose=2, validation_data=val, |
| shuffle=shuffle, callbacks=[history, RNNResetStates(), early]) |
| else: |
| dlwp.fit_generator(generator, epochs=max_epochs, verbose=2, validation_data=val_generator, |
| use_multiprocessing=True, callbacks=[history, RNNResetStates(), early]) |
| end_time = time.time() |
|
|
| |
| if model_file is not None: |
| os.makedirs(os.path.sep.join(model_file.split(os.path.sep)[:-1]), exist_ok=True) |
| save_model(dlwp, model_file, history=history) |
| print('Wrote model %s' % model_file) |
|
|
| |
| print("\nTrain time -- %s seconds --" % (end_time - start_time)) |
| 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]) |
|
|