Skip to content

Commit 70f4f55

Browse files
authored
Merge pull request #182 from metno/add-logging-timing
Add logging and timing
2 parents b28a014 + 34a0c40 commit 70f4f55

12 files changed

Lines changed: 142 additions & 49 deletions

File tree

bris/__main__.py

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import logging
22
import os
3+
import time
34
from datetime import datetime, timedelta
45

56
from anemoi.utils.dates import frequency_to_seconds
@@ -11,20 +12,22 @@
1112
from .checkpoint import Checkpoint
1213
from .inference import Inference
1314
from .utils import (
15+
LOGGER,
1416
create_config,
1517
get_all_leadtimes,
1618
parse_args,
1719
set_base_seed,
1820
set_encoder_decoder_num_chunks,
21+
setup_logging,
1922
)
2023
from .writer import CustomWriter
2124

22-
LOGGER = logging.getLogger(__name__)
23-
2425

2526
def main(arg_list: list[str] | None = None):
27+
t0 = time.perf_counter()
2628
args = parse_args(arg_list)
2729
config = create_config(args["config"], args)
30+
setup_logging(config)
2831

2932
models = list(config.checkpoints.keys())
3033
checkpoints = {
@@ -88,8 +91,8 @@ def main(arg_list: list[str] | None = None):
8891
),
8992
"%Y-%m-%dT%H:%M:%S",
9093
)
91-
LOGGER.info(
92-
"No start_date given, setting %s based on start_date and timestep.",
94+
LOGGER.warning(
95+
"No start_date given, setting %s based on end_date and timestep.",
9396
config.start_date,
9497
)
9598
else:
@@ -143,13 +146,6 @@ def main(arg_list: list[str] | None = None):
143146
)
144147
writer = CustomWriter(decoder_outputs, write_interval="batch")
145148

146-
# Set hydra defaults
147-
config.defaults = [
148-
{"override hydra/job_logging": "none"}, # disable config parsing logs
149-
{"override hydra/hydra_logging": "none"}, # disable config parsing logs
150-
"_self_",
151-
]
152-
153149
# Forecaster must know about what leadtimes to output
154150
model = instantiate(
155151
config.model,
@@ -180,9 +176,13 @@ def main(arg_list: list[str] | None = None):
180176
if is_main_thread:
181177
for decoder_output in decoder_outputs:
182178
for output in decoder_output["outputs"]:
179+
t1 = time.perf_counter()
183180
output.finalize()
181+
LOGGER.debug(
182+
f"finalizing decoder {decoder_output} output {output.filename_pattern} in {time.perf_counter() - t1:.1f}s"
183+
)
184184

185-
print("Model run completed. 🤖")
185+
LOGGER.info(f"Bris completed in {time.perf_counter() - t0:.1f}s. 🤖")
186186

187187

188188
if __name__ == "__main__":

bris/conventions/metno.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66
Additionally, the names of some dimension-variables do not use CF-names
77
"""
88

9+
from bris.utils import LOGGER
10+
911

1012
class Metno:
1113
cf_to_metno = {
@@ -46,12 +48,14 @@ def get_ncname(self, cfname: str, leveltype: str, level: int):
4648
# This is likely a forcing variable
4749
return cfname
4850
else:
49-
print(cfname, leveltype, level)
51+
LOGGER.error(
52+
f"get_ncname not implemented cfname {cfname}, leveltype {leveltype}, level {level}"
53+
)
5054
raise NotImplementedError()
5155

5256
return ncname
5357

54-
def is_single_level(self, cfname: str, leveltype: str) -> str:
58+
def is_single_level(self, cfname: str, leveltype: str) -> bool:
5559
"""Returns true if there should only be a single level in the level dimension for this
5660
variable.
5761
@@ -69,13 +73,13 @@ def is_single_level(self, cfname: str, leveltype: str) -> str:
6973
"wind_speed",
7074
] and leveltype in ["height"]
7175

72-
def get_name(self, cfname: str):
76+
def get_name(self, cfname: str) -> str:
7377
"""Get MetNorway's dimension name from cf standard name"""
7478
if cfname in self.cf_to_metno:
7579
return self.cf_to_metno[cfname]
7680
return cfname
7781

78-
def get_cfname(self, ncname):
82+
def get_cfname(self, ncname) -> str:
7983
"""Get the CF-standard name from a given MetNo name"""
8084
for k, v in self.cf_to_metno.items():
8185
if v == ncname:

bris/data/dataset/nativegrid.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -10,9 +10,7 @@
1010
from torch.utils.data import IterableDataset
1111

1212
from bris.data.grid_indices import BaseGridIndices
13-
from bris.utils import get_base_seed, get_usable_indices
14-
15-
LOGGER = logging.getLogger(__name__)
13+
from bris.utils import LOGGER, get_base_seed, get_usable_indices
1614

1715

1816
class NativeGridDataset(IterableDataset):
@@ -200,8 +198,8 @@ def per_worker_init(self, n_workers: int, worker_id: int) -> None:
200198
f"num_data_parallel = num_nodes * num_gpus_per_node / num_gpus_per_model"
201199
)
202200
if len(self.valid_date_indices) % self.ens_comm_num_groups != 0:
203-
print(
204-
f"Warning: Dataloader has {len(self.valid_date_indices)} samples, which is not divisible by "
201+
LOGGER.warning(
202+
f"Dataloader has {len(self.valid_date_indices)} samples, which is not divisible by "
205203
f"{self.ens_comm_num_groups} data parallel workers. This will lead to "
206204
f"{len(self.valid_date_indices) % self.ens_comm_num_groups} unprocessed samples.",
207205
"num_data_parallel = num_nodes * num_gpus_per_node / num_gpus_per_model",

bris/ddp_strategy.py

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -18,14 +18,15 @@
1818
from pytorch_lightning.strategies.ddp import DDPStrategy
1919
from pytorch_lightning.trainer.states import TrainerFn
2020

21-
from bris.utils import get_base_seed
22-
23-
LOGGER = logging.getLogger(__name__)
21+
from bris.utils import LOGGER, get_base_seed
2422

2523

2624
class DDPGroupStrategy(DDPStrategy):
2725
"""Distributed Data Parallel strategy with group communication."""
2826

27+
# Define type of model, set in DDPStrategy somewhere
28+
model: pl.LightningModule
29+
2930
def __init__(
3031
self,
3132
num_gpus_per_model: int,
@@ -89,6 +90,13 @@ def setup(self, trainer: pl.Trainer) -> None:
8990

9091
# set up reader groups by further splitting model_comm_group_ranks with read_group_size:
9192

93+
LOGGER.debug(
94+
"world_size %d, model_comm_group_size %d, read_group_size %d",
95+
self.world_size,
96+
self.model_comm_group_size,
97+
self.read_group_size,
98+
)
99+
92100
assert self.model_comm_group_size % self.read_group_size == 0, (
93101
f"Number of GPUs per model ({self.model_comm_group_size}) must be divisible by read_group_size "
94102
f"({self.read_group_size})."
@@ -226,7 +234,7 @@ def get_my_model_comm_group(self, num_gpus_per_model: int) -> tuple[int, int, in
226234

227235
def get_my_reader_group(
228236
self, model_comm_group_rank: int, read_group_size: int
229-
) -> tuple[int, int, int]:
237+
) -> tuple[int, int, int, int]:
230238
"""Determine tasks that work together and from a reader group.
231239
232240
Parameters

bris/inference.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import logging
2+
import time
23
from functools import cached_property
34
from typing import Any, Optional
45

@@ -7,11 +8,10 @@
78
from anemoi.utils.config import DotDict
89

910
from bris.ddp_strategy import DDPGroupStrategy
11+
from bris.utils import LOGGER
1012

1113
from .data.datamodule import DataModule
1214

13-
LOGGER = logging.getLogger(__name__)
14-
1515

1616
class Inference:
1717
def __init__(
@@ -41,7 +41,7 @@ def device(self) -> str:
4141
LOGGER.info("Specified device not set. Found GPU")
4242
return "cuda"
4343

44-
LOGGER.info("Specified device not set. Could not find gpu, using CPU")
44+
LOGGER.warning("Specified device not set. Could not find gpu, using CPU")
4545
return "cpu"
4646

4747
LOGGER.info("Using specified device: %s", self._device)
@@ -74,6 +74,9 @@ def trainer(self) -> pl.Trainer:
7474
return trainer
7575

7676
def run(self):
77+
t0 = time.perf_counter()
78+
LOGGER.debug("Bris/Inference/run Predicting")
7779
self.trainer.predict(
7880
self.model, datamodule=self.datamodule, return_predictions=False
7981
)
82+
LOGGER.debug(f"bris/Inference.run: {time.perf_counter() - t0:.1f}s")

bris/model/brispredictor.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,14 +18,13 @@
1818
get_dynamic_forcings,
1919
)
2020
from ..utils import (
21+
LOGGER,
2122
check_anemoi_training,
2223
timedelta64_from_timestep,
2324
)
2425
from .basepredictor import BasePredictor
2526
from .model_utils import get_model_static_forcings, get_variable_indices
2627

27-
LOGGER = logging.getLogger(__name__)
28-
2928

3029
class BrisPredictor(BasePredictor):
3130
"""

bris/outputs/__init__.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
import copy
2+
import time
23
from typing import Optional
34

45
import numpy as np
56

67
from bris import sources
78
from bris.predict_metadata import PredictMetadata
9+
from bris.utils import LOGGER
810

911

1012
def instantiate(name: str, predict_metadata: PredictMetadata, workdir: str, init_args):
@@ -113,6 +115,7 @@ def add_forecast(self, times: list, ensemble_member: int, pred: np.ndarray):
113115
self.pm.variables.append(name)
114116

115117
# only do this once. For multiple members, intermediate calls this several times
118+
t0 = time.perf_counter()
116119
if pred.shape[2] != len(self.pm.variables):
117120
# Append extra variables to prediction
118121
extra_pred = []
@@ -126,6 +129,9 @@ def add_forecast(self, times: list, ensemble_member: int, pred: np.ndarray):
126129
raise ValueError(f"No recipe to compute {name}")
127130

128131
pred = np.concatenate([pred] + extra_pred, axis=2)
132+
LOGGER.debug(
133+
f"outputs.add_forecast Calculate ws in {time.perf_counter() - t0:.1f}s"
134+
)
129135

130136
assert pred.shape[0] == self.pm.num_leadtimes
131137
assert pred.shape[1] == len(self.pm.lats)
@@ -136,7 +142,11 @@ def add_forecast(self, times: list, ensemble_member: int, pred: np.ndarray):
136142
assert ensemble_member >= 0
137143
assert ensemble_member < self.pm.num_members
138144

145+
t1 = time.perf_counter()
139146
self._add_forecast(times, ensemble_member, pred)
147+
LOGGER.debug(
148+
f"outputs.add_forecast called _add_forecast in {time.perf_counter() - t1:.1f}s"
149+
)
140150

141151
def _add_forecast(self, times: list, ensemble_member: int, pred: np.ndarray):
142152
"""Subclasses should implement this"""

bris/outputs/intermediate.py

Lines changed: 25 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import glob
22
import os
3+
import time
34
from typing import Optional
45

56
import numpy as np
@@ -24,17 +25,21 @@ def __init__(
2425
super().__init__(predict_metadata, extra_variables)
2526
self.workdir = workdir
2627

27-
def _add_forecast(self, times, ensemble_member, pred):
28+
def _add_forecast(self, times, ensemble_member, pred) -> None:
29+
t0 = time.perf_counter()
2830
filename = self.get_filename(times[0], ensemble_member)
2931
utils.create_directory(filename)
3032

3133
np.save(filename, pred)
34+
utils.LOGGER.debug(
35+
f"Intermediate._add_forecast for {filename} in {time.perf_counter() - t0:.1f}s"
36+
)
3237

33-
def get_filename(self, forecast_reference_time, ensemble_member):
38+
def get_filename(self, forecast_reference_time, ensemble_member) -> str:
3439
frt_ut = utils.datetime_to_unixtime(forecast_reference_time)
3540
return f"{self.workdir}/{frt_ut:.0f}_{ensemble_member:.0f}.npy"
3641

37-
def get_forecast_reference_times(self):
42+
def get_forecast_reference_times(self) -> list[np.datetime64]:
3843
"""Returns all forecast reference times that have been saved"""
3944
filenames = self.get_filenames()
4045
frts = []
@@ -48,7 +53,9 @@ def get_forecast_reference_times(self):
4853

4954
return frts
5055

51-
def get_forecast(self, forecast_reference_time, ensemble_member=None):
56+
def get_forecast(
57+
self, forecast_reference_time, ensemble_member=None
58+
) -> np.ndarray | None:
5259
"""Fetches forecasts from stored numpy files
5360
5461
Args:
@@ -61,6 +68,7 @@ def get_forecast(self, forecast_reference_time, ensemble_member=None):
6168
4D otherwise (leadtime, points, variables, members)
6269
"""
6370

71+
t0 = time.perf_counter()
6472
if ensemble_member is None:
6573
shape = [
6674
self.pm.num_leadtimes,
@@ -73,16 +81,21 @@ def get_forecast(self, forecast_reference_time, ensemble_member=None):
7381
filename = self.get_filename(forecast_reference_time, e)
7482
if os.path.exists(filename):
7583
pred[..., e] = np.load(filename)
84+
utils.LOGGER.debug(
85+
f"Intermediate.get_forecast for {filename} in {time.perf_counter() - t0:.1f}s"
86+
)
7687
else:
7788
assert isinstance(ensemble_member, int)
7889

7990
filename = self.get_filename(forecast_reference_time, ensemble_member)
8091
pred = np.load(filename) if os.path.exists(filename) else None
81-
92+
utils.LOGGER.debug(
93+
f"Intermediate.get_forecast for {filename} in {time.perf_counter() - t0:.1f}s"
94+
)
8295
return pred
8396

8497
@property
85-
def num_members(self):
98+
def num_members(self) -> int:
8699
filenames = self.get_filenames()
87100

88101
max_member = 0
@@ -91,22 +104,23 @@ def num_members(self):
91104
max_member = max(int(member), max_member)
92105
return max_member + 1
93106

94-
def get_filenames(self):
107+
def get_filenames(self) -> list[str]:
95108
return glob.glob(f"{self.workdir}/*_*.npy")
96109

97-
def cleanup(self):
110+
def cleanup(self) -> None:
98111
"""Removes up all intermediate files and removes the workdir. Called in finalize of the main output."""
99-
112+
t0 = time.perf_counter()
100113
for _filename in self.get_filenames():
101114
try:
102115
os.remove(_filename)
103116
except OSError as e:
104-
print(f"Error during cleanup of {_filename}: {e}")
117+
utils.LOGGER.warning(f"Error during cleanup of {_filename}: {e}")
105118

106119
try:
107120
os.rmdir(self.workdir)
108121
except OSError as e:
109-
print(f"Error removing workdir {self.workdir}: {e}")
122+
utils.LOGGER.warning(f"Error removing workdir {self.workdir}: {e}")
123+
utils.LOGGER.debug(f"Intermediate.cleanup in {time.perf_counter() - t0:.1f}s")
110124

111125
def finalize(self):
112126
pass

0 commit comments

Comments
 (0)