diff --git a/weather_mv/loader_pipeline/bq.py b/weather_mv/loader_pipeline/bq.py index d1e1b5f..cce33a1 100644 --- a/weather_mv/loader_pipeline/bq.py +++ b/weather_mv/loader_pipeline/bq.py @@ -54,6 +54,8 @@ GEO_POLYGON_COLUMN = 'geo_polygon' LATITUDE_RANGE = (-90, 90) LONGITUDE_RANGE = (-180, 180) +LATITUDE_COORD_CANDIDATES: t.Tuple[str, ...] = ('latitude', 'lat', 'y') +LONGITUDE_COORD_CANDIDATES: t.Tuple[str, ...] = ('longitude', 'lon', 'x') @dataclasses.dataclass @@ -115,6 +117,9 @@ class ToBigQuery(ToDataSink): skip_creating_geo_data_parquet: bool = False lat_grid_resolution: t.Optional[float] = None lon_grid_resolution: t.Optional[float] = None + lat_coord_name: str = dataclasses.field(init=False, default='latitude') + lon_coord_name: str = dataclasses.field(init=False, default='longitude') + _coord_rename_map: t.Dict[str, str] = dataclasses.field(init=False, repr=False) @classmethod def add_parser_arguments(cls, subparser: argparse.ArgumentParser): @@ -202,6 +207,8 @@ def generate_parquet( lat_grid_resolution: float, lon_grid_resolution: float, skip_creating_polygon: bool = False, + lat_column_name: str = 'latitude', + lon_column_name: str = 'longitude', ): """Generates geo data parquet.""" logger.info("Generating geo data parquet ...") @@ -210,7 +217,7 @@ def generate_parquet( # Create a temp parquet file for writing. with tempfile.NamedTemporaryFile(suffix='.parquet', mode='w+', newline='') as temp: # Define header. - header = ['latitude', 'longitude', GEO_POINT_COLUMN, GEO_POLYGON_COLUMN] + header = [lat_column_name, lon_column_name, GEO_POINT_COLUMN, GEO_POLYGON_COLUMN] data = [] for lat, lon in lat_lon_pairs: lat = float(lat) @@ -244,17 +251,23 @@ def __post_init__(self): with open_dataset(self.first_uri, self.xarray_open_dataset_kwargs, self.disable_grib_schema_normalization, self.tif_metadata_for_start_time, self.tif_metadata_for_end_time, is_zarr=self.zarr) as open_ds: + source_lat, source_lon = self._resolve_spatial_coordinate_names(open_ds) + self._coord_rename_map = self._build_coord_rename_map(source_lat, source_lon) + open_ds = self._normalize_dataset_coords(open_ds) if not self.skip_creating_polygon: - logger.warning("Assumes that equal distance between consecutive points of latitude " - "and longitude for the entire grid.") + logger.warning( + "Assumes that equal distance between consecutive points of %s and %s for the entire grid.", + self.lat_coord_name, + self.lon_coord_name, + ) # Find the grid_resolution. - if open_ds['latitude'].size > 1 and open_ds['longitude'].size > 1: - latitude_length = len(open_ds['latitude']) - longitude_length = len(open_ds['longitude']) + if open_ds[self.lat_coord_name].size > 1 and open_ds[self.lon_coord_name].size > 1: + latitude_length = len(open_ds[self.lat_coord_name]) + longitude_length = len(open_ds[self.lon_coord_name]) - latitude_range = np.ptp(open_ds["latitude"].values) - longitude_range = np.ptp(open_ds["longitude"].values) + latitude_range = np.ptp(open_ds[self.lat_coord_name].values) + longitude_range = np.ptp(open_ds[self.lon_coord_name].values) self.lat_grid_resolution = abs(latitude_range / latitude_length) / 2 self.lon_grid_resolution = abs(longitude_range / longitude_length) / 2 @@ -268,10 +281,15 @@ def __post_init__(self): if not self.skip_creating_geo_data_parquet: if self.area: n, w, s, e = self.area - open_ds = open_ds.sel(latitude=slice(n, s), longitude=slice(w, e)) + open_ds = open_ds.sel( + { + self.lat_coord_name: slice(n, s), + self.lon_coord_name: slice(w, e), + } + ) - lats = open_ds["latitude"].values.tolist() - lons = open_ds["longitude"].values.tolist() + lats = open_ds[self.lat_coord_name].values.tolist() + lons = open_ds[self.lon_coord_name].values.tolist() self.generate_parquet( self.geo_data_parquet_path, [lats] if isinstance(lats, float) else lats, @@ -279,6 +297,8 @@ def __post_init__(self): self.lat_grid_resolution, self.lon_grid_resolution, self.skip_creating_polygon, + lat_column_name=self.lat_coord_name, + lon_column_name=self.lon_coord_name, ) else: logger.info("geo data parquet is not created as '--skip_creating_geo_data_parquet' flag passed.") @@ -287,7 +307,7 @@ def __post_init__(self): if self.variables and not self.infer_schema and not open_ds.attrs['is_normalized']: logger.info('Creating schema from input variables.') table_schema = to_table_schema( - [('latitude', 'FLOAT64'), ('longitude', 'FLOAT64'), ('time', 'TIMESTAMP')] + + [(self.lat_coord_name, 'FLOAT64'), (self.lon_coord_name, 'FLOAT64'), ('time', 'TIMESTAMP')] + [(var, 'FLOAT64') for var in self.variables] ) else: @@ -314,6 +334,7 @@ def prepare_coordinates(self, uri: str) -> t.Iterator[t.Tuple[str, t.Dict]]: with open_dataset(uri, self.xarray_open_dataset_kwargs, self.disable_grib_schema_normalization, self.tif_metadata_for_start_time, self.tif_metadata_for_end_time, is_zarr=self.zarr) as ds: + ds = self._normalize_dataset_coords(ds) data_ds: xr.Dataset = _only_target_vars(ds, self.variables) for coordinate in get_coordinates(data_ds, uri): yield uri, coordinate @@ -328,10 +349,16 @@ def extract_rows(self, uri: str, coordinate: t.Dict) -> t.Iterator[t.Dict]: with open_dataset(uri, self.xarray_open_dataset_kwargs, self.disable_grib_schema_normalization, self.tif_metadata_for_start_time, self.tif_metadata_for_end_time, is_zarr=self.zarr) as ds: + ds = self._normalize_dataset_coords(ds) data_ds: xr.Dataset = _only_target_vars(ds, self.variables) if self.area: n, w, s, e = self.area - data_ds = data_ds.sel(latitude=slice(n, s), longitude=slice(w, e)) + data_ds = data_ds.sel( + { + self.lat_coord_name: slice(n, s), + self.lon_coord_name: slice(w, e), + } + ) logger.info(f'Data filtered by area, size: {data_ds.nbytes}') yield from self.to_rows(coordinate, data_ds, uri) @@ -345,10 +372,12 @@ def to_rows(self, coordinate: t.Dict, ds: xr.Dataset, uri: str) -> t.Iterator[t. selected_ds = ds.loc[coordinate] # Ensure that the latitude and longitude dimensions are in sync with the geo data parquet. - if not BQ_EXCLUDE_COORDS - set(selected_ds.dims.keys()): - selected_ds = selected_ds.transpose('latitude', 'longitude') + coord_pair = {self.lat_coord_name, self.lon_coord_name} + if coord_pair.issubset(set(selected_ds.sizes.keys())): + selected_ds = selected_ds.transpose(self.lat_coord_name, self.lon_coord_name) vector_df = pd.read_parquet(master_lat_lon) + vector_df = self._align_vector_coordinate_columns(vector_df) if self.skip_creating_polygon: vector_df[GEO_POLYGON_COLUMN] = None @@ -358,7 +387,9 @@ def to_rows(self, coordinate: t.Dict, ds: xr.Dataset, uri: str) -> t.Iterator[t. # Add un-indexed coordinates. # Filter out excluded coordinates from coords. - filtered_coords = (c for c in selected_ds.coords if c not in BQ_EXCLUDE_COORDS) + filtered_coords = ( + c for c in selected_ds.coords if c not in {self.lat_coord_name, self.lon_coord_name} + ) for c in filtered_coords: if c not in coordinate and (not self.variables or c in self.variables): vector_df[c] = to_json_serializable_type(ensure_us_time_resolution(selected_ds[c].values)) @@ -391,9 +422,81 @@ def chunks_to_rows(self, _, ds: xr.Dataset) -> t.Iterator[t.Dict]: if not self.import_time or self.zarr: self.import_time = datetime.datetime.utcnow().replace(tzinfo=datetime.timezone.utc) + ds = self._normalize_dataset_coords(ds) + for coordinate in get_coordinates(ds, uri): yield from self.to_rows(coordinate, ds, uri) + def _build_coord_rename_map(self, source_lat: str, source_lon: str) -> t.Dict[str, str]: + rename_map: t.Dict[str, str] = {} + if source_lat and source_lat != self.lat_coord_name: + rename_map[source_lat] = self.lat_coord_name + if source_lon and source_lon != self.lon_coord_name: + rename_map[source_lon] = self.lon_coord_name + return rename_map + + def _normalize_dataset_coords(self, ds: xr.Dataset) -> xr.Dataset: + """Rename dataset coordinates to the internal latitude/longitude names.""" + if not getattr(self, "_coord_rename_map", None): + return ds + missing = [source for source in self._coord_rename_map if source not in ds.coords] + if missing: + raise ValueError( + f"Dataset is missing expected coordinate(s) {missing} required for normalization." + ) + return ds.rename(self._coord_rename_map) + + def _align_vector_coordinate_columns(self, vector_df: pd.DataFrame) -> pd.DataFrame: + """Ensures the geo parquet uses the same coordinate column names as the dataset.""" + + def _find_alias(candidates: t.Tuple[str, ...]) -> t.Optional[str]: + for candidate in candidates: + if candidate in vector_df.columns: + return candidate + return None + + rename_map: t.Dict[str, str] = {} + if self.lat_coord_name not in vector_df.columns: + source = _find_alias(LATITUDE_COORD_CANDIDATES) + if source: + rename_map[source] = self.lat_coord_name + if self.lon_coord_name not in vector_df.columns: + source = _find_alias(LONGITUDE_COORD_CANDIDATES) + if source: + rename_map[source] = self.lon_coord_name + + if rename_map: + vector_df = vector_df.rename(columns=rename_map) + + missing = [ + coord for coord in (self.lat_coord_name, self.lon_coord_name) if coord not in vector_df.columns + ] + if missing: + raise ValueError( + f"Geo data parquet {self.geo_data_parquet_path} is missing coordinate columns: {missing}" + ) + + return vector_df + + def _resolve_spatial_coordinate_names(self, ds: xr.Dataset) -> t.Tuple[str, str]: + """Detects the dataset's latitude and longitude coordinate names.""" + + def _find_candidate(candidates: t.Tuple[str, ...]) -> t.Optional[str]: + for candidate in candidates: + if candidate in ds.coords or candidate in ds.dims: + return candidate + return None + + lat_name = _find_candidate(LATITUDE_COORD_CANDIDATES) + lon_name = _find_candidate(LONGITUDE_COORD_CANDIDATES) + if not lat_name or not lon_name: + raise ValueError( + f"Unable to identify spatial coordinate names. " + f"Checked latitude aliases {LATITUDE_COORD_CANDIDATES} and longitude aliases {LONGITUDE_COORD_CANDIDATES}." + ) + + return lat_name, lon_name + def expand(self, paths): """Extract rows of variables from data paths into a BigQuery table.""" if not self.zarr: diff --git a/weather_mv/loader_pipeline/bq_test.py b/weather_mv/loader_pipeline/bq_test.py index 5fe1d9d..64eec6b 100644 --- a/weather_mv/loader_pipeline/bq_test.py +++ b/weather_mv/loader_pipeline/bq_test.py @@ -11,6 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import contextlib import datetime import json import logging @@ -18,6 +19,7 @@ import tempfile import typing as t import unittest +from unittest import mock import geojson import numpy as np @@ -416,6 +418,116 @@ def test_07_extract_rows_single_point(self): } self.assertRowsEqual(actual, expected) + def test_extract_rows_with_lat_lon_aliases(self): + lat = np.array([12.0]) + lon = np.array([34.0]) + temps = np.array([[[250.0]]], dtype=np.float32) + ds = xr.Dataset( + {"temp": (('time', 'lat', 'lon'), temps)}, + coords={ + 'time': np.array(['2018-01-02T06:00:00'], dtype='datetime64[ns]'), + 'lat': lat, + 'lon': lon, + }, + attrs={'is_normalized': False}, + ) + geo_df = pd.DataFrame({ + 'lat': [lat[0]], + 'lon': [lon[0]], + 'geo_point': [geojson.dumps(geojson.Point((lon[0], lat[0])))], + 'geo_polygon': [None], + }) + temp_parquet = tempfile.NamedTemporaryFile(suffix='.parquet', delete=False) + temp_parquet.close() + self.addCleanup(lambda: os.path.exists(temp_parquet.name) and os.remove(temp_parquet.name)) + geo_df.to_parquet(temp_parquet.name, index=False) + + @contextlib.contextmanager + def _fake_open_dataset(*args, **kwargs): + yield ds.copy(deep=True) + + @contextlib.contextmanager + def _fake_open_local(*args, **kwargs): + yield temp_parquet.name + + with mock.patch('weather_mv.loader_pipeline.bq.open_dataset', side_effect=_fake_open_dataset), \ + mock.patch('weather_mv.loader_pipeline.bq.open_local', side_effect=_fake_open_local): + actual = next( + self.extract( + self.test_data_path, + geo_data_parquet_path=temp_parquet.name, + skip_creating_polygon=True, + skip_creating_geo_data_parquet=True, + ) + ) + expected = { + 'temp': 250.0, + 'data_import_time': DEFAULT_IMPORT_TIME, + 'data_first_step': '2018-01-02T06:00:00+00:00', + 'data_uri': self.test_data_path, + 'latitude': lat[0], + 'longitude': lon[0], + 'time': '2018-01-02T06:00:00+00:00', + 'geo_point': geo_df.iloc[0]['geo_point'], + 'geo_polygon': None, + } + self.assertRowsEqual(actual, expected) + + def test_extract_rows_with_xy_coordinates_and_lat_lon_parquet(self): + lat = np.array([12.0]) + lon = np.array([34.0]) + temps = np.array([[[250.0]]], dtype=np.float32) + ds = xr.Dataset( + {"temp": (('time', 'y', 'x'), temps)}, + coords={ + 'time': np.array(['2018-01-02T06:00:00'], dtype='datetime64[ns]'), + 'y': lat, + 'x': lon, + }, + attrs={'is_normalized': False}, + ) + geo_df = pd.DataFrame({ + 'latitude': [lat[0]], + 'longitude': [lon[0]], + 'geo_point': [geojson.dumps(geojson.Point((lon[0], lat[0])))], + 'geo_polygon': [None], + }) + temp_parquet = tempfile.NamedTemporaryFile(suffix='.parquet', delete=False) + temp_parquet.close() + self.addCleanup(lambda: os.path.exists(temp_parquet.name) and os.remove(temp_parquet.name)) + geo_df.to_parquet(temp_parquet.name, index=False) + + @contextlib.contextmanager + def _fake_open_dataset(*args, **kwargs): + yield ds.copy(deep=True) + + @contextlib.contextmanager + def _fake_open_local(*args, **kwargs): + yield temp_parquet.name + + with mock.patch('weather_mv.loader_pipeline.bq.open_dataset', side_effect=_fake_open_dataset), \ + mock.patch('weather_mv.loader_pipeline.bq.open_local', side_effect=_fake_open_local): + actual = next( + self.extract( + self.test_data_path, + geo_data_parquet_path=temp_parquet.name, + skip_creating_polygon=True, + skip_creating_geo_data_parquet=True, + ) + ) + expected = { + 'temp': 250.0, + 'data_import_time': DEFAULT_IMPORT_TIME, + 'data_first_step': '2018-01-02T06:00:00+00:00', + 'data_uri': self.test_data_path, + 'latitude': lat[0], + 'longitude': lon[0], + 'time': '2018-01-02T06:00:00+00:00', + 'geo_point': geo_df.iloc[0]['geo_point'], + 'geo_polygon': None, + } + self.assertRowsEqual(actual, expected) + def test_08_extract_rows_nan(self): self.test_data_path = f'{self.test_data_folder}/test_data_has_nan.nc' self.geo_data_parquet_path = f'{self.test_data_folder}/test_data_has_nan_geo_data.parquet' diff --git a/weather_mv/loader_pipeline/pipeline.py b/weather_mv/loader_pipeline/pipeline.py index 8b11c74..d497db4 100644 --- a/weather_mv/loader_pipeline/pipeline.py +++ b/weather_mv/loader_pipeline/pipeline.py @@ -25,7 +25,6 @@ from .bq import ToBigQuery from .regrid import Regrid -from .ee import ToEarthEngine from .streaming import GroupMessagesByFixedWindows, ParsePaths logger = logging.getLogger(__name__) @@ -76,6 +75,7 @@ def pipeline(known_args: argparse.Namespace, pipeline_args: t.List[str]) -> None elif known_args.subcommand == 'regrid' or known_args.subcommand == 'rg': paths | "Regrid" >> Regrid.from_kwargs(**vars(known_args)) elif known_args.subcommand == 'earthengine' or known_args.subcommand == 'ee': + from .ee import ToEarthEngine pipeline_options = PipelineOptions(pipeline_args) pipeline_options_dict = pipeline_options.get_all_options() # all_args stores all arguments passed to the pipeline.