Skip to content

Commit 0565fd3

Browse files
committed
Change dimension order in NetCDF output
Put ensemble dimension after height dimension (if height exists)
1 parent 24ceea4 commit 0565fd3

2 files changed

Lines changed: 111 additions & 8 deletions

File tree

bris/outputs/netcdf.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -479,8 +479,12 @@ def _setup_prediction_vars(
479479
shape = [len(times), len(y)]
480480

481481
if self.pm.num_members > 1:
482-
dims.insert(1, self.conv_name("ensemble_member"))
483-
shape.insert(1, self.pm.num_members)
482+
# We want the ensemble member dimension to be after height (if height exists)
483+
# I.e. we want (time, height, member, y, x) or (time, member, y, x)
484+
ens_dim_loc = 1 + (dim_name is not None)
485+
486+
dims.insert(ens_dim_loc, self.conv_name("ensemble_member"))
487+
shape.insert(ens_dim_loc, self.pm.num_members)
484488

485489
ar = np.nan * np.zeros(shape, np.float32)
486490
self.ds[ncname] = (dims, ar)
@@ -515,6 +519,7 @@ def _setup_prediction_vars(
515519
# Accumulate over lead times
516520
ar = np.cumsum(np.nan_to_num(ar, nan=0), axis=0)
517521

522+
# Rearrange to (time, member, y, [x]) or (time, y, [x])
518523
ar = np.moveaxis(ar, [-1], [1]) if self.pm.num_members > 1 else ar[..., 0]
519524

520525
cfname = cf.get_metadata(variable)["cfname"]
@@ -527,10 +532,7 @@ def _setup_prediction_vars(
527532
bris.units.convert(ar, from_units, to_units, inplace=True)
528533

529534
if level_index is not None:
530-
if self.pm.num_members > 1:
531-
self.ds[ncname][:, :, level_index, ...] = ar
532-
else:
533-
self.ds[ncname][:, level_index, ...] = ar
535+
self.ds[ncname][:, level_index, ...] = ar
534536
else:
535537
self.ds[ncname][:] = ar
536538

tests/test_outputs_netcdf.py

Lines changed: 103 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
from bris.predict_metadata import PredictMetadata
99

1010

11-
def test_1():
11+
def test_deterministic():
1212
variables = ["u_800", "u_600", "2t", "v_500", "10u", "tp", "skt"]
1313
lats = np.array([1, 2])
1414
lons = np.array([2, 4])
@@ -25,6 +25,7 @@ def test_1():
2525
times = frt + leadtimes
2626

2727
with tempfile.TemporaryDirectory() as temp_dir:
28+
temp_dir = "./"
2829
pattern = os.path.join(temp_dir, "test_%Y%m%dT%HZ.nc")
2930
workdir = os.path.join(temp_dir, "test_gridded")
3031
attrs = {"creator": "met.no"}
@@ -61,6 +62,9 @@ def test_1():
6162
assert "units" in var.attrs
6263
assert "grid_mapping" in var.attrs
6364

65+
# (time, height, y, x)
66+
assert file.variables["air_temperature_2m"].shape == (4, 1, 1, 2)
67+
6468
height_dim = file.variables["air_temperature_2m"].dims[1]
6569
levels = file.variables[height_dim].values
6670
assert levels == [2]
@@ -95,6 +99,102 @@ def test_1():
9599
assert np.min(file.variables["altitude"]) == 100
96100

97101

102+
def test_ensemble():
103+
variables = [
104+
"u_500",
105+
"u_800",
106+
"u_600",
107+
"2t",
108+
"v_500",
109+
"10u",
110+
"tp",
111+
"skt",
112+
"unknown",
113+
]
114+
lats = np.array([1, 2])
115+
lons = np.array([2, 4])
116+
altitudes = np.array([100, 200])
117+
leadtimes = np.arange(0, 3600 * 4, 3600)
118+
num_members = 2
119+
field_shape = [1, 2]
120+
pm = PredictMetadata(
121+
variables, lats, lons, altitudes, leadtimes, num_members, field_shape
122+
)
123+
124+
pred = np.random.rand(*pm.shape)
125+
frt = 1672552800
126+
times = frt + leadtimes
127+
128+
with tempfile.TemporaryDirectory() as temp_dir:
129+
temp_dir = "./"
130+
pattern = os.path.join(temp_dir, "test_%Y%m%dT%HZ.nc")
131+
workdir = os.path.join(temp_dir, "test_gridded")
132+
attrs = {"creator": "met.no"}
133+
output = Netcdf(pm, workdir, pattern, global_attributes=attrs)
134+
135+
for member in range(num_members):
136+
output.add_forecast(times, member, pred)
137+
output.finalize()
138+
139+
output_filename = os.path.join(temp_dir, "test_20230101T06Z.nc")
140+
141+
assert os.path.exists(output_filename)
142+
143+
with xr.open_dataset(output_filename) as file:
144+
# Check that global attributes are written
145+
for k, v in attrs.items():
146+
assert file.attrs[k] == v
147+
148+
for variable in [
149+
"altitude",
150+
"air_temperature_2m",
151+
"air_temperature_0m",
152+
"x_wind_pl",
153+
]:
154+
assert variable in file.variables, variable
155+
var = file.variables[variable]
156+
assert "units" in var.attrs
157+
assert "grid_mapping" in var.attrs
158+
159+
# (time, height, member, y, x)
160+
assert file.variables["air_temperature_2m"].shape == (4, 1, 2, 1, 2)
161+
162+
height_dim = file.variables["air_temperature_2m"].dims[1]
163+
levels = file.variables[height_dim].values
164+
assert levels == [2]
165+
166+
height_dim = file.variables["air_temperature_0m"].dims[1]
167+
levels = file.variables[height_dim].values
168+
assert levels == [0]
169+
170+
# (time, pressure, member, y, x)
171+
assert file.variables["x_wind_pl"].shape == (4, 3, 2, 1, 2)
172+
173+
height_dim = file.variables["x_wind_pl"].dims[1]
174+
levels = file.variables[height_dim].values
175+
assert len(levels) == 3
176+
assert levels[0] == 500
177+
assert levels[1] == 600
178+
assert levels[2] == 800
179+
assert file.variables["x_wind_pl"].values.shape[2] == num_members
180+
181+
# Test interpolation
182+
with tempfile.TemporaryDirectory() as temp_dir:
183+
pattern = os.path.join(temp_dir, "test_%Y%m%dT%HZ.nc")
184+
workdir = os.path.join(temp_dir, "test_gridded")
185+
attrs = {"creator": "met.no"}
186+
output = Netcdf(pm, workdir, pattern, interp_res=0.2)
187+
188+
for member in range(num_members):
189+
output.add_forecast(times, member, pred)
190+
output.finalize()
191+
192+
output_filename = os.path.join(temp_dir, "test_20230101T06Z.nc")
193+
with xr.open_dataset(output_filename) as file:
194+
# Check that altitude variable has attributes
195+
assert "altitude" not in file.variables
196+
197+
98198
def test_domain_name():
99199
variables = ["u_800", "u_600", "2t", "v_500", "10u"]
100200
lats = np.array([1, 2])
@@ -122,4 +222,5 @@ def test_domain_name():
122222

123223
if __name__ == "__main__":
124224
test_domain_name()
125-
test_1()
225+
test_deterministic()
226+
test_ensemble()

0 commit comments

Comments
 (0)