Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
135 changes: 119 additions & 16 deletions weather_mv/loader_pipeline/bq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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 ...")
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -268,17 +281,24 @@ 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,
[lons] if isinstance(lons, float) else lons,
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.")
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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)

Expand All @@ -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

Expand All @@ -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))
Expand Down Expand Up @@ -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:
Expand Down
112 changes: 112 additions & 0 deletions weather_mv/loader_pipeline/bq_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,15 @@
# 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
import os
import tempfile
import typing as t
import unittest
from unittest import mock

import geojson
import numpy as np
Expand Down Expand Up @@ -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'
Expand Down
2 changes: 1 addition & 1 deletion weather_mv/loader_pipeline/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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.
Expand Down