Skip to content

Commit a03a127

Browse files
Added time log for each step and removed Bfloat16cast (#68)
This PR made two changes to `run_graphcast.py`: - Added time log for each step - Removed conversion to Bfloat16. Bfloat16 should not be used for prediction. The reduced precision greatly increases run-to-run variance. --------- Co-authored-by: Russell Manser <russell.manser@noaa.gov>
1 parent e74058d commit a03a127

1 file changed

Lines changed: 31 additions & 4 deletions

File tree

‎oper/run_graphcast.py‎

Lines changed: 31 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
#!/usr/bin/env python3
1+
#!/usr/bin/env -S python3 -u
22

33
'''
44
Description: Script to call the graphcast model using gdas products
@@ -9,6 +9,7 @@
99
'''
1010
import os
1111
import argparse
12+
from time import time
1213
from datetime import timedelta
1314
import dataclasses
1415
import 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

Comments
 (0)