1- #!/usr/bin/env python3
1+ #!/usr/bin/env -S python3 -u
22
33'''
44Description: Script to call the graphcast model using gdas products
99'''
1010import os
1111import argparse
12+ from time import time
1213from datetime import timedelta
1314import dataclasses
1415import functools
@@ -163,9 +164,13 @@ def construct_wrapped_graphcast(model_config, task_config):
163164 # from/to float32 to/from BFloat16.
164165 predictor = casting .Bfloat16Cast (predictor )
165166
166- # Modify inputs/outputs to `casting.Bfloat16Cast` so the casting to/from
167- # BFloat16 happens after applying normalization to the inputs/targets.
168- predictor = normalization .InputsAndResiduals (predictor , diffs_stddev_by_level = self .diffs_stddev_by_level , mean_by_level = self .mean_by_level , stddev_by_level = self .stddev_by_level ,)
167+ # Applying normalization to the inputs/targets.
168+ predictor = normalization .InputsAndResiduals (
169+ predictor ,
170+ diffs_stddev_by_level = self .diffs_stddev_by_level ,
171+ mean_by_level = self .mean_by_level ,
172+ stddev_by_level = self .stddev_by_level ,
173+ )
169174
170175 # Wraps everything so the one-step model can produce trajectories.
171176 predictor = autoregressive .Predictor (predictor , gradient_checkpointing = True ,)
@@ -176,8 +181,11 @@ def run_forward(model_config, task_config, inputs, targets_template, forcings,):
176181 predictor = construct_wrapped_graphcast (model_config , task_config )
177182 return predictor (inputs , targets_template = targets_template , forcings = forcings ,)
178183
184+ t0 = time ()
179185 jax .jit (self ._with_configs (run_forward .init ))
180186 self .model = self ._drop_state (self ._with_params (jax .jit (self ._with_configs (run_forward .apply ))))
187+ elapsed_time = time () - t0
188+ print (f"Elapsed time for compiling the model: { elapsed_time } seconds" )
181189
182190
183191 def get_predictions (self ):
@@ -280,11 +288,30 @@ def upload_to_s3(self, keep_data):
280288 args = parser .parse_args ()
281289 runner = GraphCastModel (args .weights , args .input , args .case_name , args .config , args .output , int (args .pressure ), int (args .length ))
282290
291+ t0 = time ()
283292 runner .load_pretrained_model ()
293+ elapsed_time = time () - t0
294+ print (f"Elapsed time for loading model: { elapsed_time } seconds" )
295+
296+ t0 = time ()
284297 runner .load_gdas_data ()
298+ elapsed_time = time () - t0
299+ print (f"Elapsed time for loading input data: { elapsed_time } seconds" )
300+
301+ t0 = time ()
285302 runner .extract_inputs_targets_forcings ()
303+ elapsed_time = time () - t0
304+ print (f"Elapsed time for extracting inputs, targets, and forcings: { elapsed_time } seconds" )
305+
306+ t0 = time ()
286307 runner .load_normalization_stats ()
308+ elapsed_time = time () - t0
309+ print (f"Elapsed time for loading normalization stats: { elapsed_time } seconds" )
310+
311+ t0 = time ()
287312 runner .get_predictions ()
313+ elapsed_time = time () - t0
314+ print (f"Elapsed time for running the model: { elapsed_time } seconds" )
288315
289316 upload_data = args .upload .lower () == "yes"
290317 keep_data = args .keep .lower () == "yes"
0 commit comments