88from 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+
98198def 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
123223if __name__ == "__main__" :
124224 test_domain_name ()
125- test_1 ()
225+ test_deterministic ()
226+ test_ensemble ()
0 commit comments