commit 7ce3fe722b78d0faf07662fafb81c78376a40203 Author: nghiadang Date: Wed Aug 28 05:10:02 2024 +0000 first commit diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.dbf b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.dbf new file mode 100644 index 0000000..a806535 Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.dbf differ diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.prj b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.prj new file mode 100644 index 0000000..39dec21 --- /dev/null +++ b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.prj @@ -0,0 +1 @@ +PROJCS["Transverse_Mercator",GEOGCS["GCS_WGS_1984",DATUM["D_unknown",SPHEROID["WGS84",6378137,298.257223563]],PRIMEM["Greenwich",0],UNIT["Degree",0.017453292519943295]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105.5],PARAMETER["scale_factor",0.9999],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["Meter",1]] \ No newline at end of file diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.qpj b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.qpj new file mode 100644 index 0000000..8a1d453 --- /dev/null +++ b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.qpj @@ -0,0 +1 @@ +PROJCS["unnamed",GEOGCS["WGS 84",DATUM["unknown",SPHEROID["WGS84",6378137,298.257223563],TOWGS84[-192.873,-39.382,-111.202,-0.00205,-0.0005,0.00335,0.0188]],PRIMEM["Greenwich",0],UNIT["degree",0.0174532925199433]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105.5],PARAMETER["scale_factor",0.9999],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["Meter",1]] diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shp b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shp new file mode 100644 index 0000000..80519f5 Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shp differ diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shx b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shx new file mode 100644 index 0000000..256ef80 Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_KKDD2019.shx differ diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.dbf b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.dbf new file mode 100644 index 0000000..193f541 Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.dbf differ diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.prj b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.prj new file mode 100644 index 0000000..39dec21 --- /dev/null +++ b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.prj @@ -0,0 +1 @@ +PROJCS["Transverse_Mercator",GEOGCS["GCS_WGS_1984",DATUM["D_unknown",SPHEROID["WGS84",6378137,298.257223563]],PRIMEM["Greenwich",0],UNIT["Degree",0.017453292519943295]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105.5],PARAMETER["scale_factor",0.9999],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["Meter",1]] \ No newline at end of file diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.qpj b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.qpj new file mode 100644 index 0000000..8a1d453 --- /dev/null +++ b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.qpj @@ -0,0 +1 @@ +PROJCS["unnamed",GEOGCS["WGS 84",DATUM["unknown",SPHEROID["WGS84",6378137,298.257223563],TOWGS84[-192.873,-39.382,-111.202,-0.00205,-0.0005,0.00335,0.0188]],PRIMEM["Greenwich",0],UNIT["degree",0.0174532925199433]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105.5],PARAMETER["scale_factor",0.9999],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["Meter",1]] diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shp b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shp new file mode 100644 index 0000000..0bc61ec Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shp differ diff --git a/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shx b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shx new file mode 100644 index 0000000..3965ec7 Binary files /dev/null and b/ThuanHoa/KhoanhDat/ThuanHoa_TKDD2022.shx differ diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.cpg b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.cpg new file mode 100644 index 0000000..3ad133c --- /dev/null +++ b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.cpg @@ -0,0 +1 @@ +UTF-8 \ No newline at end of file diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.dbf b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.dbf new file mode 100644 index 0000000..cb95cc3 Binary files /dev/null and b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.dbf differ diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.prj b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.prj new file mode 100644 index 0000000..10ab055 --- /dev/null +++ b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.prj @@ -0,0 +1 @@ +PROJCS["VN-2000_TM-3_105-30",GEOGCS["GCS_VN_2000",DATUM["D_Vietnam_2000",SPHEROID["WGS_1984",6378137.0,298.257223563]],PRIMEM["Greenwich",0.0],UNIT["Degree",0.0174532925199433]],PROJECTION["Transverse_Mercator"],PARAMETER["False_Easting",500000.0],PARAMETER["False_Northing",0.0],PARAMETER["Central_Meridian",105.5],PARAMETER["Scale_Factor",0.9999],PARAMETER["Latitude_Of_Origin",0.0],UNIT["Meter",1.0]] \ No newline at end of file diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.qmd b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.qmd new file mode 100644 index 0000000..d702e82 --- /dev/null +++ b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.qmd @@ -0,0 +1,27 @@ + + + + + + + + + + + + + + + + + 0 + 0 + + + + + false + + + + diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shp b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shp new file mode 100644 index 0000000..b003a64 Binary files /dev/null and b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shp differ diff --git a/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shx b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shx new file mode 100644 index 0000000..471d202 Binary files /dev/null and b/ThuanHoa/region/ST_ThuanHoa_Boundaryofficially.shx differ diff --git a/deafrica_tools/__init__.py b/deafrica_tools/__init__.py new file mode 100644 index 0000000..610e60e --- /dev/null +++ b/deafrica_tools/__init__.py @@ -0,0 +1,26 @@ +__locales__ = __path__[0] + '/locales' + + +def set_lang(lang=None): + if lang is None: + import os + os_lang = os.getenv('LANG') + + # Just take the first 2 letters: 'fr' not 'fr_FR.UTF-8' + if os_lang is not None and len(os_lang) >=2: + lang = [os_lang[:2]] + else: + lang = [lang] + + import gettext + try: + translation = gettext.translation( + 'deafrica_tools', + localedir=__locales__, + languages=lang, + fallback=True + ) + translation.install() + + except FileNotFoundError: + print(f'Could not load lang={lang}') diff --git a/deafrica_tools/__pycache__/__init__.cpython-310.pyc b/deafrica_tools/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000..95a646e Binary files /dev/null and b/deafrica_tools/__pycache__/__init__.cpython-310.pyc differ diff --git a/deafrica_tools/__pycache__/bandindices.cpython-310.pyc b/deafrica_tools/__pycache__/bandindices.cpython-310.pyc new file mode 100644 index 0000000..f4eb3d9 Binary files /dev/null and b/deafrica_tools/__pycache__/bandindices.cpython-310.pyc differ diff --git a/deafrica_tools/app/__init__.py b/deafrica_tools/app/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/deafrica_tools/app/animations.py b/deafrica_tools/app/animations.py new file mode 100644 index 0000000..a0ad4fb --- /dev/null +++ b/deafrica_tools/app/animations.py @@ -0,0 +1,942 @@ +# -*- coding: utf-8 -*- +""" +Satellite imagery animation widget, which can be used to interactively +produce animations for multiple DE Africa products. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import fiona +import sys +import datacube +import warnings +import matplotlib.pyplot as plt +from datacube.utils.geometry import CRS +from ipyleaflet import ( + WMSLayer, + basemaps, + basemap_to_tiles, + Map, + DrawControl, + WidgetControl, + LayerGroup, + LayersControl, + GeoData, +) +from traitlets import Unicode +from ipywidgets import ( + GridspecLayout, + Button, + Layout, + HBox, + VBox, + HTML, + Output, +) +import json +import itertools +import numpy as np +import geopandas as gpd +from io import BytesIO +import ipywidgets as widgets +import datetime +from skimage import exposure +from skimage.filters import unsharp_mask + +from datacube.utils import masking +from datacube.utils.geometry import Geometry +from datacube.utils.masking import mask_invalid_data +import deafrica_tools.app.widgetconstructors as deawidgets +from deafrica_tools.dask import create_local_dask_cluster +from deafrica_tools.spatial import reverse_geocode +from deafrica_tools.datahandling import pan_sharpen_brovey + +import warnings + +warnings.filterwarnings("ignore") + + +# WMS params and satellite style bands +sat_params = { + "Landsat": { + "products": ["ls5_sr", "ls7_sr", "ls8_sr", "ls9_sr"], + "styles": { + "True colour": ("true_colour", ["red", "green", "blue"]), + "False colour": ( + "false_colour", + ["swir_1", "nir", "green"], + ), + }, + }, + "Sentinel-2": { + "products": ["s2_l2a"], + "styles": { + "True colour": ("simple_rgb", ["red", "green", "blue"]), + "False colour": ( + "infrared_green", + ["swir_2", "nir_1", "green"], + ), + }, + }, +} + + +def make_box_layout(): + return Layout( + # border='solid 1px black', + margin="0px 10px 10px 0px", + padding="5px 5px 5px 5px", + width="100%", + height="100%", + ) + + +def create_expanded_button(description, button_style): + return Button( + description=description, + button_style=button_style, + layout=Layout(width="auto", height="auto"), + ) + + +def update_map_layers(self): + """ + Updates map to add new DE Africa layers, styles or basemap when selected + using menu options. Triggers data reload by resetting load params + and output arrays. + """ + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Clear all layers and add basemap + self.map_layers.clear_layers() + self.map_layers.add_layer(self.basemap) + + +def extract_data(self): + + # Connect to datacube database + dc = datacube.Datacube(app="Exporting satellite images") + + # Configure local dask cluster + client = create_local_dask_cluster(return_client=True, display_client=True) + + # Convert to geopolygon + geopolygon = Geometry(geom=self.gdf_drawn.geometry[0], crs=self.gdf_drawn.crs) + + # Create query. + start_date = np.datetime64(self.start_date) + end_date = np.datetime64(self.end_date) + + self.query_params = { + "time": (str(start_date), str(end_date)), + "geopolygon": geopolygon, + } + + # Find matching datasets + dss = [ + dc.find_datasets(product=i, **self.query_params) + for i in sat_params[self.dealayer]["products"] + ] + dss = list(itertools.chain.from_iterable(dss)) + + # If data is found + if len(dss) > 0: + + # Get CRS + crs = str(dss[0].crs) + + self.load_params = { + "measurements": sat_params[self.dealayer]["styles"][self.style][1], + "resolution": (-self.resolution, self.resolution), + "output_crs": crs, + "group_by": "solar_day", + "dask_chunks": {"time": 1, "x": 2048, "y": 2048}, + "resampling": {"*": "cubic", "oa_fmask": "nearest", "fmask": "nearest"}, + } + + # Load data + from deafrica_tools.datahandling import load_ard + + timeseries_ds = load_ard( + dc=dc, + products=sat_params[self.dealayer]["products"], + min_gooddata=1.0 - (self.max_cloud_cover / 100), + ls7_slc_off=False, + mask_pixel_quality=self.cloud_mask, + **self.load_params, + **self.query_params, + ) + + # Set invalid nodata pixels to NaN + timeseries_ds = mask_invalid_data(timeseries_ds) + + # Else if no data is returned, return None + else: + timeseries_ds = None + + # Close down the dask client + client.close() + + return timeseries_ds.compute() + + +def plot_data(self, fname): + + # Data to plot + to_plot = self.timeseries_ds + + # If rolling median specified + if self.rolling_median: + with self.status_info: + print( + f"\nApplying rolling median ({self.rolling_median_window} timesteps window)" + ) + to_plot = to_plot.rolling( + time=int(self.rolling_median_window), center=True, min_periods=1 + ).median() + + # If resampling freq specified + if self.resample_freq: + with self.status_info: + print(f"\nResampling data to {self.resample_freq} frequency") + to_plot = to_plot.resample(time=self.resample_freq).median() + + # Raise by power to dampen bright features and enhance dark. + # Raise vmin and vmax by same amount to ensure proper stretch + if self.power < 1.0: + with self.status_info: + print(f"\nApplying power transformation ({self.power})") + to_plot = to_plot ** self.power + + # Apply unsharp masking to enhance overall dynamic range, + # and improve fine scale detail + if self.unsharp_mask: + with self.status_info: + print( + f"\nApplying unsharp masking with {self.unsharp_mask_radius} " + f"radius and {self.unsharp_mask_amount} amount" + ) + from skimage.exposure import rescale_intensity + + funcs_list = [ + rescale_intensity, + lambda x: unsharp_mask( + x, radius=self.unsharp_mask_radius, amount=self.unsharp_mask_amount + ), + ] + else: + funcs_list = None + + from deafrica_tools.plotting import xr_animation + + xr_animation( + output_path=fname, + ds=to_plot.dropna(dim="time", how="all"), + show_text="", + bands=sat_params[self.dealayer]["styles"][self.style][1], + interval=self.interval, + width_pixels=self.width, + show_gdf=deacoastlines_overlay(to_plot) if self.deacoastlines else None, + gdf_kwargs={"linewidth": 3}, + percentile_stretch=(self.vmin, self.vmax), + image_proc_funcs=funcs_list, + show_date="%Y" if self.resample_freq == "1Y" else "%b %Y", + annotation_kwargs={"fontsize": 75}, + ) + + # Add plot preview below map and finish + plt.show() + with self.status_info: + print(f"\nImage successfully exported to:\n{fname}.") + + +def deacoastlines_overlay(ds): + + import geopandas as gpd + import pandas as pd + import matplotlib + from shapely.geometry import box, Point + from deafrica_tools.coastal import get_coastlines + + # Get bounding box of data + xmin, ymin, xmax, ymax = ds.geobox.geographic_extent.boundingbox + bounds = [xmin, ymin, xmax, ymax] + + # Load data + deacl_gdf = get_coastlines(bbox=bounds) + + # Clip to extent of satellite data + bbox = gpd.GeoDataFrame(geometry=[ds.geobox.extent.geom], crs=ds.geobox.crs) + deacl_gdf = gpd.overlay(deacl_gdf, bbox.to_crs(deacl_gdf.crs)) + deacl_gdf = deacl_gdf.dissolve("year") # values("year", ascending=True) + + # Apply colours + norm = matplotlib.colors.Normalize(vmin=0, vmax=len(deacl_gdf.index)) + cmap = matplotlib.cm.get_cmap("inferno") + rgba = cmap(norm(deacl_gdf.reset_index().index)) + deacl_gdf["color"] = list(rgba) + deacl_gdf["start_time"] = pd.to_datetime(deacl_gdf.index) + pd.DateOffset(months=0) + deacl_gdf = deacl_gdf.sort_index() + + if len(deacl_gdf.index) > 0: + return deacl_gdf + else: + return None + + +class animation_app(HBox): + def __init__(self): + super().__init__() + + ###################### + # INITIAL ATTRIBUTES # + ###################### + + # Basemap + self.basemap_list = [ + ("ESRI World Imagery", basemap_to_tiles(basemaps.Esri.WorldImagery)), + ("Open Street Map", basemap_to_tiles(basemaps.OpenStreetMap.Mapnik)), + ] + self.basemap = self.basemap_list[0][1] + + # Satellite data + end_date = datetime.datetime.today() + start_date = datetime.datetime( + year=end_date.year - 3, month=end_date.month, day=end_date.day + ) + self.start_date = start_date.strftime("%Y-%m-%d") + self.end_date = end_date.strftime("%Y-%m-%d") + self.dealayer_list = [ + ("Landsat", "Landsat"), + ("Sentinel-2", "Sentinel-2"), + ] + self.dealayer = self.dealayer_list[0][1] + + # Styles + self.styles_list = ["True colour", "False colour"] + self.style = self.styles_list[0] + + # Analysis params + self.resolution = 30 + self.vmin = 0.01 + self.vmax = 0.99 + self.power = 1.0 + self.output_list = [("MP4", "mp4"), ("GIF", "gif")] + self.output_format = self.output_list[0][1] + self.rolling_median = False + self.rolling_median_window = 20 + self.unsharp_mask = False + self.unsharp_mask_radius = 20 + self.unsharp_mask_amount = 0.3 + self.max_size = False + self.width = 900 + self.interval = 100 + self.cloud_mask = False + self.max_cloud_cover = 20 + self.resample_list = [ + ("None", False), + ("Monthly", "1M"), + ("Quarterly", "Q-DEC"), + ("Yearly", "1Y"), + ] + self.resample_freq = self.resample_list[0][1] + self.deacoastlines = False + + # Drawing params + self.target = None + self.action = None + self.gdf_drawn = None + + # Data load params + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + ################## + # HEADER FOR APP # + ################## + + # Create the Header widget + header_title_text = ( + "

Digital Earth Africa satellite imagery animations

" + ) + instruction_text = ( + "

Select the desired satellite data, imagery date range " + "and image style, then zoom in and draw a rectangle to " + "select an area export as a satellite imagery time-series " + "animation.

" + ) + self.header = deawidgets.create_html(f"{header_title_text}{instruction_text}") + self.header.layout = make_box_layout() + + ##################################### + # HANDLER FUNCTION FOR DRAW CONTROL # + ##################################### + + # Define the action to take once something is drawn on the map + def update_geojson(target, action, geo_json): + + # Get data from action + self.action = action + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Convert data to geopandas + json_data = json.dumps(geo_json) + binary_data = json_data.encode() + io = BytesIO(binary_data) + io.seek(0) + gdf = gpd.read_file(io) + gdf.crs = "EPSG:4326" + + # Convert to WGS 84 / NSIDC EASE-Grid 2.0 Global and compute area + gdf_drawn_nsidc = gdf.copy().to_crs("EPSG:6933") + m2_per_ha = 10000 + area = gdf_drawn_nsidc.area.values[0] / m2_per_ha + polyarea_label = "Total area of satellite data to extract" + polyarea_text = f"{polyarea_label}: {area:.2f} ha" + + # Test area size + if self.max_size: + confirmation_text = ( + ' ' + "(Overriding maximum size limit; use with caution as may lead to memory issues)" + ) + self.header.value = ( + header_title_text + + instruction_text + + polyarea_text + + confirmation_text + ) + self.gdf_drawn = gdf + elif area <= 50000: + confirmation_text = ( + ' ' + "(Area to extract falls within " + "recommended 50000 ha limit)" + ) + self.header.value = ( + header_title_text + + instruction_text + + polyarea_text + + confirmation_text + ) + self.gdf_drawn = gdf + else: + warning_text = ( + ' ' + "(Area to extract is too large, " + "please select an area less than 50000 " + "ha)" + ) + self.header.value = ( + header_title_text + instruction_text + polyarea_text + warning_text + ) + self.gdf_drawn = None + + ########################### + # WIDGETS FOR APP OUTPUTS # + ########################### + + self.status_info = Output(layout=make_box_layout()) + self.output_plot = Output(layout=make_box_layout()) + + ######################################### + # MAP WIDGET, DRAWING TOOLS, WMS LAYERS # + ######################################### + + # Create drawing tools + desired_drawtools = ["rectangle"] + draw_control = deawidgets.create_drawcontrol(desired_drawtools) + + # Begin by displaying an empty layer group, and update the group with desired WMS on interaction. + self.map_layers = LayerGroup(layers=()) + self.map_layers.name = "Map Overlays" + + # Create map widget + self.m = deawidgets.create_map(map_center=(5.65, 26.17), zoom_level=13) + self.m.layout = make_box_layout() + + # Add tools to map widget + self.m.add_control(draw_control) + self.m.add_layer(self.map_layers) + + # Update all maps to starting defaults + update_map_layers(self) + + ############################ + # WIDGETS FOR APP CONTROLS # + ############################ + + # Create parameter widgets + dropdown_basemap = deawidgets.create_dropdown( + self.basemap_list, self.basemap_list[0][1] + ) + dropdown_dealayer = deawidgets.create_dropdown( + self.dealayer_list, self.dealayer_list[0][1] + ) + dropdown_output = deawidgets.create_dropdown( + self.output_list, self.output_list[0][1] + ) + date_picker_start = deawidgets.create_datepicker( + value=start_date, + ) + date_picker_end = deawidgets.create_datepicker( + value=end_date, + ) + dropdown_styles = deawidgets.create_dropdown( + self.styles_list, self.styles_list[0] + ) + slider_percentile = widgets.FloatRangeSlider( + value=[0.01, 0.99], + min=0, + max=1, + step=0.001, + description="", + layout={"width": "85%"}, + ) + run_button = create_expanded_button("Generate animation", "info") + + floatslider_max_cloud_cover = widgets.IntSlider( + value=20, + min=0, + max=100, + step=1, + description="", + layout={"width": "85%"}, + ) + + checkbox_rolling_median = deawidgets.create_checkbox( + self.rolling_median, + "Apply rolling median
to produce smooth,
cloud-free animations", + layout={"width": "90%", + "height": "4em"}, + ) + text_rolling_median_window = widgets.IntText( + value=20, + step=1, + description="
Rolling window (timesteps)", + layout={ + "width": "85%", + "margin": "0px", + "padding": "0px", + "display": "none", + }, + ) + + # Expandable advanced section + text_interval = widgets.IntText( + value=100, description="", step=50, layout={"width": "95%"} + ) + text_resolution = widgets.FloatText( + value=30, + description="", + layout={"width": "95%", "margin": "0px", "padding": "0px"}, + ) + text_width = widgets.IntText( + value=900, description="", step=50, layout={"width": "95%"} + ) + dropdown_resampling = deawidgets.create_dropdown( + self.resample_list, + self.resample_freq, + description="", + layout={"width": "95%"}, + ) + checkbox_cloud_mask = deawidgets.create_checkbox( + self.cloud_mask, "Mask out cloudy
pixels", layout={"width": "95%", "height": "auto"} + ) + slider_power = widgets.FloatSlider( + value=1.0, + min=0.01, + max=1.0, + step=0.01, + description="", + layout={"width": "95%"}, + ) + checkbox_unsharp_mask = deawidgets.create_checkbox( + self.unsharp_mask, "Enable", layout={"width": "95%"} + ) + text_unsharp_mask_radius = widgets.FloatText( + value=20, + step=1, + description="Radius", + layout={ + "width": "95%", + "margin": "0px", + "padding": "0px", + "display": "none", + }, + ) + text_unsharp_mask_amount = widgets.FloatText( + value=0.3, + step=0.1, + description="Amount", + layout={ + "width": "95%", + "margin": "0px", + "padding": "0px", + "display": "none", + }, + ) + checkbox_deacoastlines = deawidgets.create_checkbox( + self.deacoastlines, "Add DE Africa Coastlines overlay", layout={"width": "95%"} + ) + checkbox_max_size = deawidgets.create_checkbox( + self.max_size, "Enable", layout={"width": "95%"} + ) + expand_box = widgets.VBox( + [ + HTML("Frame interval (milliseconds):"), + text_interval, + HTML("
Resolution (metres):"), + text_resolution, + HTML("
Width of output animation in pixels:"), + text_width, + HTML("
Apply temporal resampling:"), + dropdown_resampling, + HTML("
"), + checkbox_cloud_mask, + checkbox_deacoastlines, + HTML("
Apply power transformation to darken bright features:"), + slider_power, + HTML("
Apply unsharp masking to sharpen imagery:"), + checkbox_unsharp_mask, + text_unsharp_mask_radius, + text_unsharp_mask_amount, + HTML( + "
Override maximum size limit: (use with caution; may cause memory issues/crashes)" + ), + checkbox_max_size, + ], + ) + + expand = widgets.Accordion( + children=[expand_box], + selected_index=None, + ) + expand.set_title(0, "Advanced") + + # Add specific dialogs to class so they can be modified + self.text_resolution = text_resolution + self.text_unsharp_mask_radius = text_unsharp_mask_radius + self.text_unsharp_mask_amount = text_unsharp_mask_amount + self.text_rolling_median_window = text_rolling_median_window + + #################################### + # UPDATE FUNCTIONS FOR EACH WIDGET # + #################################### + + # Run update functions whenever various widgets are changed. + date_picker_start.observe(self.update_start_date, "value") + date_picker_end.observe(self.update_end_date, "value") + dropdown_basemap.observe(self.update_basemap, "value") + dropdown_dealayer.observe(self.update_dealayer, "value") + dropdown_styles.observe(self.update_styles, "value") + + slider_percentile.observe(self.update_slider_percentile, "value") + floatslider_max_cloud_cover.observe( + self.update_floatslider_max_cloud_cover, "value" + ) + checkbox_rolling_median.observe(self.update_checkbox_rolling_median, "value") + text_rolling_median_window.observe( + self.update_text_rolling_median_window, "value" + ) + dropdown_output.observe(self.update_output, "value") + run_button.on_click(self.run_app) + draw_control.on_draw(update_geojson) + + # Advanced params + text_resolution.observe(self.update_text_resolution, "value") + slider_power.observe(self.update_slider_power, "value") + text_width.observe(self.update_width, "value") + text_interval.observe(self.update_interval, "value") + dropdown_resampling.observe(self.update_dropdown_resampling, "value") + checkbox_cloud_mask.observe(self.update_checkbox_cloud_mask, "value") + checkbox_unsharp_mask.observe(self.update_checkbox_unsharp_mask, "value") + text_unsharp_mask_radius.observe(self.update_text_unsharp_mask_radius, "value") + text_unsharp_mask_amount.observe(self.update_text_unsharp_mask_amount, "value") + checkbox_deacoastlines.observe(self.update_deacoastlines, "value") + checkbox_max_size.observe(self.update_checkbox_max_size, "value") + + ################################## + # COLLECTION OF ALL APP CONTROLS # + ################################## + + parameter_selection = VBox( + [ + HTML("Satellite imagery:"), + dropdown_dealayer, + HTML("Start date:"), + date_picker_start, + HTML("End date:"), + date_picker_end, + HTML("Style:"), + dropdown_styles, + HTML("Colour percentile stretch:"), + slider_percentile, + HTML("Maximum cloud cover (%):"), + floatslider_max_cloud_cover, + checkbox_rolling_median, + text_rolling_median_window, + HTML("
Output file format:"), + dropdown_output, + HTML("
"), + expand, + ] + ) + map_selection = VBox( + [ + HTML("
Map overlay:"), + dropdown_basemap, + ] + ) + parameter_selection.layout = make_box_layout() + map_selection.layout = make_box_layout() + + ############################### + # SPECIFICATION OF APP LAYOUT # + ############################### + + # 0 1 2 3 4 5 6 7 8 9 + # --------------------------------------------- + # 0 | Header | Map sel. | + # |-------------------------------------------| + # 1 | Params | | + # 2 | | | + # 3 | | | + # 4 | | Map | + # 5 | | | + # |--------| | + # 6 | Run | | + # |-------------------------------------------| + # 7 | Status info | Figure/output | + # 8 | | | + # 9 | | | + # 10 | | | + # 11 --------------------------------------------- + + # Create the layout #[rowspan, colspan] + grid = GridspecLayout(12, 10, height="1500px", width="auto") + + # Header and controls + grid[0, :8] = self.header + grid[0, 8:] = map_selection + grid[1:6, 0:2] = parameter_selection + grid[6, 0:2] = run_button + + # Status info, map and plot + grid[1:7, 2:] = self.m # map + grid[7:, 0:4] = self.status_info + grid[7:, 4:] = self.output_plot + + # Display using HBox children attribute + self.children = [grid] + + ###################################### + # DEFINITION OF ALL UPDATE FUNCTIONS # + ###################################### + + # Update date + def update_start_date(self, change): + self.start_date = str(change.new) + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Update date + def update_end_date(self, change): + self.end_date = str(change.new) + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Update colour stretch + def update_slider_percentile(self, change): + self.vmin, self.vmax = change.new + + # Update power transform + def update_slider_power(self, change): + self.power = change.new + + # Update good data slider + def update_floatslider_max_cloud_cover(self, change): + self.max_cloud_cover = change.new + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Enable unsharp masking and show/hide custom params + def update_checkbox_unsharp_mask(self, change): + self.unsharp_mask = change.new + + # Show unsharp masking params in menu if activated + if change.new: + self.text_unsharp_mask_radius.layout.display = "block" + self.text_unsharp_mask_amount.layout.display = "block" + else: + self.text_unsharp_mask_radius.layout.display = "none" + self.text_unsharp_mask_amount.layout.display = "none" + + # Change unsharp masking radius + def update_text_unsharp_mask_radius(self, change): + self.unsharp_mask_radius = change.new + + # Change unsharp masking amount + def update_text_unsharp_mask_amount(self, change): + self.unsharp_mask_amount = change.new + + # Enable rolling median and show/hide custom params + def update_checkbox_rolling_median(self, change): + self.rolling_median = change.new + + # Show rolling median params in menu if activated + if change.new: + self.text_rolling_median_window.layout.display = "block" + else: + self.text_rolling_median_window.layout.display = "none" + + # Change rolling median window + def update_text_rolling_median_window(self, change): + self.rolling_median_window = change.new + + # Override max size limit + def update_checkbox_max_size(self, change): + self.max_size = change.new + + # Add DE Africa Coastlines overlay + def update_deacoastlines(self, change): + self.deacoastlines = change.new + + # Apply cloud mask in load_ard + def update_checkbox_cloud_mask(self, change): + self.cloud_mask = change.new + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Override min width + def update_width(self, change): + self.width = change.new + + # Override interval + def update_interval(self, change): + self.interval = change.new + + # Update resolution + def update_text_resolution(self, change): + self.resolution = change.new + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Change layers shown on the map + def update_dealayer(self, change): + self.dealayer = change.new + + if change.new == "Landsat": + self.text_resolution.value = 30 + + else: + self.text_resolution.value = 10 + + # Update basemap + def update_basemap(self, change): + self.basemap = change.new + update_map_layers(self) + + # Set imagery style + def update_styles(self, change): + self.style = change.new + + # Clear data load params to trigger data re-load + self.timeseries_ds = None + self.load_params = None + self.query_params = None + + # Set output file format + def update_output(self, change): + self.output_format = change.new + + # Set output file format + def update_dropdown_resampling(self, change): + self.resample_freq = change.new + + def run_app(self, change): + + # Clear progress bar and output areas before running + self.status_info.clear_output() + self.output_plot.clear_output() + + # Verify that polygon was drawn + if self.gdf_drawn is not None: + + with self.status_info: + + # Load data and add to attribute + if self.timeseries_ds is None: + self.timeseries_ds = extract_data(self) + + else: + print("Using previously loaded data") + + if self.timeseries_ds is not None: + + with self.status_info: + + # Create unique file name + centre_coords = self.gdf_drawn.geometry[0].centroid.coords[0][::-1] + site = reverse_geocode(coords=centre_coords) + fname = ( + f"{self.dealayer}_{site}_{self.start_date}_" + f"{self.end_date}_{self.style}_{self.resolution:.0f}m." + f"{self.output_format}".replace(" ", "") + .replace(",", "") + .lower() + ) + + print( + f"\nExporting animation for {site}.\nThis may take several minutes..." + ) + + ############ + # Plotting # + ############ + + with self.output_plot: + plot_data(self, fname) + + else: + with self.status_info: + print( + "No satellite data found in the selected area. " + "Please select a new rectangle over an area with " + "satellite imagery." + ) + + else: + with self.status_info: + print( + 'Please draw a valid rectangle on the map, then press "Generate animation".' + ) \ No newline at end of file diff --git a/deafrica_tools/app/changefilmstrips.py b/deafrica_tools/app/changefilmstrips.py new file mode 100644 index 0000000..31fb2d8 --- /dev/null +++ b/deafrica_tools/app/changefilmstrips.py @@ -0,0 +1,275 @@ +""" +Loading and interacting with data in the change filmstrips notebook, +inside the Real_world_examples folder. +""" + +# Load modules +import os +import dask +import datacube +import warnings +import numpy as np +import pandas as pd +import xarray as xr +import matplotlib.pyplot as plt +from odc.algo import geomedian_with_mads +from odc.ui import select_on_a_map +from dask.utils import parse_bytes +from datacube.utils.geometry import CRS, assign_crs +from datacube.utils.rio import configure_s3_access +from datacube.utils.dask import start_local_dask +from ipyleaflet import basemaps, basemap_to_tiles + +# Load utility functions +from deafrica_tools.datahandling import load_ard, mostcommon_crs +from deafrica_tools.dask import create_local_dask_cluster + + +def run_filmstrip_app( + output_name, + time_range, + time_step, + tide_range=(0.0, 1.0), + resolution=(-30, 30), + max_cloud=0.5, + ls7_slc_off=False, + size_limit=10000, +): + """ + An interactive app that allows the user to select a region from a + map, then load Digital Earth Africa Landsat data and combine it + using the geometric median ("geomedian") statistic to reveal the + median or 'typical' appearance of the landscape for a series of + time periods. + + The results for each time period are combined into a 'filmstrip' + plot which visualises how the landscape has changed in appearance + across time, with a 'change heatmap' panel highlighting potential + areas of greatest change. + + For coastal applications, the analysis can be customised to select + only satellite images obtained during a specific tidal range + (e.g. low, average or high tide). + + Last modified: April 2020 + + Parameters + ---------- + output_name : str + A name that will be used to name the output filmstrip plot file. + time_range : tuple + A tuple giving the date range to analyse + (e.g. `time_range = ('1988-01-01', '2017-12-31')`). + time_step : dict + This parameter sets the length of the time periods to compare + (e.g. `time_step = {'years': 5}` will generate one filmstrip + plot for every five years of data; `time_step = {'months': 18}` + will generate one plot for each 18 month period etc. Time + periods are counted from the first value given in `time_range`. + tide_range : tuple, optional + An optional parameter that can be used to generate filmstrip + plots based on specific ocean tide conditions. This can be + valuable for analysing change consistently along the coast. + For example, `tide_range = (0.0, 0.2)` will select only + satellite images acquired at the lowest 20% of tides; + `tide_range = (0.8, 1.0)` will select images from the highest + 20% of tides. The default is `tide_range = (0.0, 1.0)` which + will select all images regardless of tide. + resolution : tuple, optional + The spatial resolution to load data. The default is + `resolution = (-30, 30)`, which will load data at 30 m pixel + resolution. Increasing this (e.g. to `resolution = (-100, 100)`) + can be useful for loading large spatial extents. + max_cloud : float, optional + This parameter can be used to exclude satellite images with + excessive cloud. The default is `0.5`, which will keep all images + with less than 50% cloud. + ls7_slc_off : bool, optional + An optional boolean indicating whether to include data from + after the Landsat 7 SLC failure (i.e. SLC-off). Defaults to + False, which removes all Landsat 7 observations > May 31 2003. + size_limit : int, optional + An optional integer (in hectares) specifying the size limit + for the data query. Queries larger than this size will receive + a warning that he data query is too large (and may + therefore result in memory errors). + + + Returns + ------- + ds_geomedian : xarray Dataset + An xarray dataset containing geomedian composites for each + timestep in the analysis. + + """ + + ######################## + # Select and load data # + ######################## + + # Define centre_coords as a global variable + global centre_coords + + # Test if centre_coords is in the global namespace; + # use default value if it isn't + if "centre_coords" not in globals(): + centre_coords = (6.587292, 1.532833) + + # Plot interactive map to select area + basemap = basemap_to_tiles(basemaps.Esri.WorldImagery) + geopolygon = select_on_a_map(height="600px", + layers=(basemap,), + center=centre_coords, + zoom=14) + + # Set centre coords based on most recent selection to re-focus + # subsequent data selections + centre_coords = geopolygon.centroid.points[0][::-1] + + # Test size of selected area + msq_per_hectare = 10000 + area = geopolygon.to_crs(crs=CRS("epsg:6933")).area / msq_per_hectare + radius = np.round(np.sqrt(size_limit), 1) + if area > size_limit: + print(f"Warning: Your selected area is {area:.00f} hectares. " + f"Please select an area of less than {size_limit} hectares." + f"\nTo select a smaller area, re-run the cell " + f"above and draw a new polygon.") + + else: + + print("Starting analysis...") + + # Connect to datacube database + dc = datacube.Datacube(app="Change_filmstrips") + + # Configure local dask cluster + client = create_local_dask_cluster(return_client=True) + + # Obtain native CRS + crs = mostcommon_crs(dc=dc, + product="ls8_sr", + query={ + "time": "2014", + "geopolygon": geopolygon + }) + + # Create query based on time range, area selected, custom params + query = { + "time": time_range, + "geopolygon": geopolygon, + "output_crs": crs, + "resolution": resolution, + "dask_chunks": { + "x": 3000, + "y": 3000 + }, + "align": (resolution[1] / 2.0, resolution[1] / 2.0), + } + + # Load data from all three Landsats + warnings.filterwarnings("ignore") + ds = load_ard( + dc=dc, + measurements=["red", "green", "blue"], + products=["ls5_sr", "ls7_sr", "ls8_sr"], + min_gooddata=max_cloud, + ls7_slc_off=ls7_slc_off, + **query, + ) + + # Optionally calculate tides for each timestep in the satellite + # dataset and drop any observations out side this range + if tide_range != (0.0, 1.0): + from deafrica_tools.coastal import tidal_tag + ds = tidal_tag(ds=ds, tidepost_lat=None, tidepost_lon=None) + min_tide, max_tide = ds.tide_height.quantile(tide_range).values + ds = ds.sel(time=(ds.tide_height >= min_tide) & + (ds.tide_height <= max_tide)) + ds = ds.drop("tide_height") + print(f" Keeping {len(ds.time)} observations with tides " + f"between {min_tide:.2f} and {max_tide:.2f} m") + + # Create time step ranges to generate filmstrips from + bins_dt = pd.date_range(start=time_range[0], + end=time_range[1], + freq=pd.DateOffset(**time_step)) + + # Bin all satellite observations by timestep. If some observations + # fall outside the upper bin, label these with the highest bin + labels = bins_dt.astype("str") + time_steps = (pd.cut(ds.time.values, bins_dt, + labels=labels[:-1]).add_categories( + labels[-1]).fillna(labels[-1])) + + time_steps_var = xr.DataArray(time_steps, [("time", ds.time.values)], + name="timestep") + + # Resample data temporally into time steps, and compute geomedians + ds_geomedian = (ds.groupby(time_steps_var).apply( + lambda ds_subset: geomedian_with_mads( + ds_subset, compute_mads=False, compute_count=False))) + + print("\nGenerating geomedian composites and plotting " + "filmstrips... (click the Dashboard link above for status)") + ds_geomedian = ds_geomedian.compute() + + # Reset CRS that is lost during geomedian compositing + ds_geomedian = assign_crs(ds_geomedian, crs=ds.geobox.crs) + + ############ + # Plotting # + ############ + + # Convert to array and extract vmin/vmax + output_array = ds_geomedian[["red", "green", "blue"]].to_array() + percentiles = output_array.quantile(q=(0.02, 0.98)).values + + # Create the plot with one subplot more than timesteps in the + # dataset. Figure width is set based on the number of subplots + # and aspect ratio + n_obs = output_array.sizes["timestep"] + ratio = output_array.sizes["x"] / output_array.sizes["y"] + fig, axes = plt.subplots(1, + n_obs + 1, + figsize=(5 * ratio * (n_obs + 1), 5)) + fig.subplots_adjust(wspace=0.05, hspace=0.05) + + # Add timesteps to the plot, set aspect to equal to preserve shape + for i, ax_i in enumerate(axes.flatten()[:n_obs]): + output_array.isel(timestep=i).plot.imshow(ax=ax_i, + vmin=percentiles[0], + vmax=percentiles[1]) + ax_i.get_xaxis().set_visible(False) + ax_i.get_yaxis().set_visible(False) + ax_i.set_aspect("equal") + + # Add change heatmap panel to final subplot. Heatmap is computed + # by first taking the log of the array (so change in dark areas + # can be identified), then computing standard deviation between + # all timesteps + (np.log(output_array).std(dim=["timestep"]).mean( + dim="variable").plot.imshow(ax=axes.flatten()[-1], + robust=True, + cmap="magma", + add_colorbar=False)) + axes.flatten()[-1].get_xaxis().set_visible(False) + axes.flatten()[-1].get_yaxis().set_visible(False) + axes.flatten()[-1].set_aspect("equal") + axes.flatten()[-1].set_title("Change heatmap") + + # Export to file + date_string = "_".join(time_range) + ts_v = list(time_step.values())[0] + ts_k = list(time_step.keys())[0] + fig.savefig( + f"filmstrip_{output_name}_{date_string}_{ts_v}{ts_k}.png", + dpi=150, + bbox_inches="tight", + pad_inches=0.1, + ) + + # close dask client + client.shutdown() + + return ds_geomedian diff --git a/deafrica_tools/app/crophealth.py b/deafrica_tools/app/crophealth.py new file mode 100644 index 0000000..f6fd091 --- /dev/null +++ b/deafrica_tools/app/crophealth.py @@ -0,0 +1,337 @@ +# crophealth.py +''' +Functions for loading and interacting with data in the crop health notebook, + inside the Real_world_examples folder. +''' + +# Load modules + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +from ipyleaflet import ( + Map, + GeoJSON, + DrawControl, + basemaps +) +import datetime as dt +import datacube +from osgeo import ogr +import matplotlib as mpl +import matplotlib.pyplot as plt +import rasterio +from rasterio.features import geometry_mask +import xarray as xr +from IPython.display import display +import warnings +import ipywidgets as widgets +import json +import geopandas as gpd +from io import BytesIO + +# Load utility functions +from deafrica_tools.datahandling import load_ard +from deafrica_tools.spatial import xr_rasterize +from deafrica_tools.bandindices import calculate_indices + + +def load_crophealth_data(lat, lon, buffer, date): + """ + Loads Sentinel-2 analysis-ready data (ARD) product for the crop health + case-study area over the last two years. + Last modified: April 2020 + + Parameters + ---------- + lat: float + The central latitude to analyse + lon: float + The central longitude to analyse + buffer: + The number of square degrees to load around the central latitude and longitude. + For reasonable loading times, set this as `0.1` or lower. + date: + The most recent date to show data for. + The app will automatically load all data available for the two years prior to this date. + + Returns + ---------- + ds: xarray.Dataset + data set containing combined, masked data + Masked values are set to 'nan' + """ + + # Suppress warnings + warnings.filterwarnings('ignore') + + # Initialise the data cube. 'app' argument is used to identify this app + dc = datacube.Datacube(app='Crophealth-app') + + # Define area to load + latitude = (lat - buffer, lat + buffer) + longitude = (lon - buffer, lon + buffer) + + # Specify the date range + # Calculated as today's date, subtract 730 days to collect two years of data + # Dates are converted to strings as required by loading function below + end_date = dt.datetime.strptime(date, "%Y-%m-%d") + start_date = end_date - dt.timedelta(days=730) + + time = (start_date.strftime("%Y-%m-%d"), end_date.strftime("%Y-%m-%d")) + + # Construct the data cube query + products = ["s2_l2a"] + + query = { + 'x': longitude, + 'y': latitude, + 'time': time, + 'measurements': [ + 'red', + 'green', + 'blue', + 'nir', + 'swir_2' + ], + 'output_crs': 'EPSG:6933', + 'resolution': (-20, 20) + } + + # Load the data and mask out bad quality pixels + ds = load_ard(dc, products=products, min_gooddata=0.5, **query) + + # Calculate the normalised difference vegetation index (NDVI) across + # all pixels for each image. + # This is stored as an attribute of the data + ds = calculate_indices(ds, index='NDVI', satellite_mission='s2') + + # Return the data + return(ds) + + +def run_crophealth_app(ds, lat, lon, buffer): + """ + Plots an interactive map of the crop health case-study area and allows + the user to draw polygons. This returns a plot of the average NDVI value + in the polygon area. + Last modified: January 2020 + + Parameters + ---------- + ds: xarray.Dataset + data set containing combined, masked data + Masked values are set to 'nan' + lat: float + The central latitude corresponding to the area of loaded ds + lon: float + The central longitude corresponding to the area of loaded ds + buffer: + The number of square degrees to load around the central latitude and longitude. + For reasonable loading times, set this as `0.1` or lower. + """ + + # Suppress warnings + warnings.filterwarnings('ignore') + + # Update plotting functionality through rcParams + mpl.rcParams.update({'figure.autolayout': True}) + + # Define polygon bounds + latitude = (lat - buffer, lat + buffer) + longitude = (lon - buffer, lon + buffer) + + # Define the bounding box that will be overlayed on the interactive map + # The bounds are hard-coded to match those from the loaded data + geom_obj = { + "type": "Feature", + "properties": { + "style": { + "stroke": True, + "color": 'red', + "weight": 4, + "opacity": 0.8, + "fill": True, + "fillColor": False, + "fillOpacity": 0, + "showArea": True, + "clickable": True + } + }, + "geometry": { + "type": "Polygon", + "coordinates": [ + [ + [ + longitude[0], + latitude[0] + ], + [ + longitude[1], + latitude[0] + ], + [ + longitude[1], + latitude[1] + ], + [ + longitude[0], + latitude[1] + ], + [ + longitude[0], + latitude[0] + ] + ] + ] + } + } + + # Create a map geometry from the geom_obj dictionary + # center specifies where the background map view should focus on + # zoom specifies how zoomed in the background map should be + loadeddata_geometry = ogr.CreateGeometryFromJson(str(geom_obj['geometry'])) + loadeddata_center = [ + loadeddata_geometry.Centroid().GetY(), + loadeddata_geometry.Centroid().GetX() + ] + loadeddata_zoom = 16 + + # define the study area map + studyarea_map = Map( + center=loadeddata_center, + zoom=loadeddata_zoom, + basemap=basemaps.Esri.WorldImagery + ) + + # define the drawing controls + studyarea_drawctrl = DrawControl( + polygon={"shapeOptions": {"fillOpacity": 0}}, + marker={}, + circle={}, + circlemarker={}, + polyline={}, + ) + + # add drawing controls and data bound geometry to the map + studyarea_map.add_control(studyarea_drawctrl) + studyarea_map.add_layer(GeoJSON(data=geom_obj)) + + # Index to count drawn polygons + polygon_number = 0 + + # Define widgets to interact with + instruction = widgets.Output(layout={'border': '1px solid black'}) + with instruction: + print("Draw a polygon within the red box to view a plot of " + "average NDVI over time in that area.") + + info = widgets.Output(layout={'border': '1px solid black'}) + with info: + print("Plot status:") + + fig_display = widgets.Output(layout=widgets.Layout( + width="50%", # proportion of horizontal space taken by plot + )) + + with fig_display: + plt.ioff() + fig, ax = plt.subplots(figsize=(8, 6)) + ax.set_ylim([0, 1]) + + colour_list = plt.rcParams['axes.prop_cycle'].by_key()['color'] + + # Function to execute each time something is drawn on the map + def handle_draw(self, action, geo_json): + nonlocal polygon_number + + # Execute behaviour based on what the user draws + if geo_json['geometry']['type'] == 'Polygon': + + info.clear_output(wait=True) # wait=True reduces flicker effect + + # Save geojson polygon to io temporary file to be rasterized later + jsonData = json.dumps(geo_json) + binaryData = jsonData.encode() + io = BytesIO(binaryData) + io.seek(0) + + # Read the polygon as a geopandas dataframe + gdf = gpd.read_file(io) + gdf.crs = "EPSG:4326" + + # Convert the drawn geometry to pixel coordinates + xr_poly = xr_rasterize(gdf, ds.NDVI.isel(time=0), crs='EPSG:6933') + + # Construct a mask to only select pixels within the drawn polygon + masked_ds = ds.NDVI.where(xr_poly) + + masked_ds_mean = masked_ds.mean(dim=['x', 'y'], skipna=True) + colour = colour_list[polygon_number % len(colour_list)] + + # Add a layer to the map to make the most recently drawn polygon + # the same colour as the line on the plot + studyarea_map.add_layer( + GeoJSON( + data=geo_json, + style={ + 'color': colour, + 'opacity': 1, + 'weight': 4.5, + 'fillOpacity': 0.0 + } + ) + ) + + # add new data to the plot + xr.plot.plot( + masked_ds_mean, + marker='*', + color=colour, + ax=ax + ) + + # reset titles back to custom + ax.set_title("Average NDVI from Sentinel-2") + ax.set_xlabel("Date") + ax.set_ylabel("NDVI") + + # refresh display + fig_display.clear_output(wait=True) # wait=True reduces flicker effect + with fig_display: + display(fig) + + with info: + print("Plot status: polygon sucessfully added to plot.") + + # Iterate the polygon number before drawing another polygon + polygon_number = polygon_number + 1 + + else: + info.clear_output(wait=True) + with info: + print("Plot status: this drawing tool is not currently " + "supported. Please use the polygon tool.") + + # call to say activate handle_draw function on draw + studyarea_drawctrl.on_draw(handle_draw) + + with fig_display: + # TODO: update with user friendly something + display(widgets.HTML("")) + + # Construct UI: + # +-----------------------+ + # | instruction | + # +-----------+-----------+ + # | map | plot | + # | | | + # +-----------+-----------+ + # | info | + # +-----------------------+ + ui = widgets.VBox([instruction, + widgets.HBox([studyarea_map, fig_display]), + info]) + display(ui) diff --git a/deafrica_tools/app/deacoastlines.py b/deafrica_tools/app/deacoastlines.py new file mode 100644 index 0000000..2b00a48 --- /dev/null +++ b/deafrica_tools/app/deacoastlines.py @@ -0,0 +1,505 @@ +""" +Digital Earth Africa Coastline widget, which can be used to +interactively extract shoreline data using transects. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import fiona +import sys +import datacube +import warnings +import matplotlib.pyplot as plt +from datacube.utils.geometry import CRS +from ipyleaflet import ( + WMSLayer, + basemaps, + basemap_to_tiles, + Map, + DrawControl, + WidgetControl, + LayerGroup, + LayersControl, + GeoData, +) +from traitlets import Unicode +from ipywidgets import ( + GridspecLayout, + Button, + Layout, + HBox, + VBox, + HTML, + Output, +) +import json +import geopandas as gpd +from io import BytesIO +import ipywidgets as widgets + +import deafrica_tools.app.widgetconstructors as deawidgets +from deafrica_tools.coastal import get_coastlines, transect_distances +from owslib.wms import WebMapService + +def make_box_layout(): + return Layout( + # border='solid 1px black', + margin='0px 10px 10px 0px', + padding='5px 5px 5px 5px', + width='100%', + height='100%', + ) + + +def create_expanded_button(description, button_style): + return Button( + description=description, + button_style=button_style, + layout=Layout(width="auto", height="auto"), + ) + + +class transect_app(HBox): + + def __init__(self): + super().__init__() + + ###################### + # INITIAL ATTRIBUTES # + ###################### + + self.output_name = "example_output" + self.export_csv = False + self.export_plot = False + self.product_list = [ + ("ESRI World Imagery", "none"), + ("Open Street Map", "open_street_map"), + ] + self.product = self.product_list[0][1] + self.mode_list = [('Distance', 'distance'), ('Width', 'width')] + self.mode = self.mode_list[0][1] + self.target = None + self.action = None + self.gdf_drawn = None + self.gdf_uploaded = None + + ################## + # HEADER FOR APP # + ################## + + # Create the Header widget + header_title_text = "

Digital Earth Africa Coastlines shoreline transect extraction

" + instruction_text = "Select parameters and draw a transect on the map to extract shoreline data. In distance mode, draw a transect line starting from land that crosses multiple shorelines.
In width mode, draw a transect line that intersects shorelines at least twice. Alternatively, upload an vector file to extract shoreline data for multiple existing transects." + self.header = deawidgets.create_html( + f"{header_title_text}

{instruction_text}

") + self.header.layout = make_box_layout() + + ##################################### + # HANDLER FUNCTION FOR DRAW CONTROL # + ##################################### + + # Define the action to take once something is drawn on the map + def update_geojson(target, action, geo_json): + + # Remove previously uploaded data if present + self.gdf_uploaded = None + fileupload_transects._counter = 0 + + # Get data from action + self.action = action + + # Convert data to geopandas + json_data = json.dumps(geo_json) + binary_data = json_data.encode() + io = BytesIO(binary_data) + io.seek(0) + gdf = gpd.read_file(io) + gdf.crs = "EPSG:4326" + + # Convert to WGS 84 / NSIDC EASE-Grid 2.0 Global and compute area + gdf_drawn_nsidc = gdf.copy().to_crs("EPSG:6933") + m2_per_km2 = 10**6 + area = gdf_drawn_nsidc.envelope.area.values[0] / m2_per_km2 + polyarea_label = 'Total area of DE Africa Coastlines data to extract' + polyarea_text = f"{polyarea_label}: {area:.2f} km2" + + # Test area size + if area <= 50000: + confirmation_text = ' (Area to extract falls within recommended limit; click "Extract shoreline data" to continue)' + self.header.value = header_title_text + polyarea_text + confirmation_text + self.gdf_drawn = gdf + else: + warning_text = ' (Area to extract is too large, please select a smaller transect)' + self.header.value = header_title_text + polyarea_text + warning_text + self.gdf_drawn = None + + ########################### + # WIDGETS FOR APP OUTPUTS # + ########################### + + self.status_info = Output(layout=make_box_layout()) + self.output_plot = Output(layout=make_box_layout()) + + ######################################### + # MAP WIDGET, DRAWING TOOLS, WMS LAYERS # + ######################################### + + # Create drawing tools + desired_drawtools = ['polyline'] + draw_control = deawidgets.create_drawcontrol(desired_drawtools) + + # Load DEACoastLines WMS + deacl_url = "https://geoserver.digitalearth.africa/geoserver/wms" + deacl_layer = "coastlines:DEAfrica_Coastlines" + deacoastlines = WMSLayer( + url=deacl_url, + layers=deacl_layer, + format='image/png', + transparent=True, + attribution='DE Africa Coastlines © 2022 Digital Earth Africa') + + # Begin by displaying an empty layer group, and update the group with desired WMS on interaction. + self.map_layers = LayerGroup(layers=(deacoastlines,)) + self.map_layers.name = 'Map Overlays' + + # Create map widget + self.m = deawidgets.create_map(map_center=(0.5273, 25.1367), + zoom_level=3, + basemap=basemaps.Esri.WorldImagery) + self.m.layout = make_box_layout() + + # Add tools to map widget + self.m.add_control(draw_control) + self.m.add_layer(self.map_layers) + + # Store current basemap for future use + self.basemap = self.m.basemap + + ############################ + # WIDGETS FOR APP CONTROLS # + ############################ + + # Create parameter widgets + text_output_name = deawidgets.create_inputtext(self.output_name, + self.output_name) + checkbox_csv = deawidgets.create_checkbox(self.export_csv, + 'Distance table (.csv)') + checkbox_plot = deawidgets.create_checkbox(self.export_plot, + 'Figure (.png)') + deaoverlay_dropdown = deawidgets.create_dropdown( + self.product_list, self.product_list[0][1]) + mode_dropdown = deawidgets.create_dropdown(self.mode_list, + self.mode_list[0][1]) + run_button = create_expanded_button("Extract shoreline data", "info") + fileupload_transects = widgets.FileUpload(accept='', multiple=True) + + #################################### + # UPDATE FUNCTIONS FOR EACH WIDGET # + #################################### + + # Run update functions whenever various widgets are changed. + text_output_name.observe(self.update_text_output_name, "value") + checkbox_csv.observe(self.update_checkbox_csv, "value") + checkbox_plot.observe(self.update_checkbox_plot, "value") + deaoverlay_dropdown.observe(self.update_deaoverlay, "value") + mode_dropdown.observe(self.update_mode, "value") + run_button.on_click(self.run_app) + draw_control.on_draw(update_geojson) + fileupload_transects.observe(self.update_fileupload_transects, "value") + + ################################## + # COLLECTION OF ALL APP CONTROLS # + ################################## + + parameter_selection = VBox([ + HTML("Output name:"), text_output_name, + HTML( + 'Transect extraction mode:
' + ), + mode_dropdown, + HTML("
Output files:
"), + checkbox_plot, + checkbox_csv, + HTML( + "
Advanced
Upload a GeoJSON or ESRI " + "Shapefile (<5 mb) containing one or more transect lines.
"), + fileupload_transects + ]) + map_selection = VBox([ + HTML("
Map overlay:"), + deaoverlay_dropdown, + ]) + parameter_selection.layout = make_box_layout() + map_selection.layout = make_box_layout() + + ############################### + # SPECIFICATION OF APP LAYOUT # + ############################### + + # 0 1 2 3 4 5 6 7 8 9 + # --------------------------------------------- + # 0 | Header | Map sel. | + # --------------------------------------------- + # 1 | Params | | + # 2 | | | + # 3 | | | + # 4 | | Map | + # 5 | | | + # ---------- | + # 6 | Run | | + # --------------------------------------------- + # 7 | Status info | + # --------------------------------------------- + # 8 | | + # 9 | Output/figure | + # 10 | | + # 11 | ------------------------------------------| + + # Create the layout #[rowspan, colspan] + grid = GridspecLayout(12, 10, height="1350px", width="auto") + + # Header and controls + grid[0, :8] = self.header + grid[0, 8:] = map_selection + grid[1:6, 0:2] = parameter_selection + grid[6, 0:2] = run_button + + # Status info, map and plot + grid[1:7, 2:] = self.m # map + grid[7:8, :] = self.status_info + grid[8:, :] = self.output_plot + + # Display using HBox children attribute + self.children = [grid] + + ###################################### + # DEFINITION OF ALL UPDATE FUNCTIONS # + ###################################### + + # Set the output csv + def update_fileupload_transects(self, change): + + # Clear any drawn data if present + self.gdf_drawn = None + + # Save to file + for uploaded_filename in change.new.keys(): + with open(uploaded_filename, "wb") as output_file: + content = change.new[uploaded_filename]['content'] + output_file.write(content) + + with self.status_info: + + try: + + print('Loading vector data...', end='\r') + valid_files = [ + file for file in change.new.keys() + if file.lower().endswith(('.shp', '.geojson')) + ] + valid_file = valid_files[0] + transect_gdf = (gpd.read_file(valid_file).to_crs( + "EPSG:4326").explode().reset_index(drop=True)) + + # Use ID column if it exists + if 'id' in transect_gdf: + transect_gdf = transect_gdf.set_index('id') + print(f"Uploaded '{valid_file}'; automatically labelling " + "transects using column 'id'.") + else: + print( + f"Uploaded '{valid_file}'; no 'id' column detected, " + f"labelling transects from 0 to {len(transect_gdf.index) - 1}." + ) + + # Create a geodata + geodata = GeoData(geo_dataframe=transect_gdf, + style={ + 'color': 'black', + 'weight': 3 + }) + + # Add to map + xmin, ymin, xmax, ymax = transect_gdf.total_bounds + self.m.fit_bounds([[ymin, xmin], [ymax, xmax]]) + self.m.add_layer(geodata) + + # If completed, add to attribute + self.gdf_uploaded = transect_gdf + + except IndexError: + print( + "Cannot read uploaded files. Please ensure that data is " + "in either GeoJSON or ESRI Shapefile format.", + end='\r') + self.gdf_uploaded = None + + except fiona.errors.DriverError: + print( + "Shapefile is invalid. Please ensure that all shapefile " + "components (e.g. .shp, .shx, .dbf, .prj) are uploaded.", + end='\r') + self.gdf_uploaded = None + + # Set output name + def update_text_output_name(self, change): + self.output_name = change.new + + # Output CSV + def update_checkbox_csv(self, change): + self.export_csv = change.new + + # Output plot + def update_checkbox_plot(self, change): + self.export_plot = change.new + + # Set mode + def update_mode(self, change): + self.mode = change.new + + # Update product + def update_deaoverlay(self, change): + + self.product = change.new + + # Load DE Africa CoastLines WMS + deacl_url = "https://geoserver.digitalearth.africa/geoserver/wms" + deacl_layer = "coastlines:DEAfrica_Coastlines" + deacoastlines = WMSLayer( + url=deacl_url, + layers=deacl_layer, + format="image/png", + transparent=True, + attribution="DE Africa Coastlines © 2022 Digital Earth Africa") + + if self.product == "none": + self.map_layers.clear_layers() + self.map_layers.add_layer(deacoastlines) + + elif self.product == "open_street_map": + self.map_layers.clear_layers() + layer = basemap_to_tiles(basemaps.OpenStreetMap.Mapnik) + self.map_layers.add_layer(layer) + self.map_layers.add_layer(deacoastlines) + + def run_app(self, change): + + # Clear progress bar and output areas before running + self.status_info.clear_output() + self.output_plot.clear_output() + + # Run DE Africa Coastlines analysis + with self.status_info: + warnings.filterwarnings("ignore") + + # Load transects from either map or uploaded files + if self.gdf_uploaded is not None: + transect_gdf = self.gdf_uploaded + run_text = 'uploaded file' + elif self.gdf_drawn is not None: + transect_gdf = self.gdf_drawn + transect_gdf.index = [self.output_name] + run_text = 'selected transect' + else: + print(f'No transect drawn or uploaded. Please select a transect on the map, or upload a GeoJSON or ESRI Shapefile.', + end='\r') + transect_gdf = None + + # If valid data was returned, load DEA Coastlines data + if transect_gdf is not None: + + # Load Coastlines data from WFS + deacl_gdf = get_coastlines(bbox=transect_gdf) + + # Test that data was correctly returned + if len(deacl_gdf.index) > 0: + + # Dissolve by year to remove duplicates, then sort by date + deacl_gdf = deacl_gdf.dissolve(by='year', as_index=False) + deacl_gdf['year'] = deacl_gdf.year.astype(int) + deacl_gdf = deacl_gdf.sort_values('year') + deacl_gdf = deacl_gdf.set_index('year') + + else: + print( + "No annual shoreline data was found near the " + "supplied transect. Please draw or select a new " + "transect.", + end='\r') + deacl_gdf = None + + # If valid DEA Coastlines data returned, calculate distances + if deacl_gdf is not None: + print(f'Analysing transect distances using "{self.mode}" mode...', + end='\r') + dist_df = transect_distances( + transect_gdf.to_crs("EPSG:6933"), + deacl_gdf.to_crs("EPSG:6933"), + mode=self.mode) + + # If valid data was produced: + if dist_df.any(axis=None): + + # Successful output + print(f'DE Africa Coastlines data successfully extracted for {run_text}.') + + # Export distance data + if self.export_csv: + + # Create folder if required and set path + out_dir = 'deacoastlines_outputs' + os.makedirs(out_dir, exist_ok=True) + csv_filename = f"{out_dir}/{self.output_name}.csv" + + # Export to file + dist_df.to_csv(csv_filename, index_label="Transect") + print(f'Distance data exported to "{csv_filename}".') + + # Generate plot + with self.output_plot: + + fig, ax = plt.subplots(constrained_layout=True, + figsize=(15, 5.5)) + dist_df.T.plot(ax=ax, linewidth=3) + + ax.legend(frameon=False, ncol=3, title='Transect') + ax.set_title(f"Digital Earth Africa Coastlines transect extraction - {self.output_name}") + ax.set_ylabel(f"Along-transect {self.mode} (m)") + ax.set_xlim(dist_df.T.index[0], dist_df.T.index[-1]) + + # Hide the right and top spines + ax.spines['right'].set_visible(False) + ax.spines['top'].set_visible(False) + + # Only show ticks on the left and bottom spines + ax.yaxis.set_ticks_position('left') + ax.xaxis.set_ticks_position('bottom') + plt.show() + + # Export plot + with self.status_info: + if self.export_plot: + + # Create folder if required and set path + out_dir = 'deacoastlines_outputs' + os.makedirs(out_dir, exist_ok=True) + figure_filename = f"{out_dir}/{self.output_name}.png" + + # Export to file + fig.savefig(figure_filename) + print(f'Figure exported to "{figure_filename}".') + + else: + print( + "No valid shoreline data intersects with the " + "supplied transect. This can occur if:\n\n" + " - the transect does not intersect with any shorelines\n" + " - the transect intersects with shorelines more than once in 'distance' mode\n" + " - the transect intersects with shorelines only once in 'width' mode\n\n" + "Please draw or upload a new transect.", + end='\r') \ No newline at end of file diff --git a/deafrica_tools/app/forestmonitoring.py b/deafrica_tools/app/forestmonitoring.py new file mode 100644 index 0000000..4e99243 --- /dev/null +++ b/deafrica_tools/app/forestmonitoring.py @@ -0,0 +1,1026 @@ +''' +Functions for loading and interacting with Global Forest Change data in the forest monitoring notebook, inside the Real_world_examples folder. +''' + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import json +import warnings +from io import BytesIO + +import deafrica_tools.app.widgetconstructors as deawidgets +import geopandas as gpd +import ipywidgets as widgets +import matplotlib.colors as mcolors +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import rioxarray +import xarray as xr +from deafrica_tools.dask import create_local_dask_cluster +from deafrica_tools.spatial import xr_rasterize +from ipyleaflet import ( + DrawControl, + GeoData, + LayerGroup, + LayersControl, + Map, + WidgetControl, + WMSLayer, + basemap_to_tiles, + basemaps, +) +from ipywidgets import HTML, Button, GridspecLayout, HBox, Layout, Output, VBox +from matplotlib.patches import Patch +from traitlets import Unicode + +# Turn off all warnings. +warnings.filterwarnings("ignore") +warnings.simplefilter("ignore") + +def make_box_layout(): + """ + Defines a number of CSS properties that impact how a widget is laid out. + """ + return Layout( # border='solid 1px black', + margin="0px 10px 10px 0px", + padding="5px 5px 5px 5px", + width="100%", + height="100%", + ) + +def create_expanded_button(description, button_style): + """ + Defines a number of CSS properties to create a button to handle mouse clicks. + """ + return Button( + description=description, + button_style=button_style, + layout=Layout(width="auto", height="auto"), + ) + +def load_gfclayer(gdf_drawn, gfclayer): + """ + Loads the selected Global Forest Change layer for the + area drawn on the map widget. + """ + # Configure local dask cluster. + client = create_local_dask_cluster(return_client=True, display_client=True) + + # Get the coordinates of the top-left corner for each Global Forest Change tile, + # covering the area of interest. + min_lat, max_lat = ( + gdf_drawn.bounds.miny.item(), + gdf_drawn.bounds.maxy.item(), + ) + min_lon, max_lon = ( + gdf_drawn.bounds.minx.item(), + gdf_drawn.bounds.maxx.item(), + ) + + lats = np.arange( + np.floor(min_lat / 10) * 10, np.ceil(max_lat / 10) * 10, 10 + ).astype(int) + lons = np.arange( + np.floor(min_lon / 10) * 10, np.ceil(max_lon / 10) * 10, 10 + ).astype(int) + + coord_list = [] + for lat in lats: + lat = lat + 10 + if lat >= 0: + lat_str = f"{lat:02d}N" + else: + lat_str = f"{abs(lat):02d}S" + for lon in lons: + if lon >= 0: + lon_str = f"{lon:03d}E" + else: + lon_str = f"{abs(lon):03d}W" + coord_str = f"{lat_str}_{lon_str}" + coord_list.append(coord_str) + + # Load each Global Forest Change tile covering the area of interest. + base_url = f"https://storage.googleapis.com/earthenginepartners-hansen/GFC-2021-v1.9/Hansen_GFC-2021-v1.9_{gfclayer}_" + dask_chunks = dict(x=2048, y=2048) + + tile_list = [] + for coord in coord_list: + tile_url = f"{base_url}{coord}.tif" + # Load the tile as an xarray.DataArray. + tile = rioxarray.open_rasterio(tile_url, chunks=dask_chunks).squeeze() + tile_list.append(tile) + + # Merge the tiles into a single xarray.DataArray. + ds = xr.combine_by_coords(tile_list) + # Clip the dataset using the bounds of the area of interest. + ds = ds.rio.clip_box( + minx=min_lon - 0.00025, + miny=min_lat - 0.00025, + maxx=max_lon + 0.00025, + maxy=max_lat + 0.00025, + ) + # Rename the y and x variables for DEA convention on xarray.DataArrays where crs="EPSG:4326". + ds = ds.rename({"y": "latitude", "x": "longitude"}) + + # Mask pixels representing no loss (encoded as 0) in the "lossyear" layer. + if gfclayer == "lossyear": + ds = ds.where(ds != 0) + # Mask pixels representing no gain (encoded as 0) in the "gain" layer. + elif gfclayer == "gain": + ds = ds.where(ds != 0) + # Mask pixels with 0 percentage tree canopy cover. + elif gfclayer == "treecover2000": + ds = ds.where(ds != 0) + + # Create a mask from the area of interest GeoDataFrame. + mask = xr_rasterize(gdf_drawn, ds) + # Mask the dataset. + ds = ds.where(mask) + # Convert the xarray.DataArray to a dataset. + ds = ds.to_dataset(name=gfclayer) + # Compute. + ds = ds.compute() + # Assign the "EPSG:4326" CRS to the dataset. + ds.rio.write_crs(4326, inplace=True) + ds = ds.transpose("latitude", "longitude") + + # Close down the dask client. + client.close() + return ds + +def load_all_gfclayers(gdf_drawn): + gfclayers = ["treecover2000", "gain", "lossyear"] + + dataset_list = [] + for layer in gfclayers: + ds = load_gfclayer(gdf_drawn, gfclayer=layer) + dataset_list.append(ds) + + dataset = xr.merge(dataset_list) + return dataset + +def get_gfclayer_treecover2000(gfclayer_ds, gfclayer="treecover2000"): + """ + Preprocess the Global Forest change "treecover2020" layer. + """ + ds = gfclayer_ds[gfclayer] + + # Check if the dataarray is empty. + condition = ds.isnull().all().item() + + if condition: + return None + else: + # Mask the dataset. + mask = np.isnan(ds) + ds_masked = ds.where(mask, 1) + + # Get the pixel count for each unique pixel value in the layer. + counts = np.unique(ds_masked, return_counts=True) + # Remove the counts for pixels with the value np.nan. + index = np.argwhere(np.isnan(counts[0])) + counts_dict = dict( + zip(np.delete(counts[0], index), np.delete(counts[1], index)) + ) + + # Reproject the dataset to EPSG:6933 which uses metres + ds_reprojected = ds_masked.rio.reproject("EPSG:6933") + # Get the area per pixel. + pixel_length = ds_reprojected.geobox.resolution[1] + m_per_km = 1000 + per_pixel_area = (pixel_length / m_per_km) ** 2 + + # Save the results as a pandas DataFrame. + df = pd.DataFrame( + data={ + "Year": ["2000"], + "Tree Cover in km$^2$": np.fromiter(counts_dict.values(), dtype=float) + * per_pixel_area, + } + ) + + # Get the total area. + print_statement = f'Total Forest Cover in {df["Year"].item()}: {round(df["Tree Cover in km$^2$"].item(), 4)} km2' + + # File name to use when exporting results. + file_name = f"forest_cover_in_2000" + + return ds, df, print_statement, file_name + +def get_gfclayer_gain(gfclayer_ds, gfclayer="gain"): + """ + Preprocess the Global Forest Change "gain" layer. + """ + ds = gfclayer_ds[gfclayer] + + # Check if the dataarray is empty. + condition = ds.isnull().all().item() + + if condition: + return None + else: + # Get the pixel count for each unique pixel value in the layer. + counts = np.unique(ds, return_counts=True) + # Remove the counts for pixels with the value np.nan. + index = np.argwhere(np.isnan(counts[0])) + counts_dict = dict( + zip(np.delete(counts[0], index), np.delete(counts[1], index)) + ) + + # Reproject the dataset to EPSG:6933 which uses metres. + ds_reprojected = ds.rio.reproject("EPSG:6933") + # Get the area per pixel. + pixel_length = ds_reprojected.geobox.resolution[1] + m_per_km = 1000 + per_pixel_area = (pixel_length / m_per_km) ** 2 + + # Save the results as a pandas DataFrame. + df = pd.DataFrame( + data={ + "Year": ["2000-2012"], + "Forest Cover Gain in km$^2$": np.fromiter( + counts_dict.values(), dtype=float + ) + * per_pixel_area, + } + ) + + # Get the total area. + print_statement = f'Total Forest Cover Gain {df["Year"].item()}: {round(df["Forest Cover Gain in km$^2$"].item(), 4)} km2' + + # File name to use when exporting results. + file_name = f"forest_cover_gain_from_2000_to_2012" + + return ds, df, print_statement, file_name + +def get_gfclayer_lossyear(gfclayer_ds, start_year, end_year, gfclayer="lossyear"): + """ + Preprocess the Global Forest Change "lossyear" layer. + """ + + ds = gfclayer_ds[gfclayer] + + # Mask the dataset to the selected time range. + selected_years = list(range(start_year, end_year + 1)) + mask = ds.isin(selected_years) + ds = ds.where(mask) + + # Check if the dataarray is empty. + condition = ds.isnull().all().item() + + if condition: + return None + else: + # Get the pixel count for each unique pixel value in the layer. + counts = np.unique(ds, return_counts=True) + # Remove the counts for pixels with the value np.nan. + index = np.argwhere(np.isnan(counts[0])) + counts_dict = dict( + zip(np.delete(counts[0], index), np.delete(counts[1], index)) + ) + + # Reproject the dataset to EPSG:6933 which uses metres + ds_reprojected = ds.rio.reproject("EPSG:6933") + # Get the area per pixel. + pixel_length = ds_reprojected.geobox.resolution[1] + m_per_km = 1000 + per_pixel_area = (pixel_length / m_per_km) ** 2 + + # For each year get the area of loss. + # Save the results as a pandas DataFrame. + df = pd.DataFrame( + { + "Year": 2000 + np.fromiter(counts_dict.keys(), dtype=int), + "Forest Cover Loss in km$^2$": np.fromiter( + counts_dict.values(), dtype=float + ) + * per_pixel_area, + } + ) + + # Get the total area. + print_statement = f'Total Forest Cover Loss from {start_year + 2000} to {end_year + 2000}: {round(df["Forest Cover Loss in km$^2$"].sum(), 4)} km2' + + # File name to use when exporting results. + file_name = f"forest_cover_loss_from_{start_year + 2000}_to_{end_year + 2000}" + + return ds, df, print_statement, file_name + +def plot_gfclayer_treecover2000(gfclayer_ds, gfclayer="treecover2000"): + """ + Plot the Global Forest Change "treecover2000" layer. + """ + + if get_gfclayer_treecover2000(gfclayer_ds) is None: + print( + f"No Global Forest Change {gfclayer} layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ds, df, print_statement, file_name = get_gfclayer_treecover2000(gfclayer_ds) + + # Export the dataframe as a csv. + df.to_csv(f"{file_name}.csv", index=False) + print(f'Table exported to "{file_name}.csv"') + + # Define the plotting parameters. + figure_width = 10 + figure_length = 10 + title = f"Tree Canopy Cover for the Year 2000" + + # Plot the dataset. + fig, ax = plt.subplots(figsize=(figure_width, figure_length)) + im = ds.plot(cmap="Greens", add_colorbar=False, ax=ax) + # Add a colorbar to the plot. + cbar = plt.colorbar(mappable=im) + cbar.set_label( + "Percentage tree canopy cover for year 2000", labelpad=-65, y=0.25 + ) + # Add a title to the plot. + plt.title(title) + # Save the plot. + plt.savefig(f"{file_name}.png") + print(f'Figure exported to "{file_name}.png"') + plt.show() + + print(print_statement) + +def plot_gfclayer_gain(gfclayer_ds, gfclayer="gain"): + """ + Plot the Global Forest Change "gain" layer. + """ + + if get_gfclayer_gain(gfclayer_ds) is None: + print( + f"No Global Forest Change {gfclayer} layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ds, df, print_statement, file_name = get_gfclayer_gain(gfclayer_ds) + + # Export the dataframe as a csv. + df.to_csv(f"{file_name}.csv", index=False) + print(f'Table exported to "{file_name}.csv"') + + # Define the plotting parameters. + color = "#6CAE75" + figure_width = 10 + figure_length = 10 + title = f"Forest Cover Gain from 2000 to 2012" + + # Plot the dataset. + fig, ax = plt.subplots(figsize=(figure_width, figure_length)) + im = ds.plot(cmap=mcolors.ListedColormap([color]), add_colorbar=False, ax=ax) + # Add a legend to the plot. + im.axes.legend( + [Patch(facecolor=color)], + ["Global forest cover gain 2000–2012"], + loc="lower left", + bbox_to_anchor=(1.0, 0.5), + frameon=False, + ) + # Add a title to the plot. + plt.title(title) + # Save the plot. + plt.savefig(f"{file_name}.png") + print(f'Figure exported to "{file_name}.png"') + plt.show() + + print(print_statement) + +def plot_gfclayer_lossyear(gfclayer_ds, start_year, end_year, gfclayer="lossyear"): + """ + Plot the Global Forest change "lossyear" layer. + """ + + if ( + get_gfclayer_lossyear(gfclayer_ds, start_year, end_year, gfclayer="lossyear") + is None + ): + print( + f"No Global Forest Change {gfclayer} layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ds, df, print_statement, file_name = get_gfclayer_lossyear( + gfclayer_ds, start_year, end_year, gfclayer="lossyear" + ) + + # Export the dataframe as a csv. + df.to_csv(f"{file_name}.csv", index=False) + print(f'Table exported to "{file_name}.csv"') + + # Define the plotting parameters. + figure_width = 10 + figure_length = 15 + nrows = 2 + ncols = 1 + title = f"Forest Cover Loss from {start_year + 2000} to {end_year + 2000}" + + # Location of transition from one color to the next on the colormap. + color_levels = list(np.arange(1 - 0.5, 22, 1)) + # Ticks to be displayed. + ticks = list(np.arange(1, 22)) + tick_labels = list(2000 + np.arange(1, 22)) + + # Define the color map to use when plotting. + color_list = [ + "#e6194b", + "#3cb44b", + "#ffe119", + "#4363d8", + "#f58231", + "#911eb4", + "#46f0f0", + "#f032e6", + "#bcf60c", + "#fabebe", + "#008080", + "#e6beff", + "#9a6324", + "#fffac8", + "#800000", + "#aaffc3", + "#808000", + "#ffd8b1", + "#000075", + "#808080", + "#7A306C", + ] + cmap = mcolors.ListedColormap(colors=color_list, N=21) + norm = mcolors.BoundaryNorm(boundaries=color_levels, ncolors=cmap.N) + + # Plot the dataset. + fig, (ax1, ax2) = plt.subplots( + nrows, ncols, figsize=(figure_width, figure_length) + ) + im = ds.plot(ax=ax1, cmap=cmap, norm=norm, add_colorbar=False) + # Add a title to the subplot. + ax1.set_title(title) + # Add a colorbar to the subplot. + cbar = plt.colorbar(mappable=im, ticks=ticks) + cbar.set_label("Year of gross forest cover loss event", labelpad=-60, y=0.25) + cbar.set_ticklabels(tick_labels) + # Plot the second subplot. + df.plot( + x="Year", + y="Forest Cover Loss in km$^2$", + ylabel="Forest Cover Loss in km$^2$", + title=title, + ax=ax2, + ) + # Save the plot. + plt.savefig(f"{file_name}.png") + print(f'Figure exported to "{file_name}.png"') + plt.show() + + print(print_statement) + +def plot_gfclayer_all(gfclayer_ds, start_year, end_year): + """ + Plot all the Global Forest Change Layers loaded. + """ + + # Define the plotting parameters. + figure_width = 10 + figure_length = 10 + treecover_color = "Greens" + gain_color = "yellow" + lossyear_color = "red" + + print_statement_list = [] + filename_list = ["\nTables exported as: "] + + figure_fn = "global_forest_change_all_layers.png" + + # Define the figure. + fig, ax = plt.subplots(figsize=(figure_width, figure_length)) + if ( + get_gfclayer_treecover2000( + gfclayer_ds[["treecover2000"]], gfclayer="treecover2000" + ) + is None + ): + print( + f"No Global Forest Change 'treecover2000' layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ( + ds_treecover2000, + df_treecover2000, + print_statement_treecover2000, + file_name_treecover2000, + ) = get_gfclayer_treecover2000( + gfclayer_ds[["treecover2000"]], gfclayer="treecover2000" + ) + # Plot the treecover2000 layer as the background layer. + background = ds_treecover2000.plot( + cmap=treecover_color, add_colorbar=False, ax=ax + ) + # Add a colorbar to the treecover2000 plot. + cbar = plt.colorbar(mappable=background) + cbar.set_label( + "Percentage tree canopy cover for year 2000", labelpad=-65, y=0.25 + ) + # Export the dataframe as a csv. + df_treecover2000.to_csv(f"{file_name_treecover2000}.csv", index=False) + # Add the print statement to the list. + print_statement_list.append(print_statement_treecover2000) + # Add the file name to the list. + filename_list.append(f'"{file_name_treecover2000}.csv"') + + if get_gfclayer_gain(gfclayer_ds[["gain"]], gfclayer="gain") is None: + print( + f"No Global Forest Change 'gain' layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ds_gain, df_gain, print_statement_gain, file_name_gain = get_gfclayer_gain( + gfclayer_ds[["gain"]], gfclayer="gain" + ) + # Plot the gain layer. + ds_gain.plot( + ax=ax, cmap=mcolors.ListedColormap([gain_color]), add_colorbar=False + ) + # Export the dataframe as a csv. + df_gain.to_csv(f"{file_name_gain}.csv", index=False) + # Add the print statement to the list. + print_statement_list.append(print_statement_gain) + # Add the file name to the list. + filename_list.append(f'"{file_name_gain}.csv"') + + if ( + get_gfclayer_lossyear( + gfclayer_ds[["lossyear"]], start_year, end_year, gfclayer="lossyear" + ) + is None + ): + print( + f"No Global Forest Change 'lossyear' layer data found in the selected area. Please select a new polygon over an area with data." + ) + else: + ( + ds_lossyear, + df_lossyear, + print_statement_lossyear, + file_name_lossyear, + ) = get_gfclayer_lossyear( + gfclayer_ds[["lossyear"]], start_year, end_year, gfclayer="lossyear" + ) + # Plot the lossyear layer. + ds_lossyear.plot( + ax=ax, cmap=mcolors.ListedColormap([lossyear_color]), add_colorbar=False + ) + # Export the dataframe as a csv. + df_lossyear.to_csv(f"{file_name_lossyear}.csv", index=False) + # Add the print statement to the list. + print_statement_list.append(print_statement_lossyear) + # Add the file name to the list. + filename_list.append(f'"{file_name_lossyear}.csv"') + + # Add a legend to the plot. + ax.legend( + [Patch(facecolor=gain_color), Patch(facecolor=lossyear_color)], + [ + "Global forest cover \n gain 2000–2012", + f"Global forest cover \n loss {str(2000+start_year)}-{str(2000+end_year)}", + ], + loc="lower right", + bbox_to_anchor=(-0.1, 0.75), + frameon=False, + ) + + plt.title("Global Forest Change Layers") + plt.savefig(figure_fn) + plt.show() + print(*print_statement_list, sep="\n") + print(*filename_list, sep="\n\t") + print(f'\nFigure saved as "{figure_fn}"'); + +def plot_gfclayer(gfclayer_ds, start_year, end_year, gfclayer): + if gfclayer == "treecover2000": + plot_gfclayer_treecover2000(gfclayer_ds, gfclayer) + elif gfclayer == "lossyear": + plot_gfclayer_lossyear(gfclayer_ds, start_year, end_year, gfclayer) + elif gfclayer == "gain": + plot_gfclayer_gain(gfclayer_ds, gfclayer) + elif gfclayer == "alllayers": + plot_gfclayer_all(gfclayer_ds, start_year, end_year) + +def update_map_layers(self): + """ + Updates map widget to add new basemap when selected + using menu options. + """ + # Clear data load parameters to trigger data reload. + self.gfclayer_ds = None + + # Remove all layers from the map_layers Layers Group. + self.map_layers.clear_layers() + # Add the selected basemap to the layer Group. + self.map_layers.add_layer(self.basemap) + +class forest_monitoring_app(HBox): + def __init__(self): + super().__init__() + + ################## + # HEADER FOR APP # + ################## + + # Create the header widget. + header_title_text = "

Digital Earth Africa Forest Change

" + instruction_text = """

Select the desired Global Forest Change layer, then zoom in and draw a polygon to + select an area for which to plot the selected Global Forest Change layer. Alternatively, upload a vector file of the area of interest.

""" + self.header = deawidgets.create_html( + value=f"{header_title_text}{instruction_text}" + ) + self.header.layout = make_box_layout() + + ############################ + # WIDGETS FOR APP CONTROLS # + ############################ + + ## Selection widget for selecting the basemap to use for the map widget. + ## and when plotting the Global Forest Change Layer. + # Basemaps available for selection for the map widget. + self.basemap_list = [ + ("Open Street Map", basemap_to_tiles(basemaps.OpenStreetMap.Mapnik)), + ("ESRI World Imagery", basemap_to_tiles(basemaps.Esri.WorldImagery)), + ] + # Set the default basemap to be used for the map widget / initial value for the widget. + self.basemap = self.basemap_list[0][1] + # Dropdown selection widget. + dropdown_basemap = deawidgets.create_dropdown( + options=self.basemap_list, value=self.basemap + ) + # Register the update function to run when a new value is selected + # on the dropdown_basemap widget. + dropdown_basemap.observe(self.update_basemap, "value") + # Text to accompany the dropdown selection widget. + basemap_selection_html = deawidgets.create_html( + value=f"
Map overlay:" + ) + # Combine the basemap_selection_html text and the dropdown_basemap widget in a single container. + basemap_selection = VBox([basemap_selection_html, dropdown_basemap]) + + ## Selection widget for selecting the Global Forest change layer to plot. + # Global Forest Change layers available plotting. + self.gfclayers_list = [ + ("Year of gross forest cover loss event", "lossyear"), + ("Global forest cover gain 2000–2012", "gain"), + ("Tree canopy cover for the year 2000", "treecover2000"), + ("All layers", "alllayers"), + ] + # Set the default GFC layer to be plotted / initial value for the widget. + self.gfclayer = self.gfclayers_list[0][1] + + ## Selection widget for the data time range. + # Set the default time range for which to load data for. + self.start_year = 1 + self.end_year = 21 + + # Create the time range selector. + time_range = list(range(self.start_year, self.end_year + 1)) + time_range_str = [str(2000 + i) for i in time_range] + + timerange_options = tuple(zip(time_range_str, time_range)) + timerange_selection_slide = widgets.SelectionRangeSlider( + options=timerange_options, + value=(self.start_year, self.end_year), + description="", + disabled=False, + ) + # Register the update function to run when a new value is selected on the slider. + timerange_selection_slide.observe(self.update_timerange, "value") + # Text to accompany the timerange_selection widget. + timerange_selection_html = deawidgets.create_html( + value=f"
Forest Cover Loss Time Range:" + ) + # Combine the timerange_selection_text and the timerange_selection_slide in a single container. + timerange_selection = VBox( + [timerange_selection_html, timerange_selection_slide] + ) + + # Set the initial parameter for the GFC layer dataset. + self.gfclayer_ds = None + # Dropdown selection widget. + dropdown_gfclayer = deawidgets.create_dropdown( + options=self.gfclayers_list, value=self.gfclayer + ) + # Register the update function to run when a new value is selected + # on the dropdown_gfclayer widget. + dropdown_gfclayer.observe(self.update_gfclayer, "value") + # Text to accompany the dropdown selection widget. + gfclayer_selection_html = deawidgets.create_html( + value=f"
Global Forest Change Layer:" + ) + # Combine the gfclayer_selection_html text and the dropdown_gfclayer widget in a single container. + gfclayer_selection = VBox([gfclayer_selection_html, dropdown_gfclayer]) + + ## Add a checkbox for whether to overide the limit to the size of polygon drawn on the + ## map widget. + # Initial value of the widget. + self.max_size = False + # CheckBox widget. + checkbox_max_size = deawidgets.create_checkbox( + value=self.max_size, description="Enable", layout={"width": "95%"} + ) + # Text to accompany the CheckBox widget. + checkbox_max_size_html = deawidgets.create_html( + value=f"""
Override maximum size limit: + (use with caution; may cause memory issues/crashes)""" + ) + # Register the update function to run when the checkbox is ticked. + # on the checkbox_max_size CheckBox + checkbox_max_size.observe(self.update_checkbox_max_size, "value") + # # Combine the checkbox_max_size_html text and the checkbox_max_size widget in a single container. + enable_max_size = VBox([checkbox_max_size_html, checkbox_max_size]) + + # Add widget to enable uploading a geojson or ESRI shapefile. + self.gdf_uploaded = None + fileupload_aoi = widgets.FileUpload(accept="", multiple=True) + # Register the update function to be called for the file upload. + fileupload_aoi.observe(self.update_fileupload_aoi, "value") + fileupload_html = deawidgets.create_html(value=f"""
Advanced
Upload a GeoJSON or ESRI Shapefile (<5 mb) containing a single area of interest.
""") + fileupload = VBox([fileupload_html, fileupload_aoi]) + + + ## Put the app controls widgets into a single container. + parameter_selection = VBox( + [ + basemap_selection, + gfclayer_selection, + timerange_selection, + enable_max_size, + fileupload + ] + ) + parameter_selection.layout = make_box_layout() + + ## Button to click to run the app. + run_button = create_expanded_button( + description="Generate plot", button_style="info" + ) + # Register the update function to be called when the run_button button + # is clicked. + run_button.on_click(self.run_app) + + + + ########################### + # WIDGETS FOR APP OUTPUTS # + ########################### + + self.status_info = Output(layout=make_box_layout()) + self.output_plot = Output(layout=make_box_layout()) + + ################################# + # MAP WIDGET WITH DRAWING TOOLS # + ################################# + + # Create the map widget. + self.m = deawidgets.create_map( + map_center=(-18.45, 28.93), + zoom_level=11, + ) + self.m.layout = make_box_layout() + + # Create an empty Layer Group. + self.map_layers = LayerGroup(layers=()) + # Name of the Layer Group layer. + self.map_layers.name = "Map Overlays" + # Add the empty Layer Group as a single layer to the map widget. + self.m.add_layer(self.map_layers) + + # Create the desired drawing tools. + desired_drawtools = ["rectangle", "polygon"] + draw_control = deawidgets.create_drawcontrol(desired_drawtools) + # Add drawing tools to the map widget. + self.m.add_control(draw_control) + # Set the initial parameters for the drawing tools. + self.target = None + self.action = None + self.gdf_drawn = None + + ##################################### + # HANDLER FUNCTION FOR DRAW CONTROL # + ##################################### + + def handle_draw(target, action, geo_json): + + """ + Defines the action to take once something is drawn on the + map widget. + """ + # Remove previously uploaded data if present + self.gdf_uploaded = None + fileupload_aoi._counter = 0 + + self.target = target + self.action = action + + # Clear data load parameters to trigger data reload. + self.gfclayer_ds = None + + # Convert the drawn polygon geojson to a GeoDataFrame. + json_data = json.dumps(geo_json) + binary_data = json_data.encode() + io = BytesIO(binary_data) + io.seek(0) + gdf = gpd.read_file(io) + gdf.crs = "EPSG:4326" + + # Convert the GeoDataFrame to WGS 84 / NSIDC EASE-Grid 2.0 Global and compute the area. + gdf_drawn_nsidc = gdf.copy().to_crs("EPSG:6933") + m2_per_ha = 10000 + area = gdf_drawn_nsidc.area.values[0] / m2_per_ha + + polyarea_label = ( + f"Total area of Global Forest Change {self.gfclayer} layer to load" + ) + polyarea_text = f"{polyarea_label}: {area:.2f} ha" + + # Test the size of the polygon drawn. + if self.max_size: + confirmation_text = """ + (Overriding maximum size limit; use with caution as may lead to memory issues)""" + self.header.value = ( + header_title_text + + instruction_text + + polyarea_text + + confirmation_text + ) + self.gdf_drawn = gdf + elif area <= 50000: + confirmation_text = """ + (Area to extract falls within + recommended 50000 ha limit)""" + self.header.value = ( + header_title_text + + instruction_text + + polyarea_text + + confirmation_text + ) + self.gdf_drawn = gdf + else: + warning_text = """ + (Area to extract is too large, + please select an area less than 50000 )""" + self.header.value = ( + header_title_text + instruction_text + polyarea_text + warning_text + ) + self.gdf_drawn = None + + # Register the handler for draw events. + draw_control.on_draw(handle_draw) + + ############################### + # SPECIFICATION OF APP LAYOUT # + ############################### + + # Create the app layout. + grid_rows = 12 + grid_columns = 11 + grid_height = "1500px" + grid_width = "auto" + grid = GridspecLayout( + grid_rows, grid_columns, height=grid_height, width=grid_width + ) + + # Place app widgets and components in app layout. + # [rows, columns] + grid[0, :] = self.header + grid[1:6, 0:4] = parameter_selection + grid[6, 0:4] = run_button + grid[7:, 0:4] = self.status_info + grid[6:, 4:] = self.output_plot + grid[1:6, 4:] = self.m + # Display using HBox children attribute + self.children = [grid] + + ###################################### + # DEFINITION OF ALL UPDATE FUNCTIONS # + ###################################### + + def update_basemap(self, change): + """ + Updates the basemap on the map widget based on the + selected value of the dropdown_basemap widget. + """ + self.basemap = change.new + self.output_plot_basemap = get_basemap(self.basemap.url) + update_map_layers(self) + + def update_gfclayer(self, change): + """ + Updates the Global Forest Change layer to be plotted + based on the selected value of the dropdown_gfclayer widget. + """ + self.gfclayer = change.new + + def update_timerange(self, change): + """Updates the time range of the data to be loaded""" + self.start_year = change.new[0] + self.end_year = change.new[1] + + def update_checkbox_max_size(self, change): + """ + Sets the value of self.max_size to True when the + checkbox_max_size CheckBox is checked. + """ + self.max_size = change.new + + def update_fileupload_aoi(self, change): + + # Clear any drawn data if present + self.gdf_drawn = None + + # Save to file + for uploaded_filename in change.new.keys(): + with open(uploaded_filename, "wb") as output_file: + content = change.new[uploaded_filename]['content'] + output_file.write(content) + + with self.status_info: + + try: + + print('Loading vector data...', end='\r') + valid_files = [ + file for file in change.new.keys() + if file.lower().endswith(('.shp', '.geojson')) + ] + valid_file = valid_files[0] + aoi_gdf = (gpd.read_file(valid_file).to_crs( + "EPSG:4326").explode().reset_index(drop=True)) + + # Create a geodata + geodata = GeoData(geo_dataframe=aoi_gdf, + style={ + 'color': 'black', + 'weight': 3 + }) + + # Add to map + xmin, ymin, xmax, ymax = aoi_gdf.total_bounds + self.m.fit_bounds([[ymin, xmin], [ymax, xmax]]) + self.m.add_layer(geodata) + + # If completed, add to attribute + self.gdf_uploaded = aoi_gdf + + except IndexError: + print( + "Cannot read uploaded files. Please ensure that data is " + "in either GeoJSON or ESRI Shapefile format.", + end='\r') + self.gdf_uploaded = None + + except fiona.errors.DriverError: + print( + "Shapefile is invalid. Please ensure that all shapefile " + "components (e.g. .shp, .shx, .dbf, .prj) are uploaded.", + end='\r') + self.gdf_uploaded = None + + def run_app(self, change): + + # Clear progress bar and output areas before running. + self.status_info.clear_output() + self.output_plot.clear_output() + + with self.status_info: + # Load the area of interest from the map or uploaded files. + if self.gdf_uploaded is not None: + aoi_gdf = self.gdf_uploaded + elif self.gdf_drawn is not None: + aoi_gdf = self.gdf_drawn + else: + print(f'No valid polygon drawn on the map or uploaded. Please draw a valid a transect on the map, or upload a GeoJSON or ESRI Shapefile.', + end='\r') + aoi_gdf = None + + # If valid area of interest data returned. Load the selected Global Forest Change data. + if aoi_gdf is not None: + + if self.gfclayer_ds is None: + if self.gfclayer != "alllayers": + self.gfclayer_ds = load_gfclayer(gdf_drawn=aoi_gdf, gfclayer=self.gfclayer) + else: + self.gfclayer_ds = load_all_gfclayers(gdf_drawn=aoi_gdf) + else: + print("Using previously loaded data") + + # Plot the selected Global Forest Change layer. + if self.gfclayer_ds is not None: + with self.output_plot: + plot_gfclayer(gfclayer_ds=self.gfclayer_ds, + start_year=self.start_year, + end_year=self.end_year, + gfclayer=self.gfclayer) + else: + with self.status_info: + print(f"No Global Forest Change {self.gfclayer} layer data found in the selected area. Please select a new polygon over an area with data.") \ No newline at end of file diff --git a/deafrica_tools/app/geomedian.py b/deafrica_tools/app/geomedian.py new file mode 100644 index 0000000..4c7f45e --- /dev/null +++ b/deafrica_tools/app/geomedian.py @@ -0,0 +1,126 @@ +""" +Geomedian widget: generates an interactive visualisation of +the geomedian summary statistic. +""" + +# Load modules +import ipywidgets as widgets +import matplotlib.pyplot as plt +from mpl_toolkits.mplot3d import Axes3D +import numpy as np +import xarray as xr +from odc.algo import xr_geomedian + +def run_app(): + + """ + An interactive app that allows users to visualise the difference between the median and geomedian time-series summary statistics. By modifying the red-green-blue values of three timesteps for a given pixel, the user changes the output summary statistics. + + This allows a visual representation of the difference through the output values, RGB colour, as well as showing values plotted as a vector on a 3-dimensional space. + + Last modified: December 2021 + """ + + # Define the red-green-blue sliders for timestep 1 + p1r = widgets.IntSlider(description='Red', max=255, value=58) + p1g = widgets.IntSlider(description='Green', max=255, value=153) + p1b = widgets.IntSlider(description='Blue', max=255, value=68) + + # Define the red-green-blue sliders for timestep 2 + p2r = widgets.IntSlider(description='Red', max=255, value=208) + p2g = widgets.IntSlider(description='Green', max=255, value=221) + p2b = widgets.IntSlider(description='Blue', max=255, value=203) + + # Define the red-green-blue sliders for timestep 3 + p3r = widgets.IntSlider(description='Red', max=255, value=202) + p3g = widgets.IntSlider(description='Green', max=255, value=82) + p3b = widgets.IntSlider(description='Blue', max=255, value=33) + + # Define the median calculation for the timesteps + def f(p1r, p1g, p1b, p2r, p2g, p2b, p3r, p3g, p3b): + print('Red Median = {}'.format(np.median([p1r, p2r, p3r]))) + print('Green Median = {}'.format(np.median([p1g, p2g, p3g]))) + print('Blue Median = {}'.format(np.median([p1b, p2b, p3b]))) + + # Define the geomedian calculation for the timesteps + def g(p1r, p1g, p1b, p2r, p2g, p2b, p3r, p3g, p3b): + print('Red Geomedian = {:.2f}'.format(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).red.values.ravel()[0])) + print('Green Geomedian = {:.2f}'.format(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).green.values.ravel()[0])) + print('Blue Geomedian = {:.2f}'.format(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).blue.values.ravel()[0])) + + # Define the Timestep 1 box colour + def h(p1r, p1g, p1b): + fig1, axes1 = plt.subplots(figsize=(2,2)) + fig1 = plt.imshow([[(p1r, p1g, p1b)]]) + axes1.set_title('Timestep 1') + axes1.axis('off') + plt.show(fig1) + + # Define the Timestep 2 box colour + def hh(p2r, p2g, p2b): + fig2, axes2 = plt.subplots(figsize=(2,2)) + fig2 = plt.imshow([[(p2r, p2g, p2b)]]) + axes2.set_title('Timestep 2') + axes2.axis('off') + plt.show(fig2) + + # Define the Timestep 3 box colour + def hhh(p3r, p3g, p3b): + fig3, axes3 = plt.subplots(figsize=(2,2)) + fig3 = plt.imshow([[(p3r, p3g, p3b)]]) + axes3.set_title('Timestep 3') + axes3.axis('off') + plt.show(fig3) + + # Define the Median RGB colour box + def i(p1r, p1g, p1b, p2r, p2g, p2b, p3r, p3g, p3b): + fig4, axes4 = plt.subplots(figsize=(3,3)) + fig4 = plt.imshow([[(int(np.median([p1r, p2r, p3r])), int(np.median([p1g, p2g, p3g])), int(np.median([p1b, p2b, p3b])))]]) + axes4.set_title('Median RGB - All timesteps') + axes4.axis('off') + plt.show(fig4) + + # Define the Geomedian RGB colour box + def ii(p1r, p1g, p1b, p2r, p2g, p2b, p3r, p3g, p3b): + fig5, axes5 = plt.subplots(figsize=(3,3)) + fig5 = plt.imshow([[(int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).red.values.ravel()[0]), int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).green.values.ravel()[0]), int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).blue.values.ravel()[0]))]]) + axes5.set_title('Geomedian RGB - All timesteps') + axes5.axis('off') + plt.show(fig5) + + # Define 3-D axis to display vectors on + def j(p1r, p1g, p1b, p2r, p2g, p2b, p3r, p3g, p3b): + fig6 = plt.figure() + axes6 = fig6.add_subplot(111, projection='3d') + x = [p1r, p2r, p3r, int(np.median([p1r, p2r, p3r])), int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).red.values.ravel()[0])] + y = [p1g, p2g, p3g, int(np.median([p1g, p2g, p3g])), int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).green.values.ravel()[0])] + z = [p1b, p2b, p3b, int(np.median([p1b, p2b, p3b])), int(xr_geomedian(xr.Dataset({"red": (("x", "y", "time"), [[[np.float32(p1r), np.float32(p2r), np.float32(p3r)]]]), "green": (("x", "y", "time"), [[[np.float32(p1g), np.float32(p2g), np.float32(p3g)]]]), "blue": (("x", "y", "time"), [[[np.float32(p1b), np.float32(p2b), np.float32(p3b)]]])})).blue.values.ravel()[0])] + labels = [' 1', ' 2', ' 3', ' median', ' geomedian'] + axes6.scatter(x, y, z, c=['black','black','black','r', 'blue'], marker='o') + axes6.set_xlabel('Red') + axes6.set_ylabel('Green') + axes6.set_zlabel('Blue') + axes6.set_xlim3d(0, 255) + axes6.set_ylim3d(0, 255) + axes6.set_zlim3d(0, 255) + for ax, ay, az, label in zip(x, y, z, labels): + axes6.text(ax, ay, az, label) + plt.title('Each band represents a dimension.') + plt.show() + + # Define outputs + outf = widgets.interactive_output(f, {'p1r': p1r, 'p2r': p2r,'p3r': p3r, 'p1g': p1g, 'p2g': p2g,'p3g': p3g, 'p1b': p1b, 'p2b': p2b,'p3b': p3b}) + outg = widgets.interactive_output(g, {'p1r': p1r, 'p2r': p2r,'p3r': p3r, 'p1g': p1g, 'p2g': p2g,'p3g': p3g, 'p1b': p1b, 'p2b': p2b,'p3b': p3b}) + + outh = widgets.interactive_output(h, {'p1r': p1r, 'p1g': p1g, 'p1b': p1b}) + outhh = widgets.interactive_output(hh, {'p2r': p2r, 'p2g': p2g, 'p2b': p2b}) + outhhh = widgets.interactive_output(hhh, {'p3r': p3r, 'p3g': p3g, 'p3b': p3b}) + + outi = widgets.interactive_output(i, {'p1r': p1r, 'p2r': p2r,'p3r': p3r, 'p1g': p1g, 'p2g': p2g,'p3g': p3g, 'p1b': p1b, 'p2b': p2b,'p3b': p3b}) + outii = widgets.interactive_output(ii, {'p1r': p1r, 'p2r': p2r,'p3r': p3r, 'p1g': p1g, 'p2g': p2g,'p3g': p3g, 'p1b': p1b, 'p2b': p2b,'p3b': p3b}) + + outj = widgets.interactive_output(j, {'p1r': p1r, 'p2r': p2r,'p3r': p3r, 'p1g': p1g, 'p2g': p2g,'p3g': p3g, 'p1b': p1b, 'p2b': p2b,'p3b': p3b}) + + app_output = widgets.HBox([widgets.VBox([widgets.HBox([outh, widgets.VBox([ p1r, p1g, p1b])]), widgets.HBox([outhh, widgets.VBox([p2r, p2g, p2b])]), widgets.HBox([outhhh, widgets.VBox([ p3r, p3g, p3b])])]), widgets.VBox([widgets.HBox([widgets.VBox([outf, outi]), widgets.VBox([outg, outii])]), outj])]) + + return app_output \ No newline at end of file diff --git a/deafrica_tools/app/imageexport.py b/deafrica_tools/app/imageexport.py new file mode 100644 index 0000000..8b0cf1c --- /dev/null +++ b/deafrica_tools/app/imageexport.py @@ -0,0 +1,372 @@ +""" +Create an interactive map for selecting satellite imagery and exporting image files. +""" + +# Load modules +import datacube +import itertools +import numpy as np +import matplotlib.pyplot as plt +from odc.ui import select_on_a_map +from datacube.utils.geometry import CRS +from datacube.utils import masking +from skimage import exposure +from ipyleaflet import (WMSLayer, basemaps, basemap_to_tiles) +from traitlets import Unicode + +from deafrica_tools.spatial import reverse_geocode +from deafrica_tools.dask import create_local_dask_cluster + + +def select_region_app(date, + satellites, + size_limit=10000): + """ + An interactive app that allows the user to select a region from a + map using imagery from Sentinel-2 and Landsat. The output of this + function is used as the input to :func:`export_image_app` to export high- + resolution satellite images. + + Last modified: September 2021 + + Parameters + ---------- + date : str + The exact date used to plot imagery on the interactive map + (e.g. ``date='1988-01-01'``). + satellites : str + The satellite data to plot on the interactive map. The + following options are supported: + + ``'Landsat-9'``: data from the Landsat 9 satellite + ``'Landsat-8'``: data from the Landsat 8 satellite + ``'Landsat-7'``: data from the Landsat 7 satellite + ``'Landsat-5'``: data from the Landsat 5 satellite + ``'Sentinel-2'``: data from Sentinel-2A and Sentinel-2B + ``'Sentinel-2 geomedian'``: data from the Sentinel-2 annual geomedian + + size_limit : int, optional + An optional size limit for the area selection in sq km. + Defaults to 10000 sq km. + + Returns + ------- + A dictionary containing: + + * 'geopolygon' (defining the area to export imagery from), + * 'date' (date used to export imagery), and + * 'satellites' (the satellites from which to extract imagery). + + These are passed to the :func:`export_image_app` function to export the image. + """ + + ######################## + # Select and load data # + ######################## + + # Load DEA WMS + class TimeWMSLayer(WMSLayer): + time = Unicode('').tag(sync=True, o=True) + + # WMS layers + wms_params = { + 'Landsat-9': 'ls9_sr', + 'Landsat-8': 'ls8_sr', + 'Landsat-7': 'ls7_sr', + 'Landsat-5': 'ls5_sr', + 'Sentinel-2': 's2_l2a', + 'Sentinel-2 geomedian': 'gm_s2_annual' + } + + time_wms = TimeWMSLayer(url='https://ows.digitalearth.africa/', + layers=wms_params[satellites], + time=date, + format='image/png', + transparent=True, + attribution='Digital Earth Africa') + + # Plot interactive map to select area + basemap = basemap_to_tiles(basemaps.OpenStreetMap.Mapnik) + geopolygon = select_on_a_map(height='1000px', + layers=( + basemap, + time_wms, + ), + center=(4, 20), + zoom=4) + + # Test size of selected area + area = geopolygon.to_crs(crs=CRS('epsg:6933')).area / 1000000 + if area > size_limit: + print(f'Warning: Your selected area is {area:.00f} sq km. ' + f'Please select an area of less than {size_limit} sq km.' + f'\nTo select a smaller area, re-run the cell ' + f'above and draw a new polygon.') + + else: + return {'geopolygon': geopolygon, + 'date': date, + 'satellites': satellites} + + +def export_image_app(geopolygon, + date, + satellites, + style='True colour', + resolution=None, + vmin=0, + vmax=2000, + percentile_stretch=None, + power=None, + image_proc_funcs=None, + output_format="jpg", + standardise_name=False): + """ + Exports Digital Earth Africa satellite data as an image file + based on the extent and time period selected using + :func:`select_region_app`. The function supports Sentinel-2 and Landsat + data, creating True and False colour images. + + By default, files are named using: + + ``" - - - .png"`` + + Set ``standardise_name=True`` for a machine-readable name: + + ``"___.png"`` + + Last modified: September 2021 + + Parameters + ---------- + geopolygon : datacube.utils.geometry object + A datacube geopolygon providing the spatial bounds used to load + satellite data. + date : str + The exact date used to extract imagery + (e.g. `date='1988-01-01'`). + satellites : str + The satellite data to be used to extract imagery. The + following options are supported: + + ``'Landsat-9'``: data from the Landsat 9 satellite + ``'Landsat-8'``: data from the Landsat 8 satellite + ``'Landsat-7'``: data from the Landsat 7 satellite + ``'Landsat-5'``: data from the Landsat 5 satellite + ``'Sentinel-2'``: data from Sentinel-2A and Sentinel-2B + ``'Sentinel-2 geomedian'``: data from the Sentinel-2 annual geomedian + + style : str, optional + The style used to produce the image. Two options are currently + supported: + + * ``'True colour'``: Creates a true colour image using the red, + green and blue satellite bands + * ``'False colour'``: Creates a false colour image using + short-wave infrared, infrared and green satellite bands. + The specific bands used vary between Landsat and Sentinel-2. + + resolution : tuple, optional + The spatial resolution to load data. By default, the tool will + automatically set the best possible resolution depending on the + satellites selected (i.e 30 m for Landsat, 10 m for Sentinel-2). + Increasing this (e.g. to ``resolution=(-100, 100)``) can be useful + for loading large spatial extents. + vmin, vmax : int or float + The minimum and maximum surface reflectance values used to + clip the resulting imagery to enhance contrast. + percentile_stretch : tuple of floats, optional + An tuple of two floats (between 0.00 and 1.00) that can be used + to clip the imagery to based on percentiles to get more control + over the brightness and contrast of the image. The default is + ``None``; ``(0.02, 0.98)`` is equivelent to ``robust=True``. If this + parameter is used, ``vmin`` and ``vmax`` will have no effect. + power : float, optional + Raises imagery by a power to reduce bright features and + enhance dark features. This can add extra definition over areas + with extremely bright features like snow, beaches or salt pans. + image_proc_funcs : list of funcs, optional + An optional list containing functions that will be applied to + the output image. This can include image processing functions + such as increasing contrast, unsharp masking, saturation etc. + The function should take AND return a `numpy.ndarray` with + shape ``[y, x, bands]``. If your function has parameters, you + can pass in custom values using a lambda function, e.g.: + ``[lambda x: skimage.filters.unsharp_mask(x, radius=5, amount=0.2)]`` + output_format : str, optional + The output file format of the image. Valid options include ``'jpg'`` + and ``'png'``. Defaults to ``'jpg'``. + standardise_name : bool, optional + Whether to export the image file with a machine-readable + file name (e.g. ``___.png``) + """ + + ########################### + # Set up satellite params # + ########################### + + sat_params = { + 'Landsat-9': { + 'products': ['ls9_sr'], + 'resolution': [-30, 30], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_1', 'nir', 'green'] + } + }, + 'Landsat-8': { + 'products': ['ls8_sr'], + 'resolution': [-30, 30], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_1', 'nir', 'green'] + } + }, + 'Landsat-7': { + 'products': ['ls7_sr'], + 'resolution': [-30, 30], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_1', 'nir', 'green'] + } + }, + 'Landsat-5': { + 'products': ['ls5_sr'], + 'resolution': [-30, 30], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_1', 'nir', 'green'] + } + }, + 'Sentinel-2': { + 'products': ['s2_l2a'], + 'resolution': [-10, 10], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_2', 'nir_1', 'green'] + } + }, + 'Sentinel-2 geomedian': { + 'products': ['gm_s2_annual'], + 'resolution': [-10, 10], + 'styles': { + 'True colour': ['red', 'green', 'blue'], + 'False colour': ['swir_2', 'nir_1', 'green'] + } + }, + } + + ############# + # Load data # + ############# + + # Connect to datacube database + dc = datacube.Datacube(app='Exporting_satellite_images') + + # Configure local dask cluster + client = create_local_dask_cluster(return_client=True) + + # Create query after adjusting interval time to UTC by + # adding a UTC offset of -10 hours. + start_date = np.datetime64(date) + query_params = { + 'time': (str(start_date)), + 'geopolygon': geopolygon + } + + # Find matching datasets + dss = [ + dc.find_datasets(product=i, **query_params) + for i in sat_params[satellites]['products'] + ] + dss = list(itertools.chain.from_iterable(dss)) + + # Get CRS and sensor + crs = str(dss[0].crs) + + if satellites == 'Sentinel-2 geomedian': + sensor = satellites + else: + sensor = dss[0].metadata_doc['properties']['eo:platform'].capitalize() + sensor = sensor[0:-1].replace('_', '-') + sensor[-1].capitalize() + + # Use resolution if provided, otherwise use default + if resolution: + sat_params[satellites]['resolution'] = resolution + + load_params = { + 'output_crs': crs, + 'resolution': sat_params[satellites]['resolution'], + 'resampling': 'bilinear' + } + + # Load data from datasets + ds = dc.load(datasets=dss, + measurements=sat_params[satellites]['styles'][style], + group_by='solar_day', + dask_chunks={ + 'time': 1, + 'x': 3000, + 'y': 3000 + }, + **load_params, + **query_params) + ds = masking.mask_invalid_data(ds) + + rgb_array = ds.isel(time=0).to_array().values + + ############ + # Plotting # + ############ + + # Create unique file name + centre_coords = geopolygon.centroid.coords[0][::-1] + site = reverse_geocode(coords=centre_coords) + fname = (f"{sensor} - {date} - {site} - {style}, " + f"{load_params['resolution'][1]} m resolution.{output_format}") + + # Remove spaces and commas if requested + if standardise_name: + fname = fname.replace(' - ', '_').replace(', ', + '-').replace(' ', + '-').lower() + + print( + f'\nExporting image to {fname}.\nThis may take several minutes to complete...' + ) + + # Convert to numpy array + rgb_array = np.transpose(rgb_array, axes=[1, 2, 0]) + + # If percentile stretch is supplied, calculate vmin and vmax + # from percentiles + if percentile_stretch: + vmin, vmax = np.nanpercentile(rgb_array, percentile_stretch) + + # Raise by power to dampen bright features and enhance dark. + # Raise vmin and vmax by same amount to ensure proper stretch + if power: + rgb_array = rgb_array**power + vmin, vmax = vmin**power, vmax**power + + # Rescale/stretch imagery between vmin and vmax + rgb_rescaled = exposure.rescale_intensity(rgb_array.astype(float), + in_range=(vmin, vmax), + out_range=(0.0, 1.0)) + + # Apply image processing funcs + if image_proc_funcs: + for i, func in enumerate(image_proc_funcs): + print(f'Applying custom function {i + 1}') + rgb_rescaled = func(rgb_rescaled) + + # Plot RGB + plt.imshow(rgb_rescaled) + + # Export to file + plt.imsave(fname=fname, arr=rgb_rescaled, format=output_format) + + # Close dask client + client.shutdown() + + print('Finished exporting image.') diff --git a/deafrica_tools/app/wetlandsinsighttool.py b/deafrica_tools/app/wetlandsinsighttool.py new file mode 100644 index 0000000..e8b6ce5 --- /dev/null +++ b/deafrica_tools/app/wetlandsinsighttool.py @@ -0,0 +1,388 @@ +""" +Wetlands insight tool widget, which can be used to run an interactive +version of the wetlands insight tool. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import datacube +import warnings +import seaborn as sns +import matplotlib.pyplot as plt +from datacube.utils.geometry import CRS +from ipyleaflet import ( + WMSLayer, + basemaps, + basemap_to_tiles, + Map, + DrawControl, + WidgetControl, + LayerGroup, + LayersControl, +) +from traitlets import Unicode +from ipywidgets import ( + GridspecLayout, + Button, + Layout, + HBox, + VBox, + HTML, + Output, +) +import json +import geopandas as gpd +from io import BytesIO +from dask.diagnostics import ProgressBar + +import deafrica_tools +from deafrica_tools.dask import create_local_dask_cluster +from deafrica_tools.wetlands import WIT_drill +import deafrica_tools.app.widgetconstructors as deawidgets + + +def make_box_layout(): + return Layout( + #border='solid 1px black', + margin='0px 10px 10px 0px', + padding='5px 5px 5px 5px', + width='100%', + height='100%', + ) + + +def create_expanded_button(description, button_style): + return Button( + description=description, + button_style=button_style, + layout=Layout(width="auto", height="auto"), + ) + + +class wit_app(HBox): + def __init__(self, lang=None): + super().__init__() + + deafrica_tools.set_lang(lang) + + ########################################################## + # INITIAL ATTRIBUTES # + + self.startdate = "2020-01-01" + self.enddate = "2020-03-01" + self.mingooddata = 0.0 + self.resamplingfreq = "1M" + self.out_csv = "example_WIT.csv" + self.out_plot = "example_WIT.png" + self.product_list = [ + (_("None"), "none"), + (_("ESRI World Imagery"), "esri_world_imagery"), + (_("Sentinel-2 Geomedian"), "gm_s2_annual"), + (_("Water Observations from Space"), "wofs_ls_summary_annual"), + + ] + self.product = self.product_list[0][1] + self.product_year = "2020-01-01" + self.target = None + self.action = None + self.gdf_drawn = None + + ########################################################## + # HEADER FOR APP # + + # Create the Header widget + header_title_text = _("Wetlands Insight Tool") + instruction_text = _("Select parameters and AOI") + self.header = deawidgets.create_html(f"

{header_title_text}

{instruction_text}

") + self.header.layout = make_box_layout() + + ########################################################## + # HANDLER FUNCTION FOR DRAW CONTROL # + + # Define the action to take once something is drawn on the map + def update_geojson(target, action, geo_json): + + self.action = action + + json_data = json.dumps(geo_json) + binary_data = json_data.encode() + io = BytesIO(binary_data) + io.seek(0) + + gdf = gpd.read_file(io) + gdf.crs = "EPSG:4326" + self.gdf_drawn = gdf + + gdf_drawn_epsg6933 = gdf.copy().to_crs("EPSG:6933") + m2_per_km2 = 10 ** 6 + area = gdf_drawn_epsg6933.area.values[0] / m2_per_km2 + polyarea_label = _('Total polygon area') + polyarea_text = f"

{polyarea_label}: {area:.2f} km2

" + + if area <= 3000: + confirmation_text = '

' + _('Area falls within recommended limit') + '

' + self.header.value = header_title_text + polyarea_text + confirmation_text + else: + warning_text = '

' + _('Area is too large, please update your polygon') + '

' + self.header.value = header_title_text + polyarea_text + warning_text + + ########################################################## + # WIDGETS FOR APP OUTPUTS # + + self.dask_client = Output(layout=make_box_layout()) + self.progress_bar = Output(layout=make_box_layout()) + self.wit_plot = Output(layout=make_box_layout()) + self.progress_header = deawidgets.create_html("") + + ########################################################## + # MAP WIDGET, DRAWING TOOLS, WMS LAYERS # + + # Create drawing tools + desired_drawtools = ['rectangle', 'polygon'] + draw_control = deawidgets.create_drawcontrol(desired_drawtools) + + # Begin by displaying an empty layer group, and update the group with desired WMS on interaction. + self.deafrica_layers = LayerGroup(layers=()) + self.deafrica_layers.name = _('Map Overlays') + + # Create map widget + self.m = deawidgets.create_map() + + self.m.layout = make_box_layout() + + # Add tools to map widget + self.m.add_control(draw_control) + self.m.add_layer(self.deafrica_layers) + + # Store current basemap for future use + self.basemap = self.m.basemap + + ########################################################## + # WIDGETS FOR APP CONTROLS # + + # Create parameter widgets + startdate_picker = deawidgets.create_datepicker() + enddate_picker = deawidgets.create_datepicker() + min_good_data = deawidgets.create_boundedfloattext(self.mingooddata, 0.0, 1.0, 0.05) + resampling_freq = deawidgets.create_inputtext(self.resamplingfreq, self.resamplingfreq) + output_csv = deawidgets.create_inputtext(self.out_csv, self.out_csv) + output_plot = deawidgets.create_inputtext(self.out_plot, self.out_plot) + deaoverlay_dropdown = deawidgets.create_dropdown(self.product_list, self.product_list[0][1]) + run_button = create_expanded_button(_("Run"), "info") + + ########################################################## + # COLLECTION OF ALL APP CONTROLS # + + parameter_selection = VBox( + [ + HTML("" + _("Map Overlay:") + ""), + deaoverlay_dropdown, + HTML("" + _("Start Date:") + ""), + startdate_picker, + HTML("" + _("End Date:") + ""), + enddate_picker, + HTML("" + _("Minimum Good Data:") + ""), + min_good_data, + HTML("" + _("Resampling Frequency:") + ""), + resampling_freq, + HTML("" + _("Output CSV:") + ""), + output_csv, + HTML("" + _("Output Plot:") + ""), + output_plot, + ] + ) + parameter_selection.layout = make_box_layout() + + ########################################################## + # SPECIFICATION OF APP LAYOUT # + + # Create the layout #[rowspan, colspan] + grid = GridspecLayout(11, 10, height="1100px", width="auto") + + # Controls and Status + grid[0, :] = self.header + grid[1:6, 0:2] = parameter_selection + grid[6, 0:2] = run_button + + # Dask and Progress info + grid[1, 7:] = self.dask_client + grid[2:7, 7:] = self.progress_bar + + # Map + grid[1:7, 2:7] = self.m + + # Plot + grid[7:, :] = self.wit_plot + + # Display using HBox children attribute + self.children = [grid] + + ########################################################## + # SPECIFICATION UPDATE FUNCTIONS FOR EACH WIDGET # + + # Run update functions whenever various widgets are changed. + startdate_picker.observe(self.update_startdate, "value") + enddate_picker.observe(self.update_enddate, "value") + min_good_data.observe(self.update_mingooddata, "value") + resampling_freq.observe(self.update_resamplingfreq, "value") + output_csv.observe(self.update_outputcsv, "value") + output_plot.observe(self.update_outputplot, "value") + deaoverlay_dropdown.observe(self.update_deaoverlay, "value") + run_button.on_click(self.run_app) + draw_control.on_draw(update_geojson) + + ############################################################## + # DEFINITION OF ALL UPDATE FUNCTIONS # + + # set the start date to the new edited date + def update_startdate(self, change): + self.startdate = change.new + + # set the end date to the new edited date + def update_enddate(self, change): + self.enddate = change.new + + # set the min good data + def update_mingooddata(self, change): + self.mingooddata = change.new + + # set the resampling frequency + def update_resamplingfreq(self, change): + self.resamplingfreq = change.new + + # set the output csv + def update_outputcsv(self, change): + self.out_csv = change.new + + # set the output plot + def update_outputplot(self, change): + self.out_plot = change.new + + # Update product + def update_deaoverlay(self, change): + + self.product = change.new + + if self.product == "none": + self.deafrica_layers.clear_layers() + elif self.product == "esri_world_imagery": + self.deafrica_layers.clear_layers() + layer = basemap_to_tiles(basemaps.Esri.WorldImagery) + self.deafrica_layers.add_layer(layer) + else: + self.deafrica_layers.clear_layers() + layer = deawidgets.create_dea_wms_layer(self.product, self.product_year) + self.deafrica_layers.add_layer(layer) + + def run_app(self, change): + + # Clear progress bar and output areas before running + self.dask_client.clear_output() + self.progress_bar.clear_output() + self.wit_plot.clear_output() + + # Connect to datacube database + dc = datacube.Datacube(app="wetland_app") + + # Configure local dask cluster + with self.dask_client: + client = create_local_dask_cluster( + return_client=True, display_client=True + ) + + # Set any defaults + TCW_threshold = -0.035 + dask_chunks = dict(x=1000, y=1000, time=1) + + #check resampling freq + if self.resamplingfreq == 'None': + rsf = None + else: + rsf = self.resamplingfreq + + self.progress_header.value = f"

"+_("Progress")+"

" + + # run wetlands polygon drill + with self.progress_bar: +# with ProgressBar(): + warnings.filterwarnings("ignore") + try: + df = WIT_drill( + gdf=self.gdf_drawn, + time=(self.startdate, self.enddate), + min_gooddata=self.mingooddata, + resample_frequency=rsf, + TCW_threshold=TCW_threshold, + export_csv=self.out_csv, + dask_chunks=dask_chunks, + verbose=False, + verbose_progress=True, + ) + print(_("WIT complete")) + except AttributeError: + print(_("No polygon selected")) + + # close down the dask client + client.shutdown() + + # save the csv + if self.out_csv: + df.to_csv(self.out_csv, index_label="Datetime") + + # ---Plotting------------------------------ + + with self.wit_plot: + + fontsize = 17 + plt.rcParams.update({"font.size": fontsize}) + # set up color palette + pal = [ + sns.xkcd_rgb["cobalt blue"], + sns.xkcd_rgb["neon blue"], + sns.xkcd_rgb["grass"], + sns.xkcd_rgb["beige"], + sns.xkcd_rgb["brown"], + ] + + # make a stacked area plot + plt.close("all") + + fig, ax = plt.subplots(constrained_layout=True, figsize=(20, 6)) + + ax.stackplot( + df.index, + df.wofs_area_percent, + df.wet_percent, + df.green_veg_percent, + df.dry_veg_percent, + df.bare_soil_percent, + labels=[ + _("open water"), + _("wet"), + _("green veg"), + _("dry veg"), + _("bare soil"), + ], + colors=pal, + alpha=0.6, + ) + + # set axis limits to the min and max + ax.set_ylim(0, 100) + ax.set_xlim(df.index[0], df.index[-1]) + ax.tick_params(axis="x", labelsize=fontsize) + + # add a legend and a tight plot box + ax.legend(loc="lower left", framealpha=0.6) + ax.set_title(_("Percentage Fractional Cover, Wetness, and Water")) + # plt.tight_layout() + plt.show() + + if self.out_plot: + # save the figure + fig.savefig(f"{self.out_plot}") diff --git a/deafrica_tools/app/widgetconstructors.py b/deafrica_tools/app/widgetconstructors.py new file mode 100644 index 0000000..6881712 --- /dev/null +++ b/deafrica_tools/app/widgetconstructors.py @@ -0,0 +1,367 @@ +""" +Functions for easily defining widgets in the context of DE Africa notebooks. + +These are largely customised wrappers around existing widgets. +""" + +import ipyleaflet as leaflet +from ipyleaflet import LayersControl +import ipywidgets as widgets +from traitlets import Unicode + + +def create_datepicker(description='', value=None, layout={'width': '85%'}): + ''' + Create a DatePicker widget + + Last modified: July 2022 + + Parameters + ---------- + description : string + descirption label to attach + layout : dictionary + any layout commands for the widget + + Returns + ------- + date_picker : ipywidgets.widgets.widget_date.DatePicker + + ''' + + date_picker = widgets.DatePicker( + description=description, + layout=layout, + disabled=False, + value=value + ) + + return date_picker + + +def create_inputtext(value, placeholder, description="", layout={'width': '85%'}): + ''' + Create a Text widget + + Last modified: October 2021 + + Parameters + ---------- + value : string + initial value of the widget + placeholder : string + placeholder text to display to the user before intput + description : string + descirption label to attach + layout : dictionary + any layout commands for the widget + + Returns + ------- + input_text : ipywidgets.widgets.widget_string.Text + + ''' + + input_text = widgets.Text( + value=value, + placeholder=placeholder, + description=description, + layout=layout, + disabled=False + ) + + return input_text + + +def create_boundedfloattext(value, min_val, max_val, step_val, description="", layout={'width': '85%'}): + ''' + Create a BoundedFloatText widget + + Last modified: October 2021 + + Parameters + ---------- + value : float + initial value of the widget + min_val : float + minimum allowed value for the float + max_val : float + maximum allowed value for the float + step_val : float + allowed increment for the float + description : string + descirption label to attach + layout : dictionary + any layout commands for the widget + + Returns + ------- + float_text : ipywidgets.widgets.widget_float.BoundedFloatText + + ''' + + float_text = widgets.BoundedFloatText( + value=value, + min=min_val, + max=max_val, + step=step_val, + description=description, + layout=layout, + disabled=False, + ) + + return float_text + + +def create_dropdown(options, value, description="", layout={'width': '85%'}): + ''' + Create a Dropdown widget + + Last modified: October 2021 + + Parameters + ---------- + options : list + a list of options for the user to select from + value : string + initial value of the widget + description : string + descirption label to attach + layout : dictionary + any layout commands for the widget + + Returns + ------- + dropdown : ipywidgets.widgets.widget_selection.Dropdown + + ''' + + dropdown = widgets.Dropdown( + options=options, + value=value, + description=description, + layout=layout, + disabled=False, + ) + + return dropdown + + +def create_html(value): + ''' + Create a HTML widget + + Last modified: October 2021 + + Parameters + ---------- + value : string + HTML text to display + + Returns + ------- + html : ipywidgets.widgets.widget_string.HTML + + ''' + + html = widgets.HTML( + value=value, + ) + + return html + + +def create_map(map_center=(4, 20), zoom_level=3, basemap=leaflet.basemaps.OpenStreetMap.Mapnik, basemap_name='Open Street Map'): + ''' + Create an interactive ipyleaflet map + + Last modified: October 2021 + + Parameters + ---------- + map_center : tuple + A tuple containing the latitude and longitude to focus on. + Defaults to center of Africa, (4, 20) + zoom_level : integer + Zoom level for the map + Defaults to 3 to view all of Africa + basemap : ipyleaflet basemap (dict) + Basemap to use, can be any from https://ipyleaflet.readthedocs.io/en/latest/api_reference/basemaps.html + Defaults to Open Street Map (basemaps.OpenStreetMap.Mapnik) + basemap_name : string + Layer name for the basemap + + Returns + ------- + m : ipyleaflet.leaflet.Map + interactive ipyleaflet map + + ''' + + basemap_tiles = leaflet.basemap_to_tiles(basemap) + basemap_tiles.name = basemap_name + + m = leaflet.Map(center=map_center, zoom=zoom_level, basemap=basemap_tiles, scroll_wheel_zoom=True) + + return m + + +def create_dea_wms_layer(product, date): + ''' + Create a Digital Earth Africa WMS layer to add to a map + + Last modified: October 2021 + + Parameters + ---------- + product : string + The Digital Earth Africa product to load + (e.g. 'gm_s2_annual') + date : string (yyyy-mm-dd format) + The date to load the product for + + Returns + ------- + time_wms : ipyleaflet WMS layer + + ''' + + + # Load DEA WMS + class TimeWMSLayer(leaflet.WMSLayer): + time = Unicode("").tag(sync=True, o=True) + + time_wms = TimeWMSLayer( + url="https://ows.digitalearth.africa/", + layers=product, + time=date, + format="image/png", + transparent=True, + attribution="Digital Earth Africa", + ) + + return time_wms + + +def create_drawcontrol( + draw_controls = ['rectangle', 'polygon', 'circle', 'polyline', 'marker', 'circlemarker'], + rectangle_options={}, + polygon_options={}, + circle_options={}, + polyline_options={}, + marker_options={}, + circlemarker_options={}, +): + ''' + Create a draw control widget to add to ipyleaflet maps + + Last modified: October 2021 + + Parameters + ---------- + draw_controls : list + List of draw controls to add to the map. Defaults to adding all + Viable options are 'rectangle', 'polygon', 'circle', 'polyline', 'marker', 'circlemarker' + rectangle_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + polygon_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + circle_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + polyline_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + marker_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + circlemarker_options : dict + Options to customise the appearence of the relevant shape + User can supply, or leave blank to get default DE Africa appearence + + + Returns + ------- + draw_control : ipyleaflet.leaflet.DrawControl + + ''' + + # Set defualt DE Africa styling options for polygons + default_shapeoptions = { + "color": "#FFFFFF", + "opacity": 0.8, + "fillColor": "#336699", + "fillOpacity": 0.4, + } + default_drawerror = { + "color": "#FF6633", + "message": "Drawing error, clear all and try again" + } + + # Set draw control appearence to DE Africa defaults + # Do this if user has requested a control, but has not provided a corresponding options dict + + if ('rectangle' in draw_controls) and (not rectangle_options): + rectangle_options = {"shapeOptions": default_shapeoptions} + + if ('polygon' in draw_controls) and (not polygon_options): + polygon_options = { + "shapeOptions": default_shapeoptions, + "drawError": default_drawerror, + "allowIntersection": False, + } + + if ('circle' in draw_controls) and (not circle_options): + circle_options = {"shapeOptions": default_shapeoptions} + + if ('polyline' in draw_controls) and (not polyline_options): + polyline_options = {"shapeOptions": default_shapeoptions} + + if ('marker' in draw_controls) and (not marker_options): + marker_options = {'shapeOptions': {'opacity': 1.0}} + + if ('circlemarker' in draw_controls) and (not circlemarker_options): + circlemarker_options = {"shapeOptions": default_shapeoptions} + + # Instantiate draw control and add options + draw_control = leaflet.DrawControl() + draw_control.rectangle = rectangle_options + draw_control.polygon = polygon_options + draw_control.marker = marker_options + draw_control.circle = circle_options + draw_control.circlemarker = circlemarker_options + draw_control.polyline = polyline_options + + return draw_control + + +def create_checkbox(value, description="", layout={'width': '85%'}): + ''' + Create a Checkbox widget + + Last modified: July 2022 + + Parameters + ---------- + value : string + initial value of the widget; True or False + description : string + description label to attach + layout : dictionary + any layout commands for the widget + + Returns + ------- + dropdown : ipywidgets.widgets.widget_selection.Dropdown + + ''' + + checklist = widgets.Checkbox(value=value, + description=description, + layout=layout, + disabled=False, + indent=False) + + return checklist \ No newline at end of file diff --git a/deafrica_tools/areaofinterest.py b/deafrica_tools/areaofinterest.py new file mode 100644 index 0000000..ed69efa --- /dev/null +++ b/deafrica_tools/areaofinterest.py @@ -0,0 +1,54 @@ +""" +Function for defining an area of interest using either a point and buffer or a vector file. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import geopandas as gpd +from shapely.geometry import box +from geojson import Feature, Point, FeatureCollection + +def define_area(lat=None, lon=None, buffer=None, vector_path=None): + ''' + Define an area of interest using either a point and buffer or a vector. + + Parameters: + ----------- + lat : float, optional + The latitude of the center point of the area of interest. + lon : float, optional + The longitude of the center point of the area of interest. + buffer : float, optional + The buffer around the center point, in degrees. + vector_path : str, optional + The path to a vector defining the area of interest. + + Returns: + -------- + feature_collection : dict + A GeoJSON feature collection representing the area of interest. + ''' + # Define area using point and buffer + if lat is not None and lon is not None and buffer is not None: + lat_range = (lat - buffer, lat + buffer) + lon_range = (lon - buffer, lon + buffer) + box_geom = box(min(lon_range), min(lat_range), max(lon_range), max(lat_range)) + aoi = gpd.GeoDataFrame(geometry=[box_geom], crs='EPSG:4326') + + # Define area using vector + elif vector_path is not None: + aoi = gpd.read_file(vector_path).to_crs("EPSG:4326") + # If neither option is provided, raise an error + else: + raise ValueError("Either lat/lon/buffer or vector_path must be provided.") + + # Convert the GeoDataFrame to a GeoJSON FeatureCollection + features = [Feature(geometry=row["geometry"], properties=row.drop("geometry").to_dict()) for _, row in aoi.iterrows()] + feature_collection = FeatureCollection(features) + + return feature_collection \ No newline at end of file diff --git a/deafrica_tools/bandindices.py b/deafrica_tools/bandindices.py new file mode 100644 index 0000000..0092925 --- /dev/null +++ b/deafrica_tools/bandindices.py @@ -0,0 +1,615 @@ +""" +Functions for computing remote sensing band indices on Digital Earth Africa +data. +""" + +# Import required packages +import warnings +import numpy as np + +# Define custom functions +def calculate_indices( + ds, + index=None, + collection=None, + satellite_mission=None, + custom_varname=None, + normalise=True, + drop=False, + deep_copy=True, +): + """ + Takes an xarray dataset containing spectral bands, calculates one of + a set of remote sensing indices, and adds the resulting array as a + new variable in the original dataset. + + Last modified: July 2022 + + Parameters + ---------- + ds : xarray Dataset + A two-dimensional or multi-dimensional array with containing the + spectral bands required to calculate the index. These bands are + used as inputs to calculate the selected water index. + + index : str or list of strs + A string giving the name of the index to calculate or a list of + strings giving the names of the indices to calculate: + + * ``'ASI'`` (Artificial Surface Index, Yongquan Zhao & Zhe Zhu 2022) + * ``'AWEI_ns'`` (Automated Water Extraction Index, no shadows, Feyisa 2014) + * ``'AWEI_sh'`` (Automated Water Extraction Index, shadows, Feyisa 2014) + * ``'BAEI'`` (Built-Up Area Extraction Index, Bouzekri et al. 2015) + * ``'BAI'`` (Burn Area Index, Martin 1998) + * ``'BSI'`` (Bare Soil Index, Rikimaru et al. 2002) + * ``'BUI'`` (Built-Up Index, He et al. 2010) + * ``'CMR'`` (Clay Minerals Ratio, Drury 1987) + * ``'ENDISI'`` (Enhanced Normalised Difference for Impervious Surfaces Index, Chen et al. 2019) + * ``'EVI'`` (Enhanced Vegetation Index, Huete 2002) + * ``'FMR'`` (Ferrous Minerals Ratio, Segal 1982) + * ``'IOR'`` (Iron Oxide Ratio, Segal 1982) + * ``'LAI'`` (Leaf Area Index, Boegh 2002) + * ``'MBI'`` (Modified Bare Soil Index, Nguyen et al. 2021) + * ``'MNDWI'`` (Modified Normalised Difference Water Index, Xu 1996) + * ``'MSAVI'`` (Modified Soil Adjusted Vegetation Index, Qi et al. 1994) + * ``'NBI'`` (New Built-Up Index, Jieli et al. 2010) + * ``'NBR'`` (Normalised Burn Ratio, Lopez Garcia 1991) + * ``'NDBI'`` (Normalised Difference Built-Up Index, Zha 2003) + * ``'NDCI'`` (Normalised Difference Chlorophyll Index, Mishra & Mishra, 2012) + * ``'NDMI'`` (Normalised Difference Moisture Index, Gao 1996) + * ``'NDSI'`` (Normalised Difference Snow Index, Hall 1995) + * ``'NDTI'`` (Normalised Difference Turbidity Index, Lacaux et al. 2007) + * ``'NDVI'`` (Normalised Difference Vegetation Index, Rouse 1973) + * ``'NDWI'`` (Normalised Difference Water Index, McFeeters 1996) + * ``'SAVI'`` (Soil Adjusted Vegetation Index, Huete 1988) + * ``'TCB'`` (Tasseled Cap Brightness, Crist 1985) + * ``'TCG'`` (Tasseled Cap Greeness, Crist 1985) + * ``'TCW'`` (Tasseled Cap Wetness, Crist 1985) + * ``'WI'`` (Water Index, Fisher 2016) + + collection : str + Deprecated in version 0.1.7. Use `satellite_mission` instead. + + Valid options are: + * ``'c2'`` (for USGS Landsat Collection 2) + If 'c2', then `satellite_mission='ls'`. + * ``'s2'`` (for Sentinel-2) + If 's2', then `satellite_mission='s2'`. + + satellite_mission : str + An string that tells the function which satellite mission's data is + being used to calculate the index. This is necessary because + different satellite missions use different names for bands covering + a similar spectra. + + Valid options are: + + * ``'ls'`` (for USGS Landsat) + * ``'s2'`` (for Copernicus Sentinel-2) + + custom_varname : str, optional + By default, the original dataset will be returned with + a new index variable named after `index` (e.g. 'NDVI'). To + specify a custom name instead, you can supply e.g. + `custom_varname='custom_name'`. Defaults to None, which uses + `index` to name the variable. + + normalise : bool, optional + Some coefficient-based indices (e.g. ``'WI'``, ``'BAEI'``, + ``'AWEI_ns'``, ``'AWEI_sh'``, ``'TCW'``, ``'TCG'``, ``'TCB'``, + ``'EVI'``, ``'LAI'``, ``'SAVI'``, ``'MSAVI'``) + produce different results if surface reflectance values are not + scaled between 0.0 and 1.0 prior to calculating the index. + Setting `normalise=True` first scales values to a 0.0-1.0 range + by dividing by 10000.0. Defaults to True. + + drop : bool, optional + Provides the option to drop the original input data, thus saving + space. If `drop=True`, returns only the index and its values. + + deep_copy: bool, optional + If `deep_copy=False`, calculate_indices will modify the original + array, adding bands to the input dataset and not removing them. + If the calculate_indices function is run more than once, variables + may be dropped incorrectly producing unexpected behaviour. This is + a bug and may be fixed in future releases. This is only a problem + when `drop=True`. + + Returns + ------- + ds : xarray Dataset + The original xarray Dataset inputted into the function, with a + new varible containing the remote sensing index as a DataArray. + If drop = True, the new variable/s as DataArrays in the + original Dataset. + """ + + # Set ds equal to a copy of itself in order to prevent the function + # from editing the input dataset. This is to prevent unexpected + # behaviour though it uses twice as much memory. + if deep_copy: + ds = ds.copy(deep=True) + + # Capture input band names in order to drop these if drop=True + if drop: + bands_to_drop = list(ds.data_vars) + print(f"Dropping bands {bands_to_drop}") + + # Dictionary containing remote sensing index band recipes + index_dict = { + # Normalised Difference Vegation Index, Rouse 1973 + "NDVI": lambda ds: (ds.nir - ds.red) / (ds.nir + ds.red), + # Enhanced Vegetation Index, Huete 2002 + "EVI": lambda ds: ( + 2.5 * ((ds.nir - ds.red) / (ds.nir + 6 * ds.red - 7.5 * ds.blue + 1)) + ), + # Leaf Area Index, Boegh 2002 + "LAI": lambda ds: ( + 3.618 + * ((2.5 * (ds.nir - ds.red)) / (ds.nir + (6 * ds.red) - (7.5 * ds.blue) + 1)) + - 0.118 + ), + # Soil Adjusted Vegetation Index, Huete 1988 + "SAVI": lambda ds: ((1.5 * (ds.nir - ds.red)) / (ds.nir + ds.red + 0.5)), + # Mod. Soil Adjusted Vegetation Index, Qi et al. 1994 + "MSAVI": lambda ds: ( + (2 * ds.nir + 1 - ((2 * ds.nir + 1) ** 2 - 8 * (ds.nir - ds.red)) ** 0.5) + / 2 + ), + # Normalised Difference Moisture Index, Gao 1996 + "NDMI": lambda ds: (ds.nir - ds.swir_1) / (ds.nir + ds.swir_1), + # Normalised Burn Ratio, Lopez Garcia 1991 + "NBR": lambda ds: (ds.nir - ds.swir_2) / (ds.nir + ds.swir_2), + # Burn Area Index, Martin 1998 + "BAI": lambda ds: (1.0 / ((0.10 - ds.red) ** 2 + (0.06 - ds.nir) ** 2)), + # Normalised Difference Chlorophyll Index, + # (Mishra & Mishra, 2012) + "NDCI": lambda ds: (ds.red_edge_1 - ds.red) / (ds.red_edge_1 + ds.red), + # Normalised Difference Snow Index, Hall 1995 + "NDSI": lambda ds: (ds.green - ds.swir_1) / (ds.green + ds.swir_1), + # Normalised Difference Water Index, McFeeters 1996 + "NDWI": lambda ds: (ds.green - ds.nir) / (ds.green + ds.nir), + # Modified Normalised Difference Water Index, Xu 2006 + "MNDWI": lambda ds: (ds.green - ds.swir_1) / (ds.green + ds.swir_1), + # Normalised Difference Built-Up Index, Zha 2003 + "NDBI": lambda ds: (ds.swir_1 - ds.nir) / (ds.swir_1 + ds.nir), + # Built-Up Index, He et al. 2010 + "BUI": lambda ds: ((ds.swir_1 - ds.nir) / (ds.swir_1 + ds.nir)) + - ((ds.nir - ds.red) / (ds.nir + ds.red)), + # Built-up Area Extraction Index, Bouzekri et al. 2015 + "BAEI": lambda ds: (ds.red + 0.3) / (ds.green + ds.swir_1), + # New Built-up Index, Jieli et al. 2010 + "NBI": lambda ds: (ds.swir_1 + ds.red) / ds.nir, + # Bare Soil Index, Rikimaru et al. 2002 + "BSI": lambda ds: ((ds.swir_1 + ds.red) - (ds.nir + ds.blue)) + / ((ds.swir_1 + ds.red) + (ds.nir + ds.blue)), + # Automated Water Extraction Index (no shadows), Feyisa 2014 + "AWEI_ns": lambda ds: ( + 4 * (ds.green - ds.swir_1) - (0.25 * ds.nir * +2.75 * ds.swir_2) + ), + # Automated Water Extraction Index (shadows), Feyisa 2014 + "AWEI_sh": lambda ds: ( + ds.blue + 2.5 * ds.green - 1.5 * (ds.nir + ds.swir_1) - 0.25 * ds.swir_2 + ), + # Water Index, Fisher 2016 + "WI": lambda ds: ( + 1.7204 + + 171 * ds.green + + 3 * ds.red + - 70 * ds.nir + - 45 * ds.swir_1 + - 71 * ds.swir_2 + ), + # Tasseled Cap Wetness, Crist 1985 + "TCW": lambda ds: ( + 0.0315 * ds.blue + + 0.2021 * ds.green + + 0.3102 * ds.red + + 0.1594 * ds.nir + + -0.6806 * ds.swir_1 + + -0.6109 * ds.swir_2 + ), + # Tasseled Cap Greeness, Crist 1985 + "TCG": lambda ds: ( + -0.1603 * ds.blue + + -0.2819 * ds.green + + -0.4934 * ds.red + + 0.7940 * ds.nir + + -0.0002 * ds.swir_1 + + -0.1446 * ds.swir_2 + ), + # Tasseled Cap Brightness, Crist 1985 + "TCB": lambda ds: ( + 0.2043 * ds.blue + + 0.4158 * ds.green + + 0.5524 * ds.red + + 0.5741 * ds.nir + + 0.3124 * ds.swir_1 + + -0.2303 * ds.swir_2 + ), + # Clay Minerals Ratio, Drury 1987 + "CMR": lambda ds: (ds.swir_1 / ds.swir_2), + # Ferrous Minerals Ratio, Segal 1982 + "FMR": lambda ds: (ds.swir_1 / ds.nir), + # Iron Oxide Ratio, Segal 1982 + "IOR": lambda ds: (ds.red / ds.blue), + # Normalized Difference Turbidity Index, Lacaux, J.P. et al. 2007 + "NDTI": lambda ds: (ds.red - ds.green) / (ds.red + ds.green), + # Modified Bare Soil Index, Nguyen et al. 2021 + "MBI": lambda ds: ((ds.swir_1 - ds.swir_2 - ds.nir) / (ds.swir_1 + ds.swir_2 + ds.nir)) + 0.5, + } + + # Enhanced Normalised Difference Impervious Surfaces Index, Chen et al. 2019 + def mndwi(ds): + return (ds.green - ds.swir_1) / (ds.green + ds.swir_1) + def swir_diff(ds): + return ds.swir_1/ds.swir_2 + def alpha(ds): + return (2*(np.mean(ds.blue)))/(np.mean(swir_diff(ds)) + np.mean(mndwi(ds)**2)) + def ENDISI(ds): + m = mndwi(ds) + s = swir_diff(ds) + a = alpha(ds) + return (ds.blue - (a)*(s + m**2))/(ds.blue + (a)*(s + m**2)) + + index_dict["ENDISI"] = ENDISI + + ## Artificial Surface Index, Yongquan Zhao & Zhe Zhu 2022 + def af(ds): + AF = (ds.nir - ds.blue) / (ds.nir + ds.blue) + AF_norm = (AF - AF.min(dim=["y","x"]))/(AF.max(dim=["y","x"]) - AF.min(dim=["y","x"])) + return AF_norm + def ndvi(ds): + return (ds.nir - ds.red) / (ds.nir + ds.red) + def msavi(ds): + return ((2 * ds.nir + 1 - ((2 * ds.nir + 1) ** 2 - 8 * (ds.nir - ds.red)) ** 0.5) / 2 ) + def vsf(ds): + NDVI = ndvi(ds) + MSAVI = msavi(ds) + VSF = 1 - NDVI * MSAVI + VSF_norm = (VSF - VSF.min(dim=["y","x"]))/(VSF.max(dim=["y","x"]) - VSF.min(dim=["y","x"])) + return VSF_norm + def mbi(ds): + return ((ds.swir_1 - ds.swir_2 - ds.nir) / (ds.swir_1 + ds.swir_2 + ds.nir)) + 0.5 + def embi(ds): + MBI = mbi(ds) + MNDWI = mndwi(ds) + return (MBI - MNDWI - 0.5) / (MBI + MNDWI + 1.5) + def ssf(ds): + EMBI = embi(ds) + SSF = 1 - EMBI + SSF_norm = (SSF - SSF.min(dim=["y","x"]))/(SSF.max(dim=["y","x"]) - SSF.min(dim=["y","x"])) + return SSF_norm + # Overall modulation using the Modulation Factor (MF). + def mf(ds): + MF = ((ds.blue + ds.green) - (ds.nir + ds.swir_1)) / ((ds.blue + ds.green) + (ds.nir + ds.swir_1)) + MF_norm = (MF - MF.min(dim=["y","x"]))/(MF.max(dim=["y","x"]) - MF.min(dim=["y","x"])) + return MF_norm + def ASI(ds): + AF = af(ds) + VSF = vsf(ds) + SSF = ssf(ds) + MF = mf(ds) + return AF * VSF * SSF * MF + + index_dict["ASI"] = ASI + + # If index supplied is not a list, convert to list. This allows us to + # iterate through either multiple or single indices in the loop below + indices = index if isinstance(index, list) else [index] + + # calculate for each index in the list of indices supplied (indexes) + for index in indices: + + # Select an index function from the dictionary + index_func = index_dict.get(str(index)) + + # If no index is provided or if no function is returned due to an + # invalid option being provided, raise an exception informing user to + # choose from the list of valid options + if index is None: + + raise ValueError( + f"No remote sensing `index` was provided. Please " + "refer to the function \ndocumentation for a full " + "list of valid options for `index` (e.g. 'NDVI')" + ) + + elif ( + index + in [ + "WI", + "BAEI", + "AWEI_ns", + "AWEI_sh", + "EVI", + "LAI", + "SAVI", + "MSAVI", + ] + and not normalise + ): + + warnings.warn( + f"\nA coefficient-based index ('{index}') normally " + "applied to surface reflectance values in the \n" + "0.0-1.0 range was applied to values in the 0-10000 " + "range. This can produce unexpected results; \nif " + "required, resolve this by setting `normalise=True`" + ) + + elif index_func is None: + + raise ValueError( + f"The selected index '{index}' is not one of the " + "valid remote sensing index options. \nPlease " + "refer to the function documentation for a full " + "list of valid options for `index`" + ) + + # Deprecation warning if `collection` is specified instead of `satellite_mission`. + if collection is not None: + warnings.warn('`collection` was deprecated in version 0.1.7. Use `satelite_mission` instead.', + DeprecationWarning, + stacklevel=2) + # Map the collection values to the valid satellite_mission values. + if collection == "c2": + satellite_mission = "ls" + elif collection == "s2": + satellite_mission = "s2" + # Raise error if no valid collection name is provided: + else: + raise ValueError( + f"'{collection}' is not a valid option for " + "`collection`. Please specify either \n" + "'c2' or 's2'.") + + + # Rename bands to a consistent format if depending on what satellite mission + # is specified in `satellite_mission`. This allows the same index calculations + # to be applied to all satellite missions. If no satellite mission was provided, + # raise an exception. + if satellite_mission is None: + + raise ValueError( + "No `satellite_mission` was provided. Please specify " + "either 'ls' or 's2' to ensure the \nfunction " + "calculates indices using the correct spectral " + "bands." + ) + + elif satellite_mission == "ls": + sr_max = 1.0 + # Dictionary mapping full data names to simpler alias names + # This only applies to properly-scaled "ls" data i.e. from + # the Landsat geomedians. calculate_indices will not show + # correct output for raw (unscaled) Landsat data (i.e. default + # outputs from dc.load) + bandnames_dict = { + "SR_B1": "blue", + "SR_B2": "green", + "SR_B3": "red", + "SR_B4": "nir", + "SR_B5": "swir_1", + "SR_B7": "swir_2", + } + + # Rename bands in dataset to use simple names (e.g. 'red') + bands_to_rename = { + a: b for a, b in bandnames_dict.items() if a in ds.variables + } + + elif satellite_mission == "s2": + sr_max = 10000 + # Dictionary mapping full data names to simpler alias names + bandnames_dict = { + "nir_1": "nir", + "B02": "blue", + "B03": "green", + "B04": "red", + "B05": "red_edge_1", + "B06": "red_edge_2", + "B07": "red_edge_3", + "B08": "nir", + "B11": "swir_1", + "B12": "swir_2", + } + + # Rename bands in dataset to use simple names (e.g. 'red') + bands_to_rename = { + a: b for a, b in bandnames_dict.items() if a in ds.variables + } + + # Raise error if no valid satellite_mission name is provided: + else: + raise ValueError( + f"'{satellite_mission}' is not a valid option for " + "`satellite_mission`. Please specify either \n" + "'ls' or 's2'" + ) + + # Apply index function + try: + # If normalised=True, divide data by 10,000 before applying func + mult = sr_max if normalise else 1.0 + index_array = index_func(ds.rename(bands_to_rename) / mult) + + except AttributeError: + raise ValueError( + f"Please verify that all bands required to " + f"compute {index} are present in `ds`." + ) + + # Add as a new variable in dataset + output_band_name = custom_varname if custom_varname else index + ds[output_band_name] = index_array + + # Once all indexes are calculated, drop input bands if drop=True + if drop: + ds = ds.drop(bands_to_drop) + + # Return input dataset with added water index variable + return ds + +def dualpol_indices( + ds, + co_pol='vv', + cross_pol='vh', + index=None, + custom_varname=None, + drop=False, + deep_copy=True, +): + """ + Takes an xarray dataset containing dual-polarization radar backscatter, + calculates one or a set of indices, and adds the resulting array as a + new variable in the original dataset. + + Last modified: July 2021 + + Parameters + ---------- + ds : xarray Dataset + A two-dimensional or multi-dimensional array containing the + two polarization bands. + + co_pol: str + Measurement name for the co-polarization band. + Default is 'vv' for Sentinel-1. + + cross_pol: str + Measurement name for the cross-polarization band. + Default is 'vh' for Sentinel-1. + + index : str or list of strs + A string giving the name of the index to calculate or a list of + strings giving the names of the indices to calculate: + + * ``'RVI'`` (Radar Vegetation Index for dual-pol, Trudel et al. 2012; Nasirzadehdizaji et al., 2019; Gururaj et al., 2019) + * ``'VDDPI'`` (Vertical dual depolarization index, Periasamy 2018) + * ``'theta'`` (pseudo scattering-type, Bhogapurapu et al. 2021) + * ``'entropy'`` (pseudo scattering entropy, Bhogapurapu et al. 2021) + * ``'purity'`` (co-pol purity, Bhogapurapu et al. 2021) + * ``'ratio'`` (cross-pol/co-pol ratio) + + custom_varname : str, optional + By default, the original dataset will be returned with + a new index variable named after `index` (e.g. 'RVI'). To + specify a custom name instead, you can supply e.g. + `custom_varname='custom_name'`. Defaults to None, which uses + `index` to name the variable. + + drop : bool, optional + Provides the option to drop the original input data, thus saving + space. If `drop=True`, returns only the index and its values. + + deep_copy: bool, optional + If `deep_copy=False`, calculate_indices will modify the original + array, adding bands to the input dataset and not removing them. + If the calculate_indices function is run more than once, variables + may be dropped incorrectly producing unexpected behaviour. This is + a bug and may be fixed in future releases. This is only a problem + when `drop=True`. + + Returns + ------- + ds : xarray Dataset + The original xarray Dataset inputted into the function, with a + new varible containing the remote sensing index as a DataArray. + If drop = True, the new variable/s as DataArrays in the + original Dataset. + """ + + if not co_pol in list(ds.data_vars): + raise ValueError(f"{co_pol} measurement is not in the dataset") + if not cross_pol in list(ds.data_vars): + raise ValueError(f"{cross_pol} measurement is not in the dataset") + + # Set ds equal to a copy of itself in order to prevent the function + # from editing the input dataset. This is to prevent unexpected + # behaviour though it uses twice as much memory. + if deep_copy: + ds = ds.copy(deep=True) + + # Capture input band names in order to drop these if drop=True + if drop: + bands_to_drop = list(ds.data_vars) + print(f"Dropping bands {bands_to_drop}") + + def ratio(ds): + return ds[cross_pol] / ds[co_pol] + + def purity(ds): + return (1 - ratio(ds)) / (1 + ratio(ds)) + + def theta(ds): + return np.arctan((1 - ratio(ds))**2 / (1 + ratio(ds)**2 - ratio(ds))) + + def P1(ds): + return 1 / (1 + ratio(ds)) + + def P2(ds): + return 1 - P1(ds) + + def entropy(ds): + return P1(ds)*np.log2(P1(ds)) + P2(ds)*np.log2(P2(ds)) + + # Dictionary containing remote sensing index band recipes + index_dict = { + # Radar Vegetation Index for dual-pol, Trudel et al. 2012 + "RVI": lambda ds: 4*ds[cross_pol] / (ds[co_pol] + ds[cross_pol]), + # Vertical dual depolarization index, Periasamy 2018 + "VDDPI": lambda ds: (ds[co_pol] + ds[cross_pol]) / ds[co_pol], + # cross-pol/co-pol ratio + "ratio": ratio, + # co-pol purity, Bhogapurapu et al. 2021 + "purity": purity, + # pseudo scattering-type, Bhogapurapu et al. 2021 + "theta": theta, + # pseudo scattering entropy, Bhogapurapu et al. 2021 + "entropy": entropy, + } + + # If index supplied is not a list, convert to list. This allows us to + # iterate through either multiple or single indices in the loop below + indices = index if isinstance(index, list) else [index] + + # calculate for each index in the list of indices supplied (indexes) + for index in indices: + + # Select an index function from the dictionary + index_func = index_dict.get(str(index)) + + # If no index is provided or if no function is returned due to an + # invalid option being provided, raise an exception informing user to + # choose from the list of valid options + if index is None: + + raise ValueError( + f"No radar `index` was provided. Please " + "refer to the function \ndocumentation for a full " + "list of valid options for `index` (e.g. 'RVI')" + ) + + elif index_func is None: + + raise ValueError( + f"The selected index '{index}' is not one of the " + "valid remote sensing index options. \nPlease " + "refer to the function documentation for a full " + "list of valid options for `index`" + ) + + # Apply index function + index_array = index_func(ds) + + # Add as a new variable in dataset + output_band_name = custom_varname if custom_varname else index + ds[output_band_name] = index_array + + # Once all indexes are calculated, drop input bands if drop=True + if drop: + ds = ds.drop(bands_to_drop) + + # Return input dataset with added water index variable + return ds diff --git a/deafrica_tools/classification.py b/deafrica_tools/classification.py new file mode 100644 index 0000000..027ba95 --- /dev/null +++ b/deafrica_tools/classification.py @@ -0,0 +1,1674 @@ +""" +Machine learning functions for classification of remote sensing data contained +in an Open Data Cube instance. +""" + +import multiprocessing as mp +import os +import sys +import time +import warnings +from abc import ABCMeta, abstractmethod +from copy import deepcopy + +import dask.array as da +import dask.distributed as dd +import joblib +import numpy as np +import pandas as pd +import xarray as xr +from dask_ml.wrappers import ParallelPostFit +from datacube.utils import geometry +from datacube.utils.geometry import assign_crs +from deafrica_tools.spatial import xr_rasterize +from sklearn.base import ClusterMixin +from sklearn.cluster import AgglomerativeClustering +from sklearn.cluster import KMeans +from sklearn.mixture import GaussianMixture +from sklearn.model_selection import BaseCrossValidator +from sklearn.model_selection import KFold, ShuffleSplit +from sklearn.utils import check_random_state +from tqdm.auto import tqdm + + +def sklearn_flatten(input_xr): + """ + Reshape a DataArray or Dataset with spatial (and optionally + temporal) structure into an np.array with the spatial and temporal + dimensions flattened into one dimension. + + This flattening procedure enables DataArrays and Datasets to be used + to train and predict with sklearn models. + + Last modified: September 2019 + + Parameters + ---------- + input_xr : xarray.DataArray or xarray.Dataset + Must have dimensions 'x' and 'y', may have dimension 'time'. + Dimensions other than 'x', 'y' and 'time' are unaffected by the + flattening. + + Returns + ---------- + input_np : numpy.array + A numpy array corresponding to input_xr.data (or + input_xr.to_array().data), with dimensions 'x','y' and 'time' + flattened into a single dimension, which is the first axis of + the returned array. input_np contains no NaNs. + + """ + # cast input Datasets to DataArray + if isinstance(input_xr, xr.Dataset): + input_xr = input_xr.to_array() + + # stack across pixel dimensions, handling timeseries if necessary + if "time" in input_xr.dims: + stacked = input_xr.stack(z=["x", "y", "time"]) + else: + stacked = input_xr.stack(z=["x", "y"]) + + # finding 'bands' dimensions in each pixel - these will not be + # flattened as their context is important for sklearn + pxdims = [] + for dim in stacked.dims: + if dim != "z": + pxdims.append(dim) + + # mask NaNs - we mask pixels with NaNs in *any* band, because + # sklearn cannot accept NaNs as input + mask = np.isnan(stacked) + if len(pxdims) != 0: + mask = mask.any(dim=pxdims) + + # turn the mask into a numpy array (boolean indexing with xarrays + # acts weird) + mask = mask.data + + # the dimension we are masking along ('z') needs to be the first + # dimension in the underlying np array for the boolean indexing to work + stacked = stacked.transpose("z", *pxdims) + input_np = stacked.data[~mask] + + return input_np + + +def sklearn_unflatten(output_np, input_xr): + """ + Reshape a numpy array with no 'missing' elements (NaNs) and + 'flattened' spatiotemporal structure into a DataArray matching the + spatiotemporal structure of the DataArray + + This enables an sklearn model's prediction to be remapped to the + correct pixels in the input DataArray or Dataset. + + Last modified: September 2019 + + Parameters + ---------- + output_np : numpy.array + The first dimension's length should correspond to the number of + valid (non-NaN) pixels in input_xr. + input_xr : xarray.DataArray or xarray.Dataset + Must have dimensions 'x' and 'y', may have dimension 'time'. + Dimensions other than 'x', 'y' and 'time' are unaffected by the + flattening. + + Returns + ---------- + output_xr : xarray.DataArray + An xarray.DataArray with the same dimensions 'x', 'y' and 'time' + as input_xr, and the same valid (non-NaN) pixels. These pixels + are set to match the data in output_np. + + """ + + # the output of a sklearn model prediction should just be a numpy array + # with size matching x*y*time for the input DataArray/Dataset. + + # cast input Datasets to DataArray + if isinstance(input_xr, xr.Dataset): + input_xr = input_xr.to_array() + + # generate the same mask we used to create the input to the sklearn model + if "time" in input_xr.dims: + stacked = input_xr.stack(z=["x", "y", "time"]) + else: + stacked = input_xr.stack(z=["x", "y"]) + + pxdims = [] + for dim in stacked.dims: + if dim != "z": + pxdims.append(dim) + + mask = np.isnan(stacked) + if len(pxdims) != 0: + mask = mask.any(dim=pxdims) + + # handle multivariable output + output_px_shape = () + if len(output_np.shape[1:]): + output_px_shape = output_np.shape[1:] + + # use the mask to put the data in all the right places + output_ma = np.ma.empty((len(stacked.z), *output_px_shape)) + output_ma[~mask] = output_np + output_ma[mask] = np.ma.masked + + # set the stacked coordinate to match the input + output_xr = xr.DataArray( + output_ma, + coords={"z": stacked["z"]}, + dims=["z", *["output_dim_" + str(idx) for idx in range(len(output_px_shape))]], + ) + + output_xr = output_xr.unstack() + + return output_xr + + +def fit_xr(model, input_xr): + """ + Utilise our wrappers to fit a vanilla sklearn model. + + Last modified: September 2019 + + Parameters + ---------- + model : scikit-learn model or compatible object + Must have a fit() method that takes numpy arrays. + input_xr : xarray.DataArray or xarray.Dataset. + Must have dimensions 'x' and 'y', may have dimension 'time'. + + Returns + ---------- + model : a scikit-learn model which has been fitted to the data in + the pixels of input_xr. + + """ + + model = model.fit(sklearn_flatten(input_xr)) + return model + + +def predict_xr( + model, + input_xr, + chunk_size=None, + persist=False, + proba=False, + clean=True, + return_input=False, +): + """ + Using dask-ml ParallelPostfit(), runs the parallel + predict and predict_proba methods of sklearn + estimators. Useful for running predictions + on a larger-than-RAM datasets. + + Last modified: September 2020 + + Parameters + ---------- + model : scikit-learn model or compatible object + Must have a .predict() method that takes numpy arrays. + input_xr : xarray.DataArray or xarray.Dataset. + Must have dimensions 'x' and 'y' + chunk_size : int + The dask chunk size to use on the flattened array. If this + is left as None, then the chunks size is inferred from the + .chunks method on the `input_xr` + persist : bool + If True, and proba=True, then 'input_xr' data will be + loaded into distributed memory. This will ensure data + is not loaded twice for the prediction of probabilities, + but this will only work if the data is not larger than + distributed RAM. + proba : bool + If True, predict probabilities + clean : bool + If True, remove Infs and NaNs from input and output arrays + return_input : bool + If True, then the data variables in the 'input_xr' dataset will + be appended to the output xarray dataset. + + Returns + ---------- + output_xr : xarray.Dataset + An xarray.Dataset containing the prediction output from model. + if proba=True then dataset will also contain probabilites, and + if return_input=True then dataset will have the input feature layers. + Has the same spatiotemporal structure as input_xr. + + """ + # if input_xr isn't dask, coerce it + dask = True + if not bool(input_xr.chunks): + dask = False + input_xr = input_xr.chunk({"x": len(input_xr.x), "y": len(input_xr.y)}) + + # set chunk size if not supplied + if chunk_size is None: + chunk_size = int(input_xr.chunks["x"][0]) * int(input_xr.chunks["y"][0]) + + def _predict_func(model, input_xr, persist, proba, clean, return_input): + x, y, crs = input_xr.x, input_xr.y, input_xr.geobox.crs + + input_data = [] + + for var_name in input_xr.data_vars: + input_data.append(input_xr[var_name]) + + input_data_flattened = [] + + for arr in input_data: + data = arr.data.flatten().rechunk(chunk_size) + input_data_flattened.append(data) + + # reshape for prediction + input_data_flattened = da.array(input_data_flattened).transpose() + + if clean == True: + input_data_flattened = da.where( + da.isfinite(input_data_flattened), input_data_flattened, 0 + ) + + if (proba == True) & (persist == True): + # persisting data so we don't require loading all the data twice + input_data_flattened = input_data_flattened.persist() + + # apply the classification + print("predicting...") + out_class = model.predict(input_data_flattened) + + # Mask out NaN or Inf values in results + if clean == True: + out_class = da.where(da.isfinite(out_class), out_class, 0) + + # Reshape when writing out + out_class = out_class.reshape(len(y), len(x)) + + # stack back into xarray + output_xr = xr.DataArray(out_class, coords={"x": x, "y": y}, dims=["y", "x"]) + + output_xr = output_xr.to_dataset(name="Predictions") + + if proba == True: + print(" probabilities...") + out_proba = model.predict_proba(input_data_flattened) + + # convert to % + out_proba = da.max(out_proba, axis=1) * 100.0 + + if clean == True: + out_proba = da.where(da.isfinite(out_proba), out_proba, 0) + + out_proba = out_proba.reshape(len(y), len(x)) + + out_proba = xr.DataArray( + out_proba, coords={"x": x, "y": y}, dims=["y", "x"] + ) + output_xr["Probabilities"] = out_proba + + if return_input == True: + print(" input features...") + # unflatten the input_data_flattened array and append + # to the output_xr containin the predictions + arr = input_xr.to_array() + stacked = arr.stack(z=["y", "x"]) + + # handle multivariable output + output_px_shape = () + if len(input_data_flattened.shape[1:]): + output_px_shape = input_data_flattened.shape[1:] + + output_features = input_data_flattened.reshape( + (len(stacked.z), *output_px_shape) + ) + + # set the stacked coordinate to match the input + output_features = xr.DataArray( + output_features, + coords={"z": stacked["z"]}, + dims=[ + "z", + *["output_dim_" + str(idx) for idx in range(len(output_px_shape))], + ], + ).unstack() + + # convert to dataset and rename arrays + output_features = output_features.to_dataset(dim="output_dim_0") + data_vars = list(input_xr.data_vars) + output_features = output_features.rename( + {i: j for i, j in zip(output_features.data_vars, data_vars)} + ) + + # merge with predictions + output_xr = xr.merge([output_xr, output_features], compat="override") + + return assign_crs(output_xr, str(crs)) + + if dask == True: + # convert model to dask predict + model = ParallelPostFit(model) + with joblib.parallel_backend("dask", wait_for_workers_timeout=20): + output_xr = _predict_func( + model, input_xr, persist, proba, clean, return_input + ) + + else: + output_xr = _predict_func( + model, input_xr, persist, proba, clean, return_input + ).compute() + + return output_xr + + +class HiddenPrints: + """ + For concealing unwanted print statements called by other functions + """ + + def __enter__(self): + self._original_stdout = sys.stdout + sys.stdout = open(os.devnull, "w") + + def __exit__(self, exc_type, exc_val, exc_tb): + sys.stdout.close() + sys.stdout = self._original_stdout + + +def _get_training_data_for_shp( + gdf, + index, + row, + out_arrs, + out_vars, + dc_query, + return_coords, + feature_func=None, + field=None, + zonal_stats=None, +): + """ + This is the core function that is triggered by `collect_training_data`. + The `collect_training_data` function loops through geometries in a geopandas + geodataframe and runs the code within `_get_training_data_for_shp`. + Parameters are inherited from `collect_training_data`. + See that function for information on the other params not listed below. + + Parameters + ---------- + index, row : iterables inherited from geopandas object + out_arrs : list + An empty list into which the training data arrays are stored. + out_vars : list + An empty list into which the data varaible names are stored. + + + Returns + -------- + Two lists, a list of numpy.arrays containing classes and extracted data for + each pixel or polygon, and another containing the data variable names. + + """ + + # prevent function altering dictionary kwargs + dc_query = deepcopy(dc_query) + + # remove dask chunks if supplied as using + # mulitprocessing for parallization + if "dask_chunks" in dc_query.keys(): + dc_query.pop("dask_chunks", None) + + # set up query based on polygon + geom = geometry.Geometry(geom=gdf.iloc[index].geometry, crs=gdf.crs) + q = {"geopolygon": geom} + + # merge polygon query with user supplied query params + dc_query.update(q) + + # Use input feature function + data = feature_func(dc_query) + + # create polygon mask + mask = xr_rasterize(gdf.iloc[[index]], data) + data = data.where(mask) + + # Check that feature_func has removed time + if "time" in data.dims: + t = data.dims["time"] + if t > 1: + raise ValueError( + "After running the feature_func, the dataset still has " + + str(t) + + " time-steps, dataset must only have" + + " x and y dimensions." + ) + + if return_coords == True: + # turn coords into a variable in the ds + data["x_coord"] = data.x + 0 * data.y + data["y_coord"] = data.y + 0 * data.x + + # append ID measurement to dataset for tracking failures + band = [m for m in data.data_vars][0] + _id = xr.zeros_like(data[band]) + data["id"] = _id + data["id"] = data["id"] + gdf.iloc[index]["id"] + + # If no zonal stats were requested then extract all pixel values + if zonal_stats is None: + flat_train = sklearn_flatten(data) + flat_val = np.repeat(row[field], flat_train.shape[0]) + stacked = np.hstack((np.expand_dims(flat_val, axis=1), flat_train)) + + elif zonal_stats in ["mean", "median", "max", "min"]: + method_to_call = getattr(data, zonal_stats) + flat_train = method_to_call() + flat_train = flat_train.to_array() + stacked = np.hstack((row[field], flat_train)) + + else: + raise Exception( + zonal_stats + + " is not one of the supported" + + " reduce functions ('mean','median','max','min')" + ) + + out_arrs.append(stacked) + out_vars.append([field] + list(data.data_vars)) + + +def _get_training_data_parallel( + gdf, dc_query, ncpus, return_coords, feature_func=None, field=None, zonal_stats=None +): + """ + Function passing the '_get_training_data_for_shp' function + to a mulitprocessing.Pool. + Inherits variables from 'collect_training_data'. + + """ + # Check if dask-client is running + try: + zx = None + zx = dd.get_client() + except: + pass + + if zx is not None: + raise ValueError( + "You have a Dask Client running, which prevents \n" + "this function from multiprocessing. Close the client." + ) + + # instantiate lists that can be shared across processes + manager = mp.Manager() + results = manager.list() + column_names = manager.list() + + # progress bar + pbar = tqdm(total=len(gdf)) + + def update(*a): + pbar.update() + + with mp.Pool(ncpus) as pool: + for index, row in gdf.iterrows(): + pool.apply_async( + _get_training_data_for_shp, + [ + gdf, + index, + row, + results, + column_names, + dc_query, + return_coords, + feature_func, + field, + zonal_stats, + ], + callback=update, + ) + + pool.close() + pool.join() + pbar.close() + + return column_names, results + + +def collect_training_data( + gdf, + dc_query, + ncpus=1, + return_coords=False, + feature_func=None, + field=None, + zonal_stats=None, + clean=True, + fail_threshold=0.02, + fail_ratio=0.5, + max_retries=3, +): + """ + This function provides methods for gathering training data from the ODC over + geometries stored within a geopandas geodataframe. The function will return a + 'model_input' array containing stacked training data arrays with all NaNs & Infs removed. + In the instance where ncpus > 1, a parallel version of the function will be run + (functions are passed to a mp.Pool()). This function can conduct zonal statistics if + the supplied shapefile contains polygons. The 'feature_func' parameter defines what + features to produce. + + Parameters + ---------- + gdf : geopandas geodataframe + geometry data in the form of a geopandas geodataframe + dc_query : dictionary + Datacube query object, should not contain lat and long (x or y) + variables as these are supplied by the 'gdf' variable + ncpus : int + The number of cpus/processes over which to parallelize the gathering + of training data (only if ncpus is > 1). Use 'mp.cpu_count()' to determine the number of + cpus available on a machine. Defaults to 1. + return_coords : bool + If True, then the training data will contain two extra columns 'x_coord' and + 'y_coord' corresponding to the x,y coordinate of each sample. This variable can + be useful for handling spatial autocorrelation between samples later in the ML workflow. + feature_func : function + A function for generating feature layers that is applied to the data within + the bounds of the input geometry. The 'feature_func' must accept a 'dc_query' + object, and return a single xarray.Dataset or xarray.DataArray containing + 2D coordinates (i.e x and y, without a third dimension). + e.g.:: + + def feature_function(query): + dc = datacube.Datacube(app='feature_layers') + ds = dc.load(**query) + ds = ds.mean('time') + return ds + + field : str + Name of the column in the gdf that contains the class labels + zonal_stats : string, optional + An optional string giving the names of zonal statistics to calculate + for each polygon. Default is None (all pixel values are returned). Supported + values are 'mean', 'median', 'max', 'min'. + clean : bool + Whether or not to remove missing values in the training dataset. If True, + training labels with any NaNs or Infs in the feature layers will be dropped + from the dataset. + fail_threshold : float, default 0.02 + Silent read fails on S3 during mulitprocessing can result in some rows + of the returned data containing NaN values. + The'fail_threshold' fraction specifies a % of acceptable fails. + e.g. Setting 'fail_threshold' to 0.05 means if >5% of the samples in the training dataset + fail then those samples will be returned to the multiprocessing queue. Below this fraction + the function will accept the failures and return the results. + fail_ratio: float + A float between 0 and 1 that defines if a given training sample has failed. + Default is 0.5, which means if 50 % of the measurements in a given sample return null + values, and the number of total fails is more than the 'fail_threshold', the sample + will be passed to the retry queue. + max_retries: int, default 3 + Maximum number of times to retry collecting samples. This number is invoked + if the 'fail_threshold' is not reached. + + Returns + -------- + Two objects are returned: + `columns_names`: a list of variable (feature) names + `model_input`: a numpy.array containing the data values for each feature extracted + + """ + + # check the dtype of the class field + if gdf[field].dtype != int: + raise ValueError( + 'The "field" column of the input vector must contain integer dtypes' + ) + + # set up some print statements + if feature_func is None: + raise ValueError( + "Please supply a feature layer function through the " + +"parameter 'feature_func'" + ) + + if zonal_stats is not None: + print("Taking zonal statistic: " + zonal_stats) + + # add unique id to gdf to help with indexing failed rows + # during multiprocessing + # if zonal_stats is not None: + gdf["id"] = range(0, len(gdf)) + + if ncpus == 1: + # progress indicator + print("Collecting training data in serial mode") + i = 0 + + # list to store results + results = [] + column_names = [] + + # loop through polys and extract training data + for index, row in gdf.iterrows(): + print(" Feature {:04}/{:04}\r".format(i + 1, len(gdf)), end="") + + _get_training_data_for_shp( + gdf, + index, + row, + results, + column_names, + dc_query, + return_coords, + feature_func, + field, + zonal_stats, + ) + i += 1 + + else: + print("Collecting training data in parallel mode") + column_names, results = _get_training_data_parallel( + gdf=gdf, + dc_query=dc_query, + ncpus=ncpus, + return_coords=return_coords, + feature_func=feature_func, + field=field, + zonal_stats=zonal_stats, + ) + + # column names are appended during each iteration + # but they are identical, grab only the first instance + column_names = column_names[0] + + # Stack the extracted training data for each feature into a single array + model_input = np.vstack(results) + + # this code block below iteratively retries failed rows + # up to max_retries or until fail_threshold is + # reached - whichever occurs first + if ncpus > 1: + i = 1 + while i <= max_retries: + # Find % of fails (null values) in data. Use Pandas for simplicity + df = pd.DataFrame(data=model_input[:, 0:-1], index=model_input[:, -1]) + # how many nan values per id? + num_nans = df.isnull().sum(axis=1) + num_nans = num_nans.groupby(num_nans.index).sum() + # how many valid values per id? + num_valid = df.notnull().sum(axis=1) + num_valid = num_valid.groupby(num_valid.index).sum() + # find fail rate + perc_fail = num_nans / (num_nans + num_valid) + fail_ids = perc_fail[perc_fail > fail_ratio] + fail_rate = len(fail_ids) / len(gdf) + + print( + "Percentage of possible fails after run " + + str(i) + + " = " + + str(round(fail_rate * 100, 2)) + + " %" + ) + + if fail_rate > fail_threshold: + print("Recollecting samples that failed") + + fail_ids = list(fail_ids.index) + # keep only the ids in model_input object that didn't fail + model_input = model_input[~np.isin(model_input[:, -1], fail_ids)] + + # index out the fail_ids from the original gdf + gdf_rerun = gdf.loc[gdf["id"].isin(fail_ids)] + gdf_rerun = gdf_rerun.reset_index(drop=True) + + time.sleep(5) # sleep for 5s to rest api + + # recollect failed rows + column_names_again, results_again = _get_training_data_parallel( + gdf=gdf_rerun, + dc_query=dc_query, + ncpus=ncpus, + return_coords=return_coords, + feature_func=feature_func, + field=field, + zonal_stats=zonal_stats, + ) + + # Stack the extracted training data for each feature into a single array + model_input_again = np.vstack(results_again) + + # merge results of the re-run with original run + model_input = np.vstack((model_input, model_input_again)) + + i += 1 + + else: + break + + # ----------------------------------------------- + + # remove id column + idx_var = column_names[0:-1] + model_col_indices = [column_names.index(var_name) for var_name in idx_var] + model_input = model_input[:, model_col_indices] + + if clean == True: + num = np.count_nonzero(np.isnan(model_input).any(axis=1)) + model_input = model_input[~np.isnan(model_input).any(axis=1)] + model_input = model_input[~np.isinf(model_input).any(axis=1)] + print("Removed " + str(num) + " rows wth NaNs &/or Infs") + print("Output shape: ", model_input.shape) + + else: + print("Returning data without cleaning") + print("Output shape: ", model_input.shape) + + return column_names[0:-1], model_input + + +class KMeans_tree(ClusterMixin): + """ + A hierarchical KMeans unsupervised clustering model. This class is + a clustering model, so it inherits scikit-learn's ClusterMixin + base class. + + Parameters + ---------- + n_levels : integer, default 2 + number of levels in the tree of clustering models. + n_clusters : integer, default 3 + Number of clusters in each of the constituent KMeans models in + the tree. + **kwargs : optional + Other keyword arguments to be passed directly to the KMeans + initialiser. + + """ + + def __init__(self, n_levels=2, n_clusters=3, **kwargs): + + assert n_levels >= 1 + + self.base_model = KMeans(n_clusters=3, **kwargs) + self.n_levels = n_levels + self.n_clusters = n_clusters + # make child models + if n_levels > 1: + self.branches = [ + KMeans_tree(n_levels=n_levels - 1, n_clusters=n_clusters, **kwargs) + for _ in range(n_clusters) + ] + + def fit(self, X, y=None, sample_weight=None): + """ + Fit the tree of KMeans models. All parameters mimic those + of KMeans.fit(). + + Parameters + ---------- + X : array-like or sparse matrix, shape=(n_samples, n_features) + Training instances to cluster. It must be noted that the + data will be converted to C ordering, which will cause a + memory copy if the given data is not C-contiguous. + y : Ignored + not used, present here for API consistency by convention. + sample_weight : array-like, shape (n_samples,), optional + The weights for each observation in X. If None, all + observations are assigned equal weight (default: None) + """ + + self.labels_ = self.base_model.fit(X, sample_weight=sample_weight).labels_ + + if self.n_levels > 1: + labels_old = np.copy(self.labels_) + # make room to add the sub-cluster labels + self.labels_ *= (self.n_clusters) ** (self.n_levels - 1) + + for clu in range(self.n_clusters): + # fit child models on their corresponding partition of the training set + self.branches[clu].fit( + X[labels_old == clu], + sample_weight=( + sample_weight[labels_old == clu] + if sample_weight is not None + else None + ), + ) + self.labels_[labels_old == clu] += self.branches[clu].labels_ + + return self + + def predict(self, X, sample_weight=None): + """ + Send X through the KMeans tree and predict the resultant + cluster. Compatible with KMeans.predict(). + + Parameters + ---------- + X : {array-like, sparse matrix}, shape = [n_samples, n_features] + New data to predict. + sample_weight : array-like, shape (n_samples,), optional + The weights for each observation in X. If None, all + observations are assigned equal weight (default: None) + + Returns + ------- + labels : array, shape [n_samples,] + Index of the cluster each sample belongs to. + """ + + result = self.base_model.predict(X, sample_weight=sample_weight) + + if self.n_levels > 1: + rescpy = np.copy(result) + + # make room to add the sub-cluster labels + result *= (self.n_clusters) ** (self.n_levels - 1) + + for clu in range(self.n_clusters): + result[rescpy == clu] += self.branches[clu].predict( + X[rescpy == clu], + sample_weight=( + sample_weight[rescpy == clu] + if sample_weight is not None + else None + ), + ) + + return result + + +def spatial_clusters( + coordinates, + method="Hierarchical", + max_distance=None, + n_groups=None, + verbose=False, + **kwargs +): + """ + Create spatial groups on coorindate data using either KMeans clustering + or a Gaussian Mixture model + + Last modified: September 2020 + + Parameters + ---------- + n_groups : int + The number of groups to create. This is passed as ``n_clusters=n_groups`` + for the KMeans algo, and ``n_components=n_groups`` for the GMM. If using + method=``'Hierarchical'`` then this parameter is ignored. + coordinates : np.array + A numpy array of coordinate values e.g.:: + + np.array([[3337270., 262400.], + [3441390., -273060.], ...]) + + method : str + Which algorithm to use to seperate data points. + Either ``'KMeans'``, ``'GMM'``, or ``'Hierarchical'``. + If using ``'Hierarchical'`` then must set max_distance. + max_distance : int + If method is set to ``'Hierarchical'`` then maximum distance describes the + maximum euclidean distances between all observations in a cluster. 'n_groups' + is ignored in this case. + **kwargs : optional, + Additional keyword arguments to pass to ``sklearn.cluster.Kmeans`` or + ``sklearn.mixture.GuassianMixture`` depending on the 'method' argument. + Returns + ------- + labels : array, shape [n_samples,] + Index of the cluster each sample belongs to. + """ + if method not in ["Hierarchical", "KMeans", "GMM"]: + raise ValueError("method must be one of: 'Hierarchical','KMeans' or 'GMM'") + + if (method in ["GMM", "KMeans"]) & (n_groups is None): + raise ValueError( + "The 'GMM' and 'KMeans' methods requires explicitly setting 'n_groups'" + ) + + if (method == "Hierarchical") & (max_distance is None): + raise ValueError("The 'Hierarchical' method requires setting max_distance") + + if method == "Hierarchical": + cluster_label = AgglomerativeClustering( + n_clusters=None, + linkage="complete", + distance_threshold=max_distance, + **kwargs + ).fit_predict(coordinates) + + if method == "KMeans": + cluster_label = KMeans(n_clusters=n_groups, **kwargs).fit_predict(coordinates) + + if method == "GMM": + cluster_label = GaussianMixture(n_components=n_groups, **kwargs).fit_predict( + coordinates + ) + if verbose: + print("n clusters = " + str(len(np.unique(cluster_label)))) + + return cluster_label + + +def SKCV( + coordinates, + n_splits, + cluster_method, + kfold_method, + test_size, + balance, + n_groups=None, + max_distance=None, + train_size=None, + random_state=None, + **kwargs +): + """ + Generate spatial k-fold cross validation indices using coordinate data. + + This function wraps the ``SpatialShuffleSplit`` and ``SpatialKFold`` classes. + These classes ingest coordinate data in the form of an + ``np.array([[eastings, northings]])`` and assign samples to a spatial cluster + using either a KMeans, Gaussain Mixture, or Agglomerative Clustering algorithm. + This cross-validator is preferred over other sklearn.model_selection methods + for spatial data to avoid overestimating cross-validation scores. + This can happen because of the inherent spatial autocorrelation that is usually + associated with this type of data. + + Last modified: Dec 2020 + + Parameters + ---------- + coordinates : np.array + A numpy array of coordinate values e.g.:: + + np.array([[3337270., 262400.], + [3441390., -273060.], ...]) + + n_splits : int + The number of test-train cross validation splits to generate. + cluster_method : str + Which algorithm to use to separate data points. Either ``'KMeans'``, + ``'GMM'``, or ``'Hierarchical'`` + kfold_method : str + One of either ``'SpatialShuffleSplit'`` or ``'SpatialKFold'``. See the docs + under class:_SpatialShuffleSplit and class:_SpatialKFold for more + information on these options. + test_size : float, int, None + If float, should be between 0.0 and 1.0 and represent the proportion + of the dataset to include in the test split. If int, represents the + absolute number of test samples. If None, the value is set to the + complement of the train size. If ``train_size`` is also None, it will + be set to 0.15. + balance : int or bool + if setting kfold_method to ``'SpatialShuffleSplit'``: int + The number of splits generated per iteration to try to balance the + amount of data in each set so that *test_size* and *train_size* are + respected. If 1, then no extra splits are generated (essentially + disabling the balacing). Must be >= 1. + + if setting kfold_method to ``'SpatialKFold'``: bool + Whether or not to split clusters into fold with approximately equal + number of data points. If False, each fold will have the same number of + clusters (which can have different number of data points in them). + + n_groups : int + The number of groups to create. This is passed as 'n_clusters=n_groups' + for the KMeans algo, and 'n_components=n_groups' for the GMM. If using + cluster_method='Hierarchical' then this parameter is ignored. + max_distance : int + If method is set to 'hierarchical' then maximum distance describes the + maximum euclidean distances between all observations in a cluster. 'n_groups' + is ignored in this case. + train_size : float, int, or None + If float, should be between 0.0 and 1.0 and represent the + proportion of the dataset to include in the train split. If + int, represents the absolute number of train samples. If None, + the value is automatically set to the complement of the test size. + random_state : int, RandomState instance or None, optional (default=None) + If int, random_state is the seed used by the random number generator; + If RandomState instance, random_state is the random number generator; + If None, the random number generator is the RandomState instance used + by ``np.random``. + **kwargs : optional, + Additional keyword arguments to pass to sklearn.cluster.Kmeans or + sklearn.mixture.GuassianMixture depending on the cluster_method argument. + Returns + -------- + generator object _BaseSpatialCrossValidator.split + + """ + # intiate a method + if kfold_method == "SpatialShuffleSplit": + splitter = _SpatialShuffleSplit( + n_groups=n_groups, + method=cluster_method, + coordinates=coordinates, + max_distance=max_distance, + test_size=test_size, + train_size=train_size, + n_splits=n_splits, + random_state=random_state, + balance=balance, + **kwargs + ) + + if kfold_method == "SpatialKFold": + splitter = _SpatialKFold( + n_groups=n_groups, + coordinates=coordinates, + max_distance=max_distance, + method=cluster_method, + test_size=test_size, + n_splits=n_splits, + random_state=random_state, + balance=balance, + **kwargs + ) + + return splitter + + +def spatial_train_test_split( + X, + y, + coordinates, + cluster_method, + kfold_method, + balance, + test_size=None, + n_splits=None, + n_groups=None, + max_distance=None, + train_size=None, + random_state=None, + **kwargs +): + """ + Split arrays into random train and test subsets. Similar to + `sklearn.model_selection.train_test_split` but instead works on + spatial coordinate data. Coordinate data is grouped according + to either a KMeans, Gaussain Mixture, or Agglomerative Clustering algorthim. + Grouping by spatial clusters is preferred over plain random splits for + spatial data to avoid overestimating validation scores due to spatial + autocorrelation. + + Parameters + ---------- + X : np.array + Training data features + y : np.array + Training data labels + coordinates : np.array + A numpy array of coordinate values e.g.:: + + np.array([[3337270., 262400.], + [3441390., -273060.], ...]) + + cluster_method : str + Which algorithm to use to seperate data points. + Either ``'KMeans'``, ``'GMM'``, or ``'Hierarchical'`` + kfold_method : str + One of either ``'SpatialShuffleSplit'`` or ``'SpatialKFold'``. + See the docs under class:_SpatialShuffleSplit and + class: _SpatialKFold for more information on these options. + balance : int or bool + if setting kfold_method to ''`SpatialShuffleSplit`'': int + The number of splits generated per iteration to try to balance the + amount of data in each set so that *test_size* and *train_size* are + respected. If 1, then no extra splits are generated (essentially + disabling the balacing). Must be >= 1. + + if setting kfold_method to ''`SpatialKFold`'': bool + Whether or not to split clusters into fold with approximately equal + number of data points. If False, each fold will have the same number of + clusters (which can have different number of data points in them). + + test_size : float, int, None + If float, should be between 0.0 and 1.0 and represent the proportion + of the dataset to include in the test split. If int, represents the + absolute number of test samples. If None, the value is set to the + complement of the train size. If ``train_size`` is also None, it will + be set to 0.15. + n_splits : int + This parameter is invoked for the 'SpatialKFold' folding method, use this + number to satisfy the train-test size ratio desired, as the 'test_size' + parameter for the KFold method often fails to get the ratio right. + n_groups : int + The number of groups to create. This is passed as 'n_clusters=n_groups' + for the KMeans algo, and 'n_components=n_groups' for the GMM. If using + cluster_method='Hierarchical' then this parameter is ignored. + max_distance : int + If method is set to 'hierarchical' then maximum distance describes the + maximum euclidean distances between all observations in a cluster. 'n_groups' + is ignored in this case. + train_size : float, int, or None + If float, should be between 0.0 and 1.0 and represent the + proportion of the dataset to include in the train split. If + int, represents the absolute number of train samples. If None, + the value is automatically set to the complement of the test size. + random_state : int, + RandomState instance or None, optional + If int, random_state is the seed used by the random number generator; + If RandomState instance, random_state is the random number generator; + If None, the random number generator is the RandomState instance used + by `np.random`. + **kwargs : optional, + Additional keyword arguments to pass to sklearn.cluster.Kmeans or + sklearn.mixture.GuassianMixture depending on the cluster_method argument. + + Returns + ------- + Tuple : + Contains four arrays in the following order: + X_train, X_test, y_train, y_test + + """ + + if kfold_method == "SpatialShuffleSplit": + splitter = _SpatialShuffleSplit( + n_groups=n_groups, + method=cluster_method, + coordinates=coordinates, + max_distance=max_distance, + test_size=test_size, + train_size=train_size, + n_splits=1 if n_splits is None else n_splits, + random_state=random_state, + balance=balance, + **kwargs + ) + + if kfold_method == "SpatialKFold": + if n_splits is None: + raise ValueError( + "n_splits parameter requires an integer value, eg. 'n_splits=5'" + ) + if (test_size is not None) or (train_size is not None): + warnings.warn( + "With the 'SpatialKFold' method, controlling the test/train ratio " + "is better achieved using the 'n_splits' parameter" + ) + + splitter = _SpatialKFold( + n_groups=n_groups, + coordinates=coordinates, + max_distance=max_distance, + method=cluster_method, + n_splits=n_splits, + random_state=random_state, + balance=balance, + **kwargs + ) + + lst = [] + for train, test in splitter.split(coordinates): + X_tr, X_tt = X[train, :], X[test, :] + y_tr, y_tt = y[train], y[test] + lst.extend([X_tr, X_tt, y_tr, y_tt]) + + return (lst[0], lst[1], lst[2], lst[3]) + + +def _partition_by_sum(array, parts): + """ + Partition an array into parts of approximately equal sum. + Does not change the order of the array elements. + Produces the partition indices on the array. Use :func:`numpy.split` to + divide the array along these indices. + Parameters + ---------- + array : array or array-like + The 1D array that will be partitioned. The array will be raveled before + computations. + parts : int + Number of parts to split the array. Can be at most the number of + elements in the array. + Returns + ------- + indices : array + The indices in which the array should be split. + Notes + ----- + Solution from https://stackoverflow.com/a/54024280 + """ + array = np.atleast_1d(array).ravel() + if parts > array.size: + raise ValueError( + "Cannot partition an array of size {} into {} parts of equal sum.".format( + array.size, parts + ) + ) + cumulative_sum = array.cumsum() + # Ideally, we want each part to have the same number of points (total / + # parts). + ideal_sum = cumulative_sum[-1] // parts + # If the parts are ideal, the cumulative sum of each part will be this + ideal_cumsum = np.arange(1, parts) * ideal_sum + indices = np.searchsorted(cumulative_sum, ideal_cumsum, side="right") + # Check for repeated split points, which indicates that there is no way to + # split the array. + if np.unique(indices).size != indices.size: + raise ValueError( + "Could not find partition points to split the array into {} parts " + "of equal sum.".format(parts) + ) + return indices + + +class _BaseSpatialCrossValidator(BaseCrossValidator, metaclass=ABCMeta): + """ + Base class for spatial cross-validators. + Parameters + ---------- + n_groups : int + The number of groups to create. This is passed as 'n_clusters=n_groups' + for the KMeans algo, and 'n_components=n_groups' for the GMM. + coordinates : np.array + A numpy array of coordinate values e.g. + np.array([[3337270., 262400.], + [3441390., -273060.], ..., + method : str + Which algorithm to use to seperate data points. Either 'KMeans' or 'GMM' + n_splits : int + Number of splitting iterations. + """ + + def __init__( + self, + n_groups=None, + coordinates=None, + method=None, + max_distance=None, + n_splits=None, + ): + + self.n_groups = n_groups + self.coordinates = coordinates + self.method = method + self.max_distance = max_distance + self.n_splits = n_splits + + def split(self, X, y=None, groups=None): + """ + Generate indices to split data into training and test set. + Parameters + ---------- + X : array-like, shape (n_samples, 2) + Columns should be the easting and northing coordinates of data + points, respectively. + y : array-like, shape (n_samples,) + The target variable for supervised learning problems. Always + ignored. + groups : array-like, with shape (n_samples,), optional + Group labels for the samples used while splitting the dataset into + train/test set. Always ignored. + Yields + ------ + train : ndarray + The training set indices for that split. + test : ndarray + The testing set indices for that split. + """ + if X.shape[1] != 2: + raise ValueError( + "X (the coordinate data) must have exactly 2 columns ({} given).".format( + X.shape[1] + ) + ) + for train, test in super().split(X, y, groups): + yield train, test + + def get_n_splits(self, X=None, y=None, groups=None): + """ + Returns the number of splitting iterations in the cross-validator + Parameters + ---------- + X : object + Always ignored, exists for compatibility. + y : object + Always ignored, exists for compatibility. + groups : object + Always ignored, exists for compatibility. + Returns + ------- + n_splits : int + Returns the number of splitting iterations in the cross-validator. + """ + return self.n_splits + + @abstractmethod + def _iter_test_indices(self, X=None, y=None, groups=None): + """ + Generates integer indices corresponding to test sets. + MUST BE IMPLEMENTED BY DERIVED CLASSES. + Parameters + ---------- + X : array-like, shape (n_samples, 2) + Columns should be the easting and northing coordinates of data + points, respectively. + y : array-like, shape (n_samples,) + The target variable for supervised learning problems. Always + ignored. + groups : array-like, with shape (n_samples,), optional + Group labels for the samples used while splitting the dataset into + train/test set. Always ignored. + Yields + ------ + test : ndarray + The testing set indices for that split. + """ + + +class _SpatialShuffleSplit(_BaseSpatialCrossValidator): + """ + Random permutation of spatial cross-validator. + Yields indices to split data into training and test sets. Data are first + grouped into clusters using either a KMeans or GMM algorithm + and are then split into testing and training sets randomly. + The proportion of clusters assigned to each set is controlled by *test_size* + and/or *train_size*. However, the total amount of actual data points in + each set could be different from these values since clusters can have + a different number of data points inside them. To guarantee that the + proportion of actual data is as close as possible to the proportion of + clusters, this cross-validator generates an extra number of splits and + selects the one with proportion of data points in each set closer to the + desired amount. The number of balance splits per + iteration is controlled by the *balance* argument. + This cross-validator is preferred over `sklearn.model_selection.ShuffleSplit` + for spatial data to avoid overestimating cross-validation scores. + This can happen because of the inherent spatial autocorrelation. + Parameters + ---------- + n_groups : int + The number of groups to create. This is passed as 'n_clusters=n_groups' + for the KMeans algo, and 'n_components=n_groups' for the GMM. If using + cluster_method='Hierarchical' then this parameter is ignored. + coordinates : np.array + A numpy array of coordinate values e.g. + np.array([[3337270., 262400.], + [3441390., -273060.], ...]) + cluster_method : str + Which algorithm to use to seperate data points. Either 'KMeans', 'GMM', or + 'Hierarchical' + max_distance : int + If method is set to 'hierarchical' then maximum distance describes the + maximum euclidean distances between all observations in a cluster. 'n_groups' + is ignored in this case. + n_splits : int, + Number of re-shuffling & splitting iterations. + test_size : float, int, None + If float, should be between 0.0 and 1.0 and represent the proportion + of the dataset to include in the test split. If int, represents the + absolute number of test samples. If None, the value is set to the + complement of the train size. If ``train_size`` is also None, it will + be set to 0.1. + train_size : float, int, or None + If float, should be between 0.0 and 1.0 and represent the + proportion of the dataset to include in the train split. If + int, represents the absolute number of train samples. If None, + the value is automatically set to the complement of the test size. + random_state : int, RandomState instance or None, optional (default=None) + If int, random_state is the seed used by the random number generator; + If RandomState instance, random_state is the random number generator; + If None, the random number generator is the RandomState instance used + by `np.random`. + balance : int + The number of splits generated per iteration to try to balance the + amount of data in each set so that *test_size* and *train_size* are + respected. If 1, then no extra splits are generated (essentially + disabling the balacing). Must be >= 1. + **kwargs : optional, + Additional keyword arguments to pass to sklearn.cluster.Kmeans or + sklearn.mixture.GuassianMixture depending on the cluster_method argument. + Returns + -------- + generator + containing indices to split data into training and test sets + """ + + def __init__( + self, + n_groups=None, + coordinates=None, + method="Heirachical", + max_distance=None, + n_splits=None, + test_size=0.15, + train_size=None, + random_state=None, + balance=10, + **kwargs + ): + super().__init__( + n_groups=n_groups, + coordinates=coordinates, + method=method, + max_distance=max_distance, + n_splits=n_splits, + **kwargs + ) + if balance < 1: + raise ValueError( + "The *balance* argument must be >= 1. To disable balance, use 1." + ) + self.test_size = test_size + self.train_size = train_size + self.random_state = random_state + self.balance = balance + self.kwargs = kwargs + + def _iter_test_indices(self, X=None, y=None, groups=None): + """ + Generates integer indices corresponding to test sets. + Runs several iterations until a split is found that yields clusters with + the right amount of data points in it. + Parameters + ---------- + X : array-like, shape (n_samples, 2) + Columns should be the easting and northing coordinates of data + points, respectively. + y : array-like, shape (n_samples,) + The target variable for supervised learning problems. Always + ignored. + groups : array-like, with shape (n_samples,), optional + Group labels for the samples used while splitting the dataset into + train/test set. Always ignored. + Yields + ------ + test : ndarray + The testing set indices for that split. + """ + labels = spatial_clusters( + n_groups=self.n_groups, + coordinates=self.coordinates, + method=self.method, + max_distance=self.max_distance, + **self.kwargs + ) + + cluster_ids = np.unique(labels) + # Generate many more splits so that we can pick and choose the ones + # that have the right balance of training and testing data. + shuffle = ShuffleSplit( + n_splits=self.n_splits * self.balance, + test_size=self.test_size, + train_size=self.train_size, + random_state=self.random_state, + ).split(cluster_ids) + + for _ in range(self.n_splits): + test_sets, balance = [], [] + for _ in range(self.balance): + # This is a false positive in pylint which is why the warning + # is disabled at the top of this file: + # https://github.com/PyCQA/pylint/issues/1830 + # pylint: disable=stop-iteration-return + train_clusters, test_clusters = next(shuffle) + # pylint: enable=stop-iteration-return + train_points = np.where(np.isin(labels, cluster_ids[train_clusters]))[0] + test_points = np.where(np.isin(labels, cluster_ids[test_clusters]))[0] + # The proportion of data points assigned to each group should + # be close the proportion of clusters assigned to each group. + balance.append( + abs( + train_points.size / test_points.size + - train_clusters.size / test_clusters.size + ) + ) + test_sets.append(test_points) + best = np.argmin(balance) + yield test_sets[best] + + +class _SpatialKFold(_BaseSpatialCrossValidator): + """ + Spatial K-Folds cross-validator. + Yields indices to split data into training and test sets. Data are first + grouped into clusters using either a KMeans or GMM algorithm + clusters. The clusters are then split into testing and training sets iteratively + along k folds of the data (k is given by *n_splits*). + By default, the clusters are split into folds in a way that makes each fold + have approximately the same number of data points. Sometimes this might not + be possible, which can happen if the number of splits is close to the + number of clusters. In these cases, each fold will have the same number of + clusters regardless of how many data points are in each cluster. This + behaviour can also be disabled by setting ``balance=False``. + This cross-validator is preferred over `sklearn.model_selection.KFold` for + spatial data to avoid overestimating cross-validation scores. This can happen + because of the inherent autocorrelation that is usually associated with + this type of data. + Parameters + ---------- + n_groups : int + The number of groups to create. This is passed as 'n_clusters=n_groups' + for the KMeans algo, and 'n_components=n_groups' for the GMM. If using + cluster_method='Hierarchical' then this parameter is ignored. + coordinates : np.array + A numpy array of coordinate values e.g. + np.array([[3337270., 262400.], + [3441390., -273060.], ...]) + cluster_method : str + Which algorithm to use to seperate data points. Either 'KMeans', 'GMM', or + 'Hierarchical' + max_distance : int + If method is set to 'hierarchical' then maximum distance describes the + maximum euclidean distances between all observations in a cluster. 'n_groups' + is ignored in this case. + n_splits : int + Number of folds. Must be at least 2. + shuffle : bool + Whether to shuffle the data before splitting into batches. + random_state : int, RandomState instance or None, optional (defasult=None) + If int, random_state is the seed used by the random number generator; + If RandomState instance, random_state is the random number generator; + If None, the random number generator is the RandomState instance used + by `np.random`. + balance : bool + Whether or not to split clusters into fold with approximately equal + number of data points. If False, each fold will have the same number of + clusters (which can have different number of data points in them). + **kwargs : optional, + Additional keyword arguments to pass to sklearn.cluster.Kmeans or + sklearn.mixture.GuassianMixture depending on the cluster_method argument. + """ + + def __init__( + self, + n_groups=None, + coordinates=None, + method="Heirachical", + max_distance=None, + n_splits=5, + test_size=0.15, + train_size=None, + shuffle=True, + random_state=None, + balance=True, + **kwargs + ): + super().__init__( + n_groups=n_groups, + coordinates=coordinates, + method=method, + max_distance=max_distance, + n_splits=n_splits, + **kwargs + ) + + if n_splits < 2: + raise ValueError( + "Number of splits must be >=2 for clusterKFold. Given {}.".format( + n_splits + ) + ) + self.test_size = test_size + self.shuffle = shuffle + self.random_state = random_state + self.balance = balance + self.kwargs = kwargs + + def _iter_test_indices(self, X=None, y=None, groups=None): + """ + Generates integer indices corresponding to test sets. + Parameters + ---------- + X : array-like, shape (n_samples, 2) + Columns should be the easting and northing coordinates of data + points, respectively. + y : array-like, shape (n_samples,) + The target variable for supervised learning problems. Always + ignored. + groups : array-like, with shape (n_samples,), optional + Group labels for the samples used while splitting the dataset into + train/test set. Always ignored. + Yields + ------ + test : ndarray + The testing set indices for that split. + """ + labels = spatial_clusters( + n_groups=self.n_groups, + coordinates=self.coordinates, + method=self.method, + max_distance=self.max_distance, + **self.kwargs + ) + + cluster_ids = np.unique(labels) + if self.n_splits > cluster_ids.size: + raise ValueError( + "Number of k-fold splits ({}) cannot be greater than the number of " + "clusters ({}). Either decrease n_splits or increase the number of " + "clusters.".format(self.n_splits, cluster_ids.size) + ) + if self.shuffle: + check_random_state(self.random_state).shuffle(cluster_ids) + if self.balance: + cluster_sizes = [np.isin(labels, i).sum() for i in cluster_ids] + try: + split_points = _partition_by_sum(cluster_sizes, parts=self.n_splits) + folds = np.split(np.arange(cluster_ids.size), split_points) + except ValueError: + warnings.warn( + "Could not balance folds to have approximately the same " + "number of data points. Dividing into folds with equal " + "number of clusters instead. Decreasing n_splits or increasing " + "the number of clusters may help.", + UserWarning, + ) + folds = [i for _, i in KFold(n_splits=self.n_splits).split(cluster_ids)] + else: + folds = [i for _, i in KFold(n_splits=self.n_splits).split(cluster_ids)] + for test_clusters in folds: + test_points = np.where(np.isin(labels, cluster_ids[test_clusters]))[0] + yield test_points diff --git a/deafrica_tools/coastal.py b/deafrica_tools/coastal.py new file mode 100644 index 0000000..c76d3db --- /dev/null +++ b/deafrica_tools/coastal.py @@ -0,0 +1,1123 @@ +""" +Coastal analyses on Digital Earth Africa data. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import requests +import numpy as np +import xarray as xr +import pandas as pd +import geopandas as gpd +import matplotlib.pyplot as plt +from scipy import stats +from otps import TimePoint +from otps import predict_tide +from shapely.geometry import box +from datacube.utils.geometry import CRS +from owslib.wfs import WebFeatureService + +from deafrica_tools.datahandling import parallel_apply + + +# Fix converters for tidal plot +from pandas.plotting import register_matplotlib_converters +register_matplotlib_converters() + + +# URL for the DE Africa Coastlines data on Geoserver. +WFS_ADDRESS = "https://geoserver.digitalearth.africa/geoserver/wfs" + +def model_tides( + x, + y, + time, + model="FES2014", + directory="/var/share/tide_models", + epsg=4326, + method="bilinear", + extrapolate=True, + cutoff=10.0, +): + """ + Compute tides at points and times using tidal harmonics. + If multiple x, y points are provided, tides will be + computed for all timesteps at each point. + + This function supports any tidal model supported by + `pyTMD`, including the FES2014 Finite Element Solution + tide model, and the TPXO8-atlas and TPXO9-atlas-v5 + TOPEX/POSEIDON global tide models. + + This function requires access to tide model data files + to work. These should be placed in a folder with + subfolders matching the formats specified by `pyTMD`: + https://pytmd.readthedocs.io/en/latest/getting_started/Getting-Started.html#directories + + For FES2014 (https://www.aviso.altimetry.fr/es/data/products/auxiliary-products/global-tide-fes/description-fes2014.html): + - {directory}/fes2014/ocean_tide/ + {directory}/fes2014/load_tide/ + + For TPXO8-atlas (https://www.tpxo.net/tpxo-products-and-registration): + - {directory}/tpxo8_atlas/ + + For TPXO9-atlas-v5 (https://www.tpxo.net/tpxo-products-and-registration): + - {directory}/TPXO9_atlas_v5/ + + This function is a minor modification of the `pyTMD` + package's `compute_tide_corrections` function, adapted + to process multiple timesteps for multiple input point + locations. For more info: + https://pytmd.readthedocs.io/en/stable/user_guide/compute_tide_corrections.html + + Parameters: + ----------- + x, y : float or list of floats + One or more x and y coordinates used to define + the location at which to model tides. By default these + coordinates should be lat/lon; use `epsg` if they + are in a custom coordinate reference system. + time : A datetime array or pandas.DatetimeIndex + An array containing 'datetime64[ns]' values or a + 'pandas.DatetimeIndex' providing the times at which to + model tides in UTC time. + model : string + The tide model used to model tides. Options include: + - "FES2014" (only pre-configured option on DEA Sandbox) + - "TPXO8-atlas" + - "TPXO9-atlas-v5" + directory : string + The directory containing tide model data files. These + data files should be stored in sub-folders for each + model that match the structure provided by `pyTMD`: + https://pytmd.readthedocs.io/en/latest/getting_started/Getting-Started.html#directories + For example: + - {directory}/fes2014/ocean_tide/ + {directory}/fes2014/load_tide/ + - {directory}/tpxo8_atlas/ + - {directory}/TPXO9_atlas_v5/ + epsg : int + Input coordinate system for 'x' and 'y' coordinates. + Defaults to 4326 (WGS84). + method : string + Method used to interpolate tidal contsituents + from model files. Options include: + - bilinear: quick bilinear interpolation + - spline: scipy bivariate spline interpolation + - linear, nearest: scipy regular grid interpolations + extrapolate : bool + Whether to extrapolate tides for locations outside of + the tide modelling domain using nearest-neighbor + cutoff : int or float + Extrapolation cutoff in kilometers. Set to `np.inf` + to extrapolate for all points. + + Returns + ------- + A pandas.DataFrame containing tide heights for every + combination of time and point coordinates. + """ + + import os + import pyproj + import numpy as np + import pyTMD.time + import pyTMD.model + import pyTMD.utilities + from pyTMD.calc_delta_time import calc_delta_time + from pyTMD.infer_minor_corrections import infer_minor_corrections + from pyTMD.predict_tide_drift import predict_tide_drift + from pyTMD.read_tide_model import extract_tidal_constants + from pyTMD.read_netcdf_model import extract_netcdf_constants + from pyTMD.read_GOT_model import extract_GOT_constants + from pyTMD.read_FES_model import extract_FES_constants + + # Check that tide directory is accessible + try: + os.access(directory, os.F_OK) + except: + raise FileNotFoundError("Invalid tide directory") + + # Get parameters for tide model + model = pyTMD.model(directory, format="netcdf", compressed=False).elevation(model) + + # If time passed as a single Timestamp, convert to datetime64 + if isinstance(time, pd.Timestamp): + time = time.to_datetime64() + + # Handle numeric or array inputs + x = np.atleast_1d(x) + y = np.atleast_1d(y) + time = np.atleast_1d(time) + + # Determine point and time counts + assert len(x) == len(y), "x and y must be the same length" + n_points = len(x) + n_times = len(time) + + # Converting x,y from EPSG to latitude/longitude + try: + # EPSG projection code string or int + crs1 = pyproj.CRS.from_string("epsg:{0:d}".format(int(epsg))) + except (ValueError, pyproj.exceptions.CRSError): + # Projection SRS string + crs1 = pyproj.CRS.from_string(epsg) + + crs2 = pyproj.CRS.from_string("epsg:{0:d}".format(4326)) + transformer = pyproj.Transformer.from_crs(crs1, crs2, always_xy=True) + lon, lat = transformer.transform(x.flatten(), y.flatten()) + + # Assert delta time is an array and convert datetime + time = np.atleast_1d(time) + t = pyTMD.time.convert_datetime(time, epoch=(1992, 1, 1, 0, 0, 0)) / 86400.0 + + # Delta time (TT - UT1) file + delta_file = pyTMD.utilities.get_data_path(["data", "merged_deltat.data"]) + + # Read tidal constants and interpolate to grid points + if model.format in ("OTIS", "ATLAS"): + amp, ph, D, c = extract_tidal_constants( + lon, + lat, + model.grid_file, + model.model_file, + model.projection, + TYPE=model.type, + METHOD=method, + EXTRAPOLATE=extrapolate, + CUTOFF=cutoff, + GRID=model.format, + ) + deltat = np.zeros_like(t) + + elif model.format == "netcdf": + amp, ph, D, c = extract_netcdf_constants( + lon, + lat, + model.grid_file, + model.model_file, + TYPE=model.type, + METHOD=method, + EXTRAPOLATE=extrapolate, + CUTOFF=cutoff, + SCALE=model.scale, + GZIP=model.compressed, + ) + deltat = np.zeros_like(t) + + elif model.format == "GOT": + amp, ph, c = extract_GOT_constants( + lon, + lat, + model.model_file, + METHOD=method, + EXTRAPOLATE=extrapolate, + CUTOFF=cutoff, + SCALE=model.scale, + GZIP=model.compressed, + ) + + # Interpolate delta times from calendar dates to tide time + deltat = calc_delta_time(delta_file, t) + + elif model.format == "FES": + amp, ph = extract_FES_constants( + lon, + lat, + model.model_file, + TYPE=model.type, + VERSION=model.version, + METHOD=method, + EXTRAPOLATE=extrapolate, + CUTOFF=cutoff, + SCALE=model.scale, + GZIP=model.compressed, + ) + + # Available model constituents + c = model.constituents + + # Interpolate delta times from calendar dates to tide time + deltat = calc_delta_time(delta_file, t) + + # Calculate complex phase in radians for Euler's + cph = -1j * ph * np.pi / 180.0 + + # Calculate constituent oscillation + hc = amp * np.exp(cph) + + # Repeat constituents to length of time and number of input + # coords before passing to `predict_tide_drift` + t, hc, deltat = ( + np.tile(t, n_points), + hc.repeat(n_times, axis=0), + np.tile(deltat, n_points), + ) + + # Predict tidal elevations at time and infer minor corrections + npts = len(t) + tide = np.ma.zeros((npts), fill_value=np.nan) + tide.mask = np.any(hc.mask, axis=1) + + # Depending on pyTMD version (<=1.06 vs > 1.06), use different params + # TODO: Remove once Sandbox is updated to use pyTMD version 1.0.9 + try: + tide.data[:] = predict_tide_drift( + t, hc, c, deltat=deltat, corrections=model.format + ) + minor = infer_minor_corrections( + t, hc, c, deltat=deltat, corrections=model.format + ) + except: + tide.data[:] = predict_tide_drift( + t, hc, c, DELTAT=deltat, CORRECTIONS=model.format + ) + minor = infer_minor_corrections( + t, hc, c, DELTAT=deltat, CORRECTIONS=model.format + ) + tide.data[:] += minor.data[:] + + # Replace invalid values with fill value + tide.data[tide.mask] = tide.fill_value + + # Export data as a dataframe + return pd.DataFrame( + { + "time": np.tile(time, n_points), + "x": np.repeat(x, n_times), + "y": np.repeat(y, n_times), + "tide_m": tide, + } + ).set_index("time") + + +def pixel_tides( + ds, + times=None, + resample=True, + calculate_quantiles=None, + resolution=None, + buffer=None, + resample_method="bilinear", + **model_tides_kwargs, +): + """ + Obtain tide heights for each pixel in a dataset by modelling + tides into a low-resolution grid surrounding the dataset, + then (optionally) spatially resample this low-res data back + into the original higher resolution dataset extent and resolution. + + Parameters: + ----------- + ds : xarray.Dataset + A dataset whose geobox (`ds.odc.geobox`) will be used to define + the spatial extent of the low resolution tide modelling grid. + times : pandas.DatetimeIndex or list of pandas.Timestamps, optional + By default, the function will model tides using the times + contained in the `time` dimension of `ds`. Alternatively, this + param can be used to model tides for a custom set of times + instead. For example: + `times=pd.date_range(start="2000", end="2001", freq="5h")` + resample : bool, optional + Whether to resample low resolution tides back into `ds`'s original + higher resolution grid. Set this to `False` if you do not want + low resolution tides to be re-projected back to higher resolution. + calculate_quantiles : list or np.array, optional + Rather than returning all individual tides, low-resolution tides + can be first aggregated using a quantile calculation by passing in + a list or array of quantiles to compute. For example, this could + be used to calculate the min/max tide across all times: + `calculate_quantiles=[0.0, 1.0]`. + resolution: int, optional + The desired resolution of the low-resolution grid used for tide + modelling. The default None will create a 5000 m resolution grid + if `ds` has a projected CRS (i.e. metre units), or a 0.05 degree + resolution grid if `ds` has a geographic CRS (e.g. degree units). + Note: higher resolutions do not necessarily provide better + tide modelling performance, as results will be limited by the + resolution of the underlying global tide model (e.g. 1/16th + degree / ~5 km resolution grid for FES2014). + buffer : int, optional + The amount by which to buffer the higher resolution grid extent + when creating the new low resolution grid. This buffering is + important as it ensures that ensure pixel-based tides are seamless + across dataset boundaries. This buffer will eventually be clipped + away when the low-resolution data is re-projected back to the + resolution and extent of the higher resolution dataset. To + ensure that at least two pixels occur outside of the dataset + bounds, the default None applies a 12000 m buffer if `ds` has a + projected CRS (i.e. metre units), or a 0.12 degree buffer if + `ds` has a geographic CRS (e.g. degree units). + resample_method : string, optional + If resampling is requested (see `resample` above), use this + resampling method when converting from low resolution to high + resolution pixels. Defaults to "bilinear"; valid options include + "nearest", "cubic", "min", "max", "average" etc. + **model_tides_kwargs : + Optional parameters passed to the `dea_tools.coastal.model_tides` + function. Important parameters include "model" and "directory", + used to specify the tide model to use and the location of its files. + + Returns: + -------- + If `resample` is True: + + tides_lowres : xr.DataArray + A low resolution data array giving either tide heights every + timestep in `ds` (if `times` is None), tide heights at every + time in `times` (if `times` is not None), or tide height quantiles + for every quantile provided by `calculate_quantiles`. + + If `resample` is False: + + tides_highres, tides_lowres : tuple of xr.DataArrays + In addition to `tides_lowres` (see above), a high resolution + array of tide heights will be generated that matches the + exact spatial resolution and extent of `ds`. This will contain + either tide heights every timestep in `ds` (if `times` is None), + tide heights at every time in `times` (if `times` is not None), + or tide height quantiles for every quantile provided by + `calculate_quantiles`. + """ + + import odc.geo.xr + from odc.geo.geobox import GeoBox + + # First test if no time dimension and nothing passed to `times` + if ('time' not in ds.dims) & (times is None): + raise ValueError( + "`ds` does not contain a 'time' dimension. Times are required " + "for modelling tides: please pass in a set of custom tides " + "using the `times` parameter. For example: " + "`times=pd.date_range(start='2000', end='2001', freq='5h')`" + ) + + # If custom times are provided, convert them to a consistent + # pandas.DatatimeIndex format + if times is not None: + if isinstance(times, list): + time_coords = pd.DatetimeIndex(times) + elif isinstance(times, pd.Timestamp): + time_coords = pd.DatetimeIndex([times]) + else: + time_coords = times + + # Otherwise, use times from `ds` directly + else: + time_coords = ds.coords["time"] + + # Determine spatial dimensions + y_dim, x_dim = ds.odc.spatial_dims + + # Determine resolution and buffer, using different defaults for + # geographic (i.e. degrees) and projected (i.e. metres) CRSs: + crs_units = ds.odc.geobox.crs.units[0][0:6] + if ds.odc.geobox.crs.geographic: + if resolution is None: + resolution = 0.05 + elif resolution > 360: + raise ValueError(f"A resolution of greater than 360 was " + f"provided, but `ds` has a geographic CRS " + f"in {crs_units} units. Did you accidently " + f"provide a resolution in projected " + f"(i.e. metre) units?") + if buffer is None: + buffer = 0.12 + else: + if resolution is None: + resolution = 5000 + elif resolution < 1: + raise ValueError(f"A resolution of less than 1 was provided, " + f"but `ds` has a projected CRS in " + f"{crs_units} units. Did you accidently " + f"provide a resolution in geographic " + f"(degree) units?") + if buffer is None: + buffer = 12000 + + # Raise error if resolution is less than dataset resolution + dataset_res = ds.odc.geobox.resolution.x + if resolution < dataset_res: + raise ValueError(f"The resolution of the low-resolution tide " + f"modelling grid ({resolution:.2f}) is less " + f"than `ds`'s pixel resolution ({dataset_res:.2f}). " + f"This can cause extremely slow tide modelling " + f"performance. Please select provide a resolution " + f"greater than {dataset_res:.2f} using " + f"`pixel_tides`'s 'resolution' parameter.") + + # Create a new reduced resolution tide modelling grid after + # first buffering the grid + print(f"Creating reduced resolution {resolution} x {resolution} " + f"{crs_units} tide modelling array") + buffered_geobox = ds.odc.geobox.buffered(buffer) + rescaled_geobox = GeoBox.from_bbox( + bbox=buffered_geobox.boundingbox, resolution=resolution + ) + rescaled_ds = odc.geo.xr.xr_zeros(rescaled_geobox) + + # Flatten grid to 1D, then add time dimension + flattened_ds = rescaled_ds.stack(z=(x_dim, y_dim)) + flattened_ds = flattened_ds.expand_dims(dim={"time": time_coords.values}) + + # Model tides for each timestep + model = ( + "FES2014" if "model" not in model_tides_kwargs else model_tides_kwargs["model"] + ) + print(f"Modelling tides using {model} tide model") + tide_df = model_tides( + x=flattened_ds[x_dim], + y=flattened_ds[y_dim], + time=flattened_ds.time, + epsg=ds.odc.geobox.crs.epsg, + **model_tides_kwargs, + ) + + # Rename x and y coordinates to match satellite array + tide_df = tide_df.rename({"x": x_dim, "y": y_dim}, axis=1) + + # Insert modelled tide values back into flattened array, then unstack + # back to 3D (y, x, time) + tides_lowres = ( + + # Convert dataframe to xarray format + tide_df.set_index([x_dim, y_dim], append=True) + .to_xarray() + + # Re-index and transpose back into 3D + .tide_m.reindex_like(rescaled_ds) + .transpose("time", y_dim, x_dim) + .astype(np.float32) + ) + + # Optionally calculate and return quantiles rather than raw data + if calculate_quantiles is not None: + + print("Computing tide quantiles") + tides_lowres = tides_lowres.quantile(q=calculate_quantiles, dim="time") + reproject_dim = "quantile" + + else: + reproject_dim = "time" + + # Ensure CRS is present + tides_lowres = tides_lowres.odc.assign_crs(ds.odc.geobox.crs) + + # Reproject each timestep into original high resolution grid + if resample: + + print("Reprojecting tides into original array") + tides_highres = parallel_apply( + tides_lowres, + reproject_dim, + odc.algo.xr_reproject, + ds.odc.geobox.compat, + resample_method, + ) + + return tides_highres, tides_lowres + + else: + print("Returning low resolution tide array") + return tides_lowres + +def tidal_tag( + ds, + ebb_flow=False, + swap_dims=False, + tidepost_lat=None, + tidepost_lon=None, + return_tideposts=False, + **model_tides_kwargs, +): + """ + Takes an xarray.Dataset and returns the same dataset with a new + `tide_m` variable giving the height of the tide at the exact + moment of each satellite acquisition. + + The function models tides at the centroid of the dataset by default, + but a custom tidal modelling location can be specified using + `tidepost_lat` and `tidepost_lon`. + + The default settings use the FES2014 global tidal model, implemented + using the pyTMD Python package. FES2014 was produced by NOVELTIS, + LEGOS, CLS Space Oceanography Division and CNES. It is distributed + by AVISO, with support from CNES (http://www.aviso.altimetry.fr/). + + Parameters + ---------- + ds : xarray.Dataset + An xarray.Dataset object with x, y and time dimensions + ebb_flow : bool, optional + An optional boolean indicating whether to compute if the + tide phase was ebbing (falling) or flowing (rising) for each + observation. The default is False; if set to True, a new + `ebb_flow` variable will be added to the dataset with each + observation labelled with 'Ebb' or 'Flow'. + swap_dims : bool, optional + An optional boolean indicating whether to swap the `time` + dimension in the original xarray.Dataset to the new + `tide_m` variable. Defaults to False. + tidepost_lat, tidepost_lon : float or int, optional + Optional coordinates used to model tides. The default is None, + which uses the centroid of the dataset as the tide modelling + location. + return_tideposts : bool, optional + An optional boolean indicating whether to return the `tidepost_lat` + and `tidepost_lon` location used to model tides in addition to the + xarray.Dataset. Defaults to False. + **model_tides_kwargs : + Optional parameters passed to the `dea_tools.coastal.model_tides` + function. Important parameters include "model" and "directory", + used to specify the tide model to use and the location of its files. + + Returns + ------- + The original xarray.Dataset with a new `tide_m` variable giving + the height of the tide (and optionally, its ebb-flow phase) at the + exact moment of each satellite acquisition (if `return_tideposts=True`, + the function will also return the `tidepost_lon` and `tidepost_lat` + location used in the analysis). + + """ + + import odc.geo.xr + + # If custom tide modelling locations are not provided, use the + # dataset centroid + if not tidepost_lat or not tidepost_lon: + + tidepost_lon, tidepost_lat = ds.odc.geobox.geographic_extent.centroid.coords[0] + print( + f"Setting tide modelling location from dataset centroid: " + f"{tidepost_lon:.2f}, {tidepost_lat:.2f}" + ) + + else: + print( + f"Using user-supplied tide modelling location: " + f"{tidepost_lon:.2f}, {tidepost_lat:.2f}" + ) + + # Use tidal model to compute tide heights for each observation: + model = ( + "FES2014" if "model" not in model_tides_kwargs else model_tides_kwargs["model"] + ) + print(f"Modelling tides using {model} tidal model") + tide_df = model_tides( + x=tidepost_lon, + y=tidepost_lat, + time=ds.time, + epsg="EPSG:4326", + **model_tides_kwargs, + ) + + # If tides cannot be successfully modeled (e.g. if the centre of the + # xarray dataset is located is over land), raise an exception + if tide_df.tide_m.isnull().all(): + + raise ValueError( + f"Tides could not be modelled for dataset centroid located " + f"at {tidepost_lon:.2f}, {tidepost_lat:.2f}. This can occur if " + f"this coordinate occurs over land. Please manually specify " + f"a tide modelling location located over water using the " + f"`tidepost_lat` and `tidepost_lon` parameters." + ) + + # Assign tide heights to the dataset as a new variable + ds["tide_m"] = xr.DataArray(tide_df.tide_m, coords=[ds.time]) + + # Optionally calculate the tide phase for each observation + if ebb_flow: + + # Model tides for a time 15 minutes prior to each previously + # modelled satellite acquisition time. This allows us to compare + # tide heights to see if they are rising or falling. + print("Modelling tidal phase (e.g. ebb or flow)") + tide_pre_df = model_tides( + x=tidepost_lon, + y=tidepost_lat, + time=(ds.time - pd.Timedelta("15 min")), + epsg="EPSG:4326", + **model_tides_kwargs, + ) + + # Compare tides computed for each timestep. If the previous tide + # was higher than the current tide, the tide is 'ebbing'. If the + # previous tide was lower, the tide is 'flowing' + tidal_phase = [ + "Ebb" if i else "Flow" + for i in tide_pre_df.tide_m.values > tide_df.tide_m.values + ] + + # Assign tide phase to the dataset as a new variable + ds["ebb_flow"] = xr.DataArray(tidal_phase, coords=[ds.time]) + + # If swap_dims = True, make tide height the primary dimension + # instead of time + if swap_dims: + + # Swap dimensions and sort by tide height + ds = ds.swap_dims({"time": "tide_m"}) + ds = ds.sortby("tide_m") + ds = ds.drop_vars("time") + + if return_tideposts: + return ds, tidepost_lon, tidepost_lat + else: + return ds + + +def tidal_stats( + ds, + tidepost_lat=None, + tidepost_lon=None, + plain_english=True, + plot=True, + modelled_freq="2h", + linear_reg=False, + round_stats=3, + **model_tides_kwargs, +): + """ + Takes an xarray.Dataset and statistically compares the tides + modelled for each satellite observation against the full modelled + tidal range. This comparison can be used to evaluate whether the + tides observed by satellites (e.g. Landsat) are biased compared to + the natural tidal range (e.g. fail to observe either the highest or + lowest tides etc). + + For more information about the tidal statistics computed by this + function, refer to Figure 8 in Bishop-Taylor et al. 2018: + https://www.sciencedirect.com/science/article/pii/S0272771418308783#fig8 + + The function models tides at the centroid of the dataset by default, + but a custom tidal modelling location can be specified using + `tidepost_lat` and `tidepost_lon`. + + The default settings use the FES2014 global tidal model, implemented + using the pyTMD Python package. FES2014 was produced by NOVELTIS, + LEGOS, CLS Space Oceanography Division and CNES. It is distributed + by AVISO, with support from CNES (http://www.aviso.altimetry.fr/). + + Parameters + ---------- + ds : xarray.Dataset + An xarray.Dataset object with x, y and time dimensions + tidepost_lat, tidepost_lon : float or int, optional + Optional coordinates used to model tides. The default is None, + which uses the centroid of the dataset as the tide modelling + location. + plain_english : bool, optional + An optional boolean indicating whether to print a plain english + version of the tidal statistics to the screen. Defaults to True. + plot : bool, optional + An optional boolean indicating whether to plot how satellite- + observed tide heights compare against the full tidal range. + Defaults to True. + modelled_freq : str, optional + An optional string giving the frequency at which to model tides + when computing the full modelled tidal range. Defaults to '2h', + which computes a tide height for every two hours across the + temporal extent of `ds`. + linear_reg: bool, optional + Experimental: whether to return linear regression stats that + assess whether dstellite-observed and all available tides show + any decreasing or increasing trends over time. Not currently + recommended as all observed regressions always return as + significant due to far larger sample size. + round_stats : int, optional + The number of decimal places used to round the output statistics. + Defaults to 3. + **model_tides_kwargs : + Optional parameters passed to the `dea_tools.coastal.model_tides` + function. Important parameters include "model" and "directory", + used to specify the tide model to use and the location of its files. + + Returns + ------- + A pandas.Series object containing the following statistics: + + tidepost_lat: latitude used for modelling tide heights + tidepost_lon: longitude used for modelling tide heights + observed_min_m: minimum tide height observed by the satellite + all_min_m: minimum tide height from all available tides + observed_max_m: maximum tide height observed by the satellite + all_max_m: maximum tide height from all available tides + observed_range_m: tidal range observed by the satellite + all_range_m: full astronomical tidal range based on all + available tides + spread_m: proportion of the full astronomical tidal range observed + by the satellite (see Bishop-Taylor et al. 2018) + low_tide_offset: proportion of the lowest tides never observed + by the satellite (see Bishop-Taylor et al. 2018) + high_tide_offset: proportion of the highest tides never observed + by the satellite (see Bishop-Taylor et al. 2018) + + If `linear_reg = True`, the output will also contain: + + observed_slope: slope of any relationship between observed tide + heights and time + all_slope: slope of any relationship between all available tide + heights and time + observed_pval: significance/p-value of any relationship between + observed tide heights and time + all_pval: significance/p-value of any relationship between + all available tide heights and time + + """ + + # Model tides for each observation in the supplied xarray object + ds_tides, tidepost_lon, tidepost_lat = tidal_tag( + ds, + tidepost_lat=tidepost_lat, + tidepost_lon=tidepost_lon, + return_tideposts=True, + **model_tides_kwargs, + ) + + # Drop spatial ref for nicer plotting + if "spatial_ref" in ds_tides: + ds_tides = ds_tides.drop_vars("spatial_ref") + + # Generate range of times covering entire period of satellite record + all_timerange = pd.date_range( + start=ds_tides.time.min().item(), + end=ds_tides.time.max().item(), + freq=modelled_freq, + ) + + # Model tides for each timestep + all_tides_df = model_tides( + x=tidepost_lon, + y=tidepost_lat, + time=all_timerange, + epsg="EPSG:4326", + **model_tides_kwargs, + ) + + # Get coarse statistics on all and observed tidal ranges + obs_mean = ds_tides.tide_m.mean().item() + all_mean = all_tides_df.tide_m.mean() + obs_min, obs_max = ds_tides.tide_m.quantile([0.0, 1.0]).values + all_min, all_max = all_tides_df.tide_m.quantile([0.0, 1.0]).values + + # Calculate tidal range + obs_range = obs_max - obs_min + all_range = all_max - all_min + + # Calculate Bishop-Taylor et al. 2018 tidal metrics + spread = obs_range / all_range + low_tide_offset = abs(all_min - obs_min) / all_range + high_tide_offset = abs(all_max - obs_max) / all_range + + # Extract x (time in decimal years) and y (distance) values + all_x = ( + all_tides_df.index.year + + ((all_tides_df.index.dayofyear - 1) / 365) + + ((all_tides_df.index.hour - 1) / 24) + ) + all_y = all_tides_df.tide_m.values.astype(np.float32) + time_period = all_x.max() - all_x.min() + + # Extract x (time in decimal years) and y (distance) values + obs_x = ( + ds_tides.time.dt.year + + ((ds_tides.time.dt.dayofyear - 1) / 365) + + ((ds_tides.time.dt.hour - 1) / 24) + ) + obs_y = ds_tides.tide_m.values.astype(np.float32) + + + # Compute linear regression + obs_linreg = stats.linregress(x=obs_x, y=obs_y) + all_linreg = stats.linregress(x=all_x, y=all_y) + + if plain_english: + + print( + f"\n{spread:.0%} of the {all_range:.2f} m modelled astronomical " + f"tidal range is observed at this location.\nThe lowest " + f"{low_tide_offset:.0%} and highest {high_tide_offset:.0%} " + f"of astronomical tides are never observed.\n" + ) + + if linear_reg: + + if obs_linreg.pvalue > 0.05: + print( + f"Observed tides show no significant trends " + f"over the ~{time_period:.0f} year period." + ) + else: + obs_slope_desc = "decrease" if obs_linreg.slope < 0 else "increase" + print( + f"Observed tides {obs_slope_desc} significantly " + f"(p={obs_linreg.pvalue:.3f}) over time by " + f"{obs_linreg.slope:.03f} m per year (i.e. a " + f"~{time_period * obs_linreg.slope:.2f} m " + f"{obs_slope_desc} over the ~{time_period:.0f} year period)." + ) + + if all_linreg.pvalue > 0.05: + print( + f"All tides show no significant trends " + f"over the ~{time_period:.0f} year period." + ) + else: + all_slope_desc = "decrease" if all_linreg.slope < 0 else "increase" + print( + f"All tides {all_slope_desc} significantly " + f"(p={all_linreg.pvalue:.3f}) over time by " + f"{all_linreg.slope:.03f} m per year (i.e. a " + f"~{time_period * all_linreg.slope:.2f} m " + f"{all_slope_desc} over the ~{time_period:.0f} year period)." + ) + + if plot: + + # Create plot and add all time and observed tide data + fig, ax = plt.subplots(figsize=(10, 5)) + all_tides_df.tide_m.plot(ax=ax, alpha=0.4) + ds_tides.tide_m.plot.line( + ax=ax, marker="o", linewidth=0.0, color="black", markersize=2 + ) + + # Add horizontal lines for spread/offsets + ax.axhline(obs_min, color="black", linestyle=":", linewidth=1) + ax.axhline(obs_max, color="black", linestyle=":", linewidth=1) + ax.axhline(all_min, color="black", linestyle=":", linewidth=1) + ax.axhline(all_max, color="black", linestyle=":", linewidth=1) + + # Add text annotations for spread/offsets + ax.annotate( + f" High tide\n offset ({high_tide_offset:.0%})", + xy=(all_timerange.max(), np.mean([all_max, obs_max])), + va="center", + ) + ax.annotate( + f" Spread\n ({spread:.0%})", + xy=(all_timerange.max(), np.mean([obs_min, obs_max])), + va="center", + ) + ax.annotate( + f" Low tide\n offset ({low_tide_offset:.0%})", + xy=(all_timerange.max(), np.mean([all_min, obs_min])), + ) + + # Remove top right axes and add labels + ax.spines["right"].set_visible(False) + ax.spines["top"].set_visible(False) + ax.set_ylabel("Tide height (m)") + ax.set_xlabel("") + ax.margins(x=0.015) + + # Export pandas.Series containing tidal stats + output_stats = { + "tidepost_lat": tidepost_lat, + "tidepost_lon": tidepost_lon, + "observed_mean_m": obs_mean, + "all_mean_m": all_mean, + "observed_min_m": obs_min, + "all_min_m": all_min, + "observed_max_m": obs_max, + "all_max_m": all_max, + "observed_range_m": obs_range, + "all_range_m": all_range, + "spread": spread, + "low_tide_offset": low_tide_offset, + "high_tide_offset": high_tide_offset, + } + + if linear_reg: + output_stats.update( + { + "observed_slope": obs_linreg.slope, + "all_slope": all_linreg.slope, + "observed_pval": obs_linreg.pvalue, + "all_pval": all_linreg.pvalue, + } + ) + + return pd.Series(output_stats).round(round_stats) + + +def transect_distances(transects_gdf, lines_gdf, mode='distance'): + """ + Take a set of transects (e.g. shore-normal beach survey lines), and + determine the distance along the transect to each object in a set of + lines (e.g. shorelines). Distances are measured in the CRS of the + input datasets. + + For coastal applications, transects should be drawn from land to + water (with the first point being on land so that it can be used + as a consistent location from which to measure distances. + + The distance calculation can be performed using two modes: + - 'distance': Distances are measured from the start of the + transect to where it intersects with each line. Any transect + that intersects a line more than once is ignored. This mode is + useful for measuring e.g. the distance to the shoreline over + time from a consistent starting location. + - 'width' Distances are measured between the first and last + intersection between a transect and each line. Any transect + that intersects a line only once is ignored. This is useful + for e.g. measuring the width of a narrow area of coastline over + time, e.g. the neck of a spit or tombolo. + + Parameters + ---------- + transects_gdf : geopandas.GeoDataFrame + A GeoDataFrame containing one or multiple vector profile lines. + The GeoDataFrame's index column will be used to name the rows in + the output distance table. + lines_gdf : geopandas.GeoDataFrame + A GeoDataFrame containing one or multiple vector line features + that intersect the profile lines supplied to `transects_gdf`. + The GeoDataFrame's index column will be used to name the columns + in the output distance table. + mode : string, optional + Whether to use 'distance' (for measuring distances from the + start of a profile) or 'width' mode (for measuring the width + between two profile intersections). See docstring above for more + info; defaults to 'distance'. + + Returns + ------- + distance_df : pandas.DataFrame + A DataFrame containing distance measurements for each profile + line (rows) and line feature (columns). + """ + + import warnings + from shapely.errors import ShapelyDeprecationWarning + from shapely.geometry import Point + + def _intersect_dist(transect_gdf, lines_gdf, mode=mode): + """ + Take an individual transect, and determine the distance along + the transect to each object in a set of lines (e.g. shorelines). + """ + + # Identify intersections between transects and lines + intersect_points = lines_gdf.apply( + lambda x: x.geometry.intersection(transect_gdf.geometry), axis=1) + + # In distance mode, identify transects with one intersection only, + # and use this as the end point and the start of the transect as the + # start point when measuring distances + if mode == 'distance': + start_point = Point(transect_gdf.geometry.coords[0]) + point_df = intersect_points.apply( + lambda x: pd.Series({'start': start_point, 'end': x}) + if x.type == 'Point' + else pd.Series({'start': None, 'end': None})) + + # In width mode, identify transects with multiple intersections, and + # use the first intersection as the start point and the second + # intersection for the end point when measuring distances + if mode == 'width': + point_df = intersect_points.apply( + lambda x: pd.Series({'start': x.geoms[0], 'end': x.geoms[-1]}) + if x.type == 'MultiPoint' + else pd.Series({'start': None, 'end': None})) + + # Calculate distances between valid start and end points + distance_df = point_df.apply( + lambda x: x.start.distance(x.end) if x.start else None, axis=1) + + return distance_df + + # Run code after ignoring Shapely pre-v2.0 warnings + with warnings.catch_warnings(): + warnings.filterwarnings("ignore", category=ShapelyDeprecationWarning) + + # Assert that both datasets use the same CRS + assert transects_gdf.crs == lines_gdf.crs, ('Please ensure both ' + 'input datasets use the same CRS.') + + # Run distance calculations + distance_df = transects_gdf.apply( + lambda x: _intersect_dist(x, lines_gdf), axis=1) + + return pd.DataFrame(distance_df) + + +def get_coastlines(bbox: tuple, + crs="EPSG:4326", + layer="shorelines", + drop_wms=True) -> gpd.GeoDataFrame: + """ + Get DE Africa Coastlines data for a provided bounding box using WFS. + + For a full description of the DE Africa Coastlines dataset, refer to the + official Digital Earth Africa product description: + + Parameters + ---------- + bbox : (xmin, ymin, xmax, ymax), or geopandas object + Bounding box expressed as a tuple. Alternatively, a bounding + box can be automatically extracted by suppling a + geopandas.GeoDataFrame or geopandas.GeoSeries. + crs : str, optional + Optional CRS for the bounding box. This is ignored if `bbox` + is provided as a geopandas object. + layer : str, optional + Which DE Africa Coastlines layer to load. Options include the annual + shoreline vectors ("shorelines") and the rates of change + statistics points ("statistics"). Defaults to "shorelines". + drop_wms : bool, optional + Whether to drop WMS-specific attribute columns from the data. + These columns are used for visualising the dataset on DE Africa Maps, + and are unlikely to be useful for scientific analysis. Defaults + to True. + + Returns + ------- + gpd.GeoDataFrame + A GeoDataFrame containing shoreline or point features and + associated metadata. + """ + + # If bbox is a geopandas object, convert to bbox. + try: + crs = str(bbox.crs) + bbox = bbox.total_bounds + except: + pass + + # Get the available layers in the coastlines:DEAfrica_Coastlines group. + describe_layer_url = "https://geoserver.digitalearth.africa/geoserver/wms?service=WMS&version=1.1.1&request=DescribeLayer&layers=coastlines:DEAfrica_Coastlines&outputFormat=application/json" + describe_layer_response = requests.get(describe_layer_url).json() + available_layers = [layer["layerName"] for layer in describe_layer_response['layerDescriptions']] + + # Get the layer name. + if layer == "shorelines": + layer_name = [i for i in available_layers if "shorelines" in i] + else: + layer_name = [i for i in available_layers if "rates_of_change" in i] + + # Query WFS. + wfs = WebFeatureService(url=WFS_ADDRESS, version="1.1.0") + response = wfs.getfeature(typename=layer_name, + bbox=tuple(bbox) + (crs,), + outputFormat="json") + + # Load data as a geopandas.GeoDataFrame. + coastlines_gdf = gpd.read_file(response) + + # Clip to extent of bounding box. + extent = gpd.GeoSeries(box(*bbox), crs=crs).to_crs(coastlines_gdf.crs) + coastlines_gdf = coastlines_gdf.clip(extent) + + # Optionally drop WMS-specific columns. + if drop_wms: + coastlines_gdf = coastlines_gdf.loc[:, ~coastlines_gdf.columns.str.contains("wms_")] + + return coastlines_gdf + diff --git a/deafrica_tools/dask.py b/deafrica_tools/dask.py new file mode 100644 index 0000000..c29ed54 --- /dev/null +++ b/deafrica_tools/dask.py @@ -0,0 +1,106 @@ +""" +Functions for simplifying the creation of a local dask cluster. +""" + +from importlib.util import find_spec +import os +import dask +from aiohttp import ClientConnectionError +from datacube.utils.dask import start_local_dask +from datacube.utils.rio import configure_s3_access + +_HAVE_PROXY = bool(find_spec('jupyter_server_proxy')) +_IS_AWS = ('AWS_ACCESS_KEY_ID' in os.environ or + 'AWS_DEFAULT_REGION' in os.environ) + + +def create_local_dask_cluster(spare_mem='3Gb', display_client=True, return_client=False): + """ + Using the datacube utils function `start_local_dask`, generate + a local dask cluster. Automatically detects if on AWS or NCI. + + Parameters + ---------- + spare_mem : String, optional + The amount of memory, in Gb, to leave for the notebook to run. + This memory will not be used by the cluster. e.g '3Gb' + display_client : Bool, optional + An optional boolean indicating whether to display a summary of + the dask client, including a link to monitor progress of the + analysis. Set to False to hide this display. + return_client : Bool, optional + An optional boolean indicating whether to return the dask client + object. + + """ + + if _HAVE_PROXY: + # Configure dashboard link to go over proxy + prefix = os.environ.get('JUPYTERHUB_SERVICE_PREFIX', '/') + dask.config.set({"distributed.dashboard.link": + prefix + "proxy/{port}/status"}) + + # Start up a local cluster + client = start_local_dask(mem_safety_margin=spare_mem) + + if _IS_AWS: + # Configure GDAL for s3 access + configure_s3_access(aws_unsigned=True, + client=client) + + # Show the dask cluster settings + if display_client: + from IPython.display import display + display(client) + + # return the client as an object + if return_client: + return client + + +try: + from dask_gateway import Gateway + + def create_dask_gateway_cluster(profile='r5_L', workers=2): + """ + Create a cluster in our internal dask cluster. + + Parameters + ---------- + profile : str + Possible values are: + - r5_L (2 cores, 15GB memory) + - r5_XL (4 cores, 31GB memory) + - r5_2XL (8 cores, 63GB memory) + - r5_4XL (16 cores, 127GB memory) + + workers : int + Number of workers in the cluster. + """ + try: + gateway = Gateway() + + # Close any existing clusters + cluster_names = gateway.list_clusters() + if len(cluster_names) > 0: + print("Cluster(s) still running:", cluster_names) + for n in cluster_names: + cluster = gateway.connect(n.name) + cluster.shutdown() + + options = gateway.cluster_options() + options['profile'] = profile + + # limit username to alphanumeric characters + # kubernetes pods won't launch if labels contain anything other than [a-Z, -, _] + options['jupyterhub_user'] = ''.join(c if c.isalnum() else '-' for c in os.getenv('JUPYTERHUB_USER')) + + cluster = gateway.new_cluster(options) + cluster.scale(workers) + return cluster + except ClientConnectionError: + raise ConnectionError("access to dask gateway cluster unauthorized") + +except ImportError: + def create_dask_gateway_cluster(*args, **kwargs): + raise NotImplementedError \ No newline at end of file diff --git a/deafrica_tools/datahandling.py b/deafrica_tools/datahandling.py new file mode 100644 index 0000000..085952a --- /dev/null +++ b/deafrica_tools/datahandling.py @@ -0,0 +1,1523 @@ +""" +Functions for loading and handling Digital Earth Africa data. +""" + +# Import required packages +import os +from osgeo import gdal +import requests +import zipfile +import warnings +import numpy as np +import xarray as xr +import pandas as pd +import datetime +import pytz + +from collections import Counter +from datacube.utils import masking +from scipy.ndimage import binary_dilation +from odc.algo import mask_cleanup +from copy import deepcopy +import odc.algo + +from skimage.morphology import binary_erosion,binary_dilation,disk +from scipy.ndimage.filters import uniform_filter +from scipy.ndimage.measurements import variance +from datetime import datetime +from dateutil import parser +from deafrica_tools.bandindices import calculate_indices + +def _dc_query_only(**kw): + """ + Remove load-only parameters, the rest + can be passed to Query + + Returns + ======= + + dict of query parameters + """ + + def _impl( + measurements=None, + output_crs=None, + resolution=None, + resampling=None, + skip_broken_datasets=None, + dask_chunks=None, + fuse_func=None, + align=None, + datasets=None, + progress_cbk=None, + group_by=None, + **query, + ): + return query + + return _impl(**kw) + + +def _common_bands(dc, products): + """ + Takes a list of products and returns a list of measurements/bands + that are present in all products + Returns + ------- + List of band names + """ + common = None + bands = None + + for p in products: + p = dc.index.products.get_by_name(p) + if common is None: + common = set(p.measurements) + bands = list(p.measurements) + else: + common = common.intersection(set(p.measurements)) + return [band for band in bands if band in common] + + +def load_ard( + dc, + products=None, + min_gooddata=0.0, + categories_to_mask_ls=dict( + cloud="high_confidence", cloud_shadow="high_confidence" + ), + categories_to_mask_s2=[ + "cloud high probability", + "cloud medium probability", + "thin cirrus", + "cloud shadows", + "saturated or defective", + ], + categories_to_mask_s1=["invalid data"], + mask_filters=None, + mask_pixel_quality=True, + ls7_slc_off=True, + predicate=None, + dtype="auto", + verbose=True, + **kwargs, +): + """ + Loads analysis ready data. + + Loads and combines Landsat USGS Collections 2, Sentinel-2, and Sentinel-1 for + multiple sensors (i.e. ls5t, ls7e, ls8c and ls9 for Landsat; s2a and s2b for Sentinel-2), + optionally applies pixel quality masks, and drops time steps that + contain greater than a minimum proportion of good quality (e.g. non- + cloudy or shadowed) pixels. + + The function supports loading the following DE Africa products: + + Landsat: + * ls5_sr ('sr' denotes surface reflectance) + * ls7_sr + * ls8_sr + * ls9_sr + * ls5_st ('st' denotes surface temperature) + * ls7_st + * ls8_st + * ls9_st + + Sentinel-2: + * s2_l2a + + Sentinel-1: + * s1_rtc + + Last modified: Feb 2021 + + Parameters + ---------- + dc : datacube Datacube object + The Datacube to connect to, i.e. `dc = datacube.Datacube()`. + This allows you to also use development datacubes if required. + products : list + A list of product names to load data from. For example: + + * Landsat C2: ``['ls5_sr', 'ls7_sr', 'ls8_sr', 'ls9_sr']`` + * Sentinel-2: ``['s2_l2a']`` + * Sentinel-1: ``['s1_rtc']`` + + min_gooddata : float, optional + An optional float giving the minimum percentage of good quality + pixels required for a satellite observation to be loaded. + Defaults to 0.0 which will return all observations regardless of + pixel quality (set to e.g. 0.99 to return only observations with + more than 99% good quality pixels). + categories_to_mask_ls : dict, optional + An optional dictionary that is used to identify poor quality pixels + for masking. This mask is used for both masking out low + quality pixels (e.g. cloud or shadow), and for dropping + observations entirely based on the `min_gooddata` calculation. + categories_to_mask_s2 : list, optional + An optional list of Sentinel-2 Scene Classification Layer (SCL) names + that identify poor quality pixels for masking. + categories_to_mask_s1 : list, optional + An optional list of Sentinel-1 mask names that identify poor + quality pixels for masking. + mask_filters : iterable of tuples, optional + Iterable tuples of morphological operations - ("", ) + to apply on mask, where: + + operation: string, can be one of these morphological operations: + * ``'closing'`` = remove small holes in cloud - morphological closing + * ``'opening'`` = shrinks away small areas of the mask + * ``'dilation'`` = adds padding to the mask + * ``'erosion'`` = shrinks bright regions and enlarges dark regions + + radius: int + e.g. ``mask_filters=[('erosion', 5),("opening", 2),("dilation", 2)]`` + mask_pixel_quality : bool, optional + An optional boolean indicating whether to apply the poor data + mask to all observations that were not filtered out for having + less good quality pixels than ``min_gooddata``. E.g. if + ``min_gooddata=0.99``, the filtered observations may still contain + up to 1% poor quality pixels. The default of ``False`` simply + returns the resulting observations without masking out these + pixels; ``True`` masks them and sets them to NaN using the poor data + mask. This will convert numeric values to floating point values + which can cause memory issues, set to False to prevent this. + ls7_slc_off : bool, optional + An optional boolean indicating whether to include data from + after the Landsat 7 SLC failure (i.e. SLC-off). Defaults to + ``True``, which keeps all Landsat 7 observations > May 31 2003. + predicate : function, optional + An optional function that can be passed in to restrict the + datasets that are loaded by the function. A filter function + should take a `datacube.model.Dataset` object as an input (i.e. + as returned from `dc.find_datasets`), and return a boolean. + For example, a filter function could be used to return True on + only datasets acquired in January: + ``dataset.time.begin.month == 1`` + dtype : string, optional + An optional parameter that controls the data type/dtype that + layers are coerced to after loading. Valid values: ''`native`'', + ``'auto'``, ``'float{16|32|64}'``. + When ``'auto'`` is used, the data will be + converted to ``'float32'`` if masking is used, otherwise data will + be returned in the native data type of the data. Be aware that + if data is loaded in its native dtype, nodata and masked + pixels will be returned with the data's native nodata value + (typically ``-999``), not ``NaN``. + NOTE: If loading Landsat, the data is automatically rescaled so + 'native' dtype will return a value error. + verbose : bool, optional + If True, print progress statements during loading + **kwargs : dict, optional + A set of keyword arguments to ``dc.load`` that define the + spatiotemporal query used to extract data. This typically + includes ``measurements``, ``x`, ``y``, ``time``, ``resolution``, + ``resampling``, ``group_by`` and ``crs``. Keyword arguments can + either be listed directly in the ``load_ard`` call like any + other parameter (e.g. ``measurements=['red']``), or by + passing in a query kwarg dictionary (e.g. ``**query``). For a + list of possible options, see the ``dc.load`` documentation: + https://datacube-core.readthedocs.io/en/latest/dev/api/generate/datacube.Datacube.load.html + + Returns + ------- + combined_ds : xarray Dataset + An xarray dataset containing only satellite observations that + contains greater than `min_gooddata` proportion of good quality + pixels. + + """ + + ######### + # Setup # + ######### + # prevent function altering original query object + kwargs = deepcopy(kwargs) + + # We deal with `dask_chunks` separately + dask_chunks = kwargs.pop("dask_chunks", None) + requested_measurements = kwargs.pop("measurements", None) + + # Warn user if they combine lazy load with min_gooddata + if verbose: + if (min_gooddata > 0.0) and dask_chunks is not None: + warnings.warn( + "Setting 'min_gooddata' percentage to > 0.0 " + "will cause dask arrays to compute when " + "loading pixel-quality data to calculate " + "'good pixel' percentage. This can " + "slow the return of your dataset." + ) + + # Verify that products were provided and determine if Sentinel-2 + # or Landsat data is being loaded + if not products: + raise ValueError( + "Please provide a list of product names to load data from. " + "Valid options are: Landsat C2 SR: ['ls5_sr', 'ls7_sr', 'ls8_sr', 'ls9_sr'], or " + "Landsat C2 ST: ['ls5_st', 'ls7_st', 'ls8_st', 'ls9_st'], or " + "Sentinel-2: ['s2_l2a'], or" + "Sentinel-1: ['s1_rtc'], or" + ) + + # convert products to list if user passed as a string + if type(products) == str: + products=[products] + + if all(["ls" in product for product in products]): + product_type = "ls" + elif all(["s2" in product for product in products]): + product_type = "s2" + elif all(["s1" in product for product in products]): + product_type = "s1" + + # check if the landsat product is surface temperature + st = False + if (product_type == "ls") & (all(["st" in product for product in products])): + st = True + + # Check some parameters before proceeding + if (product_type == "ls") & (dtype == "native"): + raise ValueError( + "Cannot load Landsat bands in native dtype " + "as values require rescaling which converts dtype to float" + ) + + if product_type == "ls": + if any(k in categories_to_mask_ls for k in ("cirrus", "cirrus_confidence")): + raise ValueError( + "'cirrus' categories for the pixel quality mask" + " are not supported by load_ard" + ) + + # If `measurements` are specified but do not include pixel quality bands, + # add these to `measurements` according to collection + if product_type == "ls": + if verbose: + print("Using pixel quality parameters for USGS Collection 2") + fmask_band = "pixel_quality" + + elif product_type == "s2": + if verbose: + print("Using pixel quality parameters for Sentinel 2") + fmask_band = "SCL" + + elif product_type == "s1": + if verbose: + print("Using pixel quality parameters for Sentinel 1") + fmask_band = "mask" + + measurements = requested_measurements.copy() if requested_measurements else None + + # define a list of acceptable aliases to load landsat. We can't rely on 'common' + # measurements as native band names have the same name for different measurements. + ls_aliases = ["pixel_quality", "radiometric_saturation"] + if st: + ls_aliases = [ + "surface_temperature", + "surface_temperature_quality", + "atmospheric_transmittance", + "thermal_radiance", + "emissivity", + "emissivity_stddev", + "cloud_distance", + "upwell_radiance", + "downwell_radiance", + ] + ls_aliases + else: + ls_aliases = ["red", "green", "blue", "nir", "swir_1", "swir_2"] + ls_aliases + + if measurements is not None: + if product_type == "ls": + + # check we aren't loading aerosol bands from LS8 + aerosol_bands = [ + "aerosol_qa", + "qa_aerosol", + "atmos_opacity", + "coastal_aerosol", + "SR_QA_AEROSOL", + ] + if any(b in aerosol_bands for b in measurements): + raise ValueError( + "load_ard doesn't support loading aerosol or " + "atmospeheric opacity related bands " + "for Landsat, instead use dc.load()" + ) + + # check measurements are in acceptable aliases list for landsat + if set(measurements).issubset(ls_aliases): + pass + else: + raise ValueError( + "load_ard does not support all band aliases for Landsat, " + "use only the following band names to load Landsat data: " + + str(ls_aliases) + ) + + # Deal with "load all" case: pick a set of bands common across + # all products + if measurements is None: + if product_type == "ls": + measurements = ls_aliases + else: + measurements = _common_bands(dc, products) + + # If `measurements` are specified but do not include pq, add. + if measurements: + if fmask_band not in measurements: + measurements.append(fmask_band) + + # Get list of data and mask bands so that we can later exclude + # mask bands from being masked themselves (also handle the case of rad_sat) + data_bands = [ + band + for band in measurements + if band not in (fmask_band, "radiometric_saturation") + ] + mask_bands = [band for band in measurements if band not in data_bands] + + ################# + # Find datasets # + ################# + + # Pull out query params only to pass to dc.find_datasets + query = _dc_query_only(**kwargs) + + # Extract datasets for each product using subset of dcload_kwargs + dataset_list = [] + + # Get list of datasets for each product + if verbose: + print("Finding datasets") + for product in products: + + # Obtain list of datasets for product + if verbose: + print(f" {product}") + + if product_type == "ls": + # handle LS seperately to S2/S1 due to collection_category + # force the user to load Tier 1 + datasets = dc.find_datasets( + product=product, collection_category='T1', **query + ) + else: + datasets = dc.find_datasets(product=product, **query) + + # Remove Landsat 7 SLC-off observations if ls7_slc_off=False + if not ls7_slc_off and product in ["ls7_sr"]: + if verbose: + print(" Ignoring SLC-off observations for ls7") + datasets = [ + i + for i in datasets + if i.time.begin < datetime.datetime(2003, 5, 31, tzinfo=pytz.UTC) + ] + + # Add any returned datasets to list + dataset_list.extend(datasets) + + # Raise exception if no datasets are returned + if len(dataset_list) == 0: + raise ValueError( + "No data available for query: ensure that " + "the products specified have data for the " + "time and location requested" + ) + + # If predicate is specified, use this function to filter the list + # of datasets prior to load (this now redundant as dc.load now supports + # a predicate filter) + if predicate: + if verbose: + print(f"Filtering datasets using filter function") + dataset_list = [ds for ds in dataset_list if predicate(ds)] + + # Raise exception if filtering removes all datasets + if len(dataset_list) == 0: + raise ValueError("No data available after filtering with " "filter function") + + ############# + # Load data # + ############# + + # Note we always load using dask here so that + # we can lazy load data before filtering by good data + ds = dc.load( + datasets=dataset_list, + measurements=measurements, + dask_chunks={} if dask_chunks is None else dask_chunks, + **kwargs, + ) + #print(ds) + #################### + # Filter good data # + #################### + + # need to distinguish between products due to different + # pq band properties + + # collection 2 USGS + if product_type == "ls": + mask, _ = masking.create_mask_value( + ds[fmask_band].attrs["flags_definition"], **categories_to_mask_ls + ) + + pq_mask = (ds[fmask_band] & mask) != 0 + + # only run if data bands are present + if len(data_bands) > 0: + + # identify pixels that will become negative after rescaling (but not 0 values) + invalid = ( + ((ds[data_bands] < (-1.0 * -0.2 / 0.0000275)) & (ds[data_bands] > 0)) + .to_array(dim="band") + .any(dim="band") + ) + + #merge masks + pq_mask = np.logical_or(pq_mask, pq_mask) + + # sentinel 2 + if product_type == "s2": + pq_mask = odc.algo.enum_to_bool(mask=ds[fmask_band], + categories=categories_to_mask_s2) + + # sentinel 1 + if product_type == "s1": + pq_mask = odc.algo.enum_to_bool(mask=ds[fmask_band], + categories=categories_to_mask_s1) + #print(pq_mask) + # The good data percentage calculation has to load in all `fmask` + # data, which can be slow. If the user has chosen no filtering + # by using the default `min_gooddata = 0`, we can skip this step + # completely to save processing time + if min_gooddata > 0.0: + + # Compute good data for each observation as % of total pixels. + # Inveerting the pq_mask for this because cloud=True in pq_mask + # and we want to sum good pixels + if verbose: + print("Counting good quality pixels for each time step") + data_perc = (~pq_mask).sum(axis=[1, 2], dtype="int32") / ( + pq_mask.shape[1] * pq_mask.shape[2] + ) + + keep = (data_perc >= min_gooddata).persist() + + # Filter by `min_gooddata` to drop low quality observations + total_obs = len(ds.time) + ds = ds.sel(time=keep) + pq_mask = pq_mask.sel(time=keep) + + if verbose: + print( + f"Filtering to {len(ds.time)} out of {total_obs} " + f"time steps with at least {min_gooddata:.1%} " + f"good quality pixels" + ) + + # morpholigcal filtering on cloud masks + if (mask_filters is not None) & (mask_pixel_quality): + if verbose: + print(f"Applying morphological filters to pq mask {mask_filters}") + pq_mask = mask_cleanup(pq_mask, mask_filters=mask_filters) + + ############### + # Apply masks # + ############### + + # Generate good quality data mask + mask = None + if mask_pixel_quality: + if verbose: + print("Applying pixel quality/cloud mask") + mask = pq_mask + + # Split into data/masks bands, as conversion to float and masking + # should only be applied to data bands + ds_data = ds[data_bands] + ds_masks = ds[mask_bands] + + # Remove sentinel-2 pixels valued 1 (scene edges, terrain shadow) + if product_type == "s2": + valid_data_mask = (ds_data > 1).to_array(dim="band").all(dim="band") + ds_data = odc.algo.keep_good_only(ds_data, where=valid_data_mask) + + # Mask data if either of the above masks were generated + if mask is not None: + ds_data = odc.algo.erase_bad(ds_data, where=mask) + + # Automatically set dtype to either native or float32 depending + # on whether masking was requested + if dtype == "auto": + dtype = "native" if mask is None else "float32" + + # Set nodata values using odc.algo tools to reduce peak memory + # use when converting data dtype + if dtype != "native": + ds_data = odc.algo.to_float(ds_data, dtype=dtype) + + # Put data and mask bands back together + attrs = ds.attrs + ds = xr.merge([ds_data, ds_masks]) + ds.attrs.update(attrs) + + ############### + # Return data # + ############### + + # Drop bands not originally requested by user + if requested_measurements: + ds = ds[requested_measurements] + + # Apply the scale and offset factors to Collection 2 Landsat. We need + # different factors for different bands. Also handle the case where + # masking_pixel_quaity = False, in which case the dtype is still + # in int, so we convert it to float + if product_type == "ls": + if verbose: + print("Re-scaling Landsat C2 data") + + sr_bands = ["red", "green", "blue", "nir", "swir_1", "swir_2"] + radiance_bands = ["thermal_radiance", "upwell_radiance", "downwell_radiance"] + trans_emiss = ["atmospheric_transmittance", "emissivity", "emissivity_stddev"] + qa = ["pixel_quality", "radiometric_saturation"] + + if mask_pixel_quality == False: + # set nodata to NaNs before rescaling + # in the case where masking hasn't already done this + for band in ds.data_vars: + if band not in qa: + ds[band] = odc.algo.to_f32(ds[band]) + + for band in ds.data_vars: + if band == "cloud_distance": + ds[band] = 0.01 * ds[band] + + if band == "surface_temperature_quality": + ds[band] = 0.01 * ds[band] + + if band in radiance_bands: + ds[band] = 0.001 * ds[band] + + if band in trans_emiss: + ds[band] = 0.0001 * ds[band] + + if band in sr_bands: + ds[band] = 2.75e-5 * ds[band] - 0.2 + + if band == "surface_temperature": + ds[band] = ds[band] * 0.00341802 + 149.0 + + # add back attrs that are lost during scaling calcs + for band in ds.data_vars: + ds[band].attrs.update(attrs) + + # If user supplied dask_chunks, return data as a dask array without + # actually loading it in + if dask_chunks is not None: + if verbose: + print(f"Returning {len(ds.time)} time steps as a dask array") + return ds + else: + if verbose: + print(f"Loading {len(ds.time)} time steps") + return ds.compute() + + +def array_to_geotiff( + fname, data, geo_transform, projection, nodata_val=0, dtype=gdal.GDT_Float32 +): + """ + Create a single band GeoTIFF file with data from an array. + + Because this works with simple arrays rather than xarray datasets + from DEA, it requires geotransform info (`(upleft_x, x_size, + x_rotation, upleft_y, y_rotation, y_size)`) and projection data + (in "WKT" format) for the output raster. These are typically + obtained from an existing raster using the following GDAL calls: + + >>> from osgeo import gdal + >>> gdal_dataset = gdal.Open(raster_path) + >>> geotrans = gdal_dataset.GetGeoTransform() + >>> prj = gdal_dataset.GetProjection() + + or alternatively, directly from an xarray dataset: + + >>> geotrans = xarraydataset.geobox.transform.to_gdal() + >>> prj = xarraydataset.geobox.crs.wkt + + + Parameters + ---------- + fname : str + Output geotiff file path including extension + data : numpy array + Input array to export as a geotiff + geo_transform : tuple + Geotransform for output raster; e.g. `(upleft_x, x_size, + x_rotation, upleft_y, y_rotation, y_size)` + projection : str + Projection for output raster (in "WKT" format) + nodata_val : int, optional + Value to convert to nodata in the output raster; default 0 + dtype : gdal dtype object, optional + Optionally set the dtype of the output raster; can be + useful when exporting an array of float or integer values. + Defaults to `gdal.GDT_Float32` + + """ + + # Set up driver + driver = gdal.GetDriverByName("GTiff") + + # Create raster of given size and projection + rows, cols = data.shape + dataset = driver.Create(fname, cols, rows, 1, dtype) + dataset.SetGeoTransform(geo_transform) + dataset.SetProjection(projection) + + # Write data to array and set nodata values + band = dataset.GetRasterBand(1) + band.WriteArray(data) + band.SetNoDataValue(nodata_val) + + # Close file + dataset = None + + +def mostcommon_crs(dc, product, query): + """ + Takes a given query and returns the most common CRS for observations + returned for that spatial extent. This can be useful when your study + area lies on the boundary of two UTM zones, forcing you to decide + which CRS to use for your `output_crs` in `dc.load`. + + Parameters + ---------- + dc : datacube Datacube object + The Datacube to connect to, i.e. `dc = datacube.Datacube()`. + This allows you to also use development datacubes if required. + product : str + A product name to load CRSs from + query : dict + A datacube query including x, y and time range to assess for the + most common CRS + + Returns + ------- + str + A EPSG string giving the most common CRS from all datasets returned + by the query above + + """ + + # remove dask_chunks & align to prevent func failing + # prevent function altering dictionary kwargs + query = deepcopy(query) + if "dask_chunks" in query: + query.pop("dask_chunks", None) + + if "align" in query: + query.pop("align", None) + + # List of matching products + matching_datasets = dc.find_datasets(product=product, **query) + + # Extract all CRSs + crs_list = [str(i.crs) for i in matching_datasets] + + # Identify most common CRS + crs_counts = Counter(crs_list) + crs_mostcommon = crs_counts.most_common(1)[0][0] + + # Warn user if multiple CRSs are encountered + if len(crs_counts.keys()) > 1: + + warnings.warn( + f"Multiple UTM zones {list(crs_counts.keys())} " + f"were returned for this query. Defaulting to " + f"the most common zone: {crs_mostcommon}", + UserWarning, + ) + + return crs_mostcommon + + +def download_unzip(url, output_dir=None, remove_zip=True): + """ + Downloads and unzips a .zip file from an external URL to a local + directory. + + Parameters + ---------- + url : str + A string giving a URL path to the zip file you wish to download + and unzip + output_dir : str, optional + An optional string giving the directory to unzip files into. + Defaults to None, which will unzip files in the current working + directory + remove_zip : bool, optional + An optional boolean indicating whether to remove the downloaded + .zip file after files are unzipped. Defaults to True, which will + delete the .zip file. + + """ + + # Get basename for zip file + zip_name = os.path.basename(url) + + # Raise exception if the file is not of type .zip + if not zip_name.endswith(".zip"): + raise ValueError( + f"The URL provided does not point to a .zip " + f"file (e.g. {zip_name}). Please specify a " + f"URL path to a valid .zip file" + ) + + # Download zip file + print(f"Downloading {zip_name}") + r = requests.get(url) + with open(zip_name, "wb") as f: + f.write(r.content) + + # Extract into output_dir + with zipfile.ZipFile(zip_name, "r") as zip_ref: + zip_ref.extractall(output_dir) + print( + f"Unzipping output files to: " + f"{output_dir if output_dir else os.getcwd()}" + ) + + # Optionally cleanup + if remove_zip: + os.remove(zip_name) + + +def wofs_fuser(dest, src): + """ + Fuse two WOfS water measurements represented as `ndarray` objects. + + Note: this is a copy of the function located here: + https://github.com/GeoscienceAustralia/digitalearthau/blob/develop/digitalearthau/utils.py + """ + empty = (dest & 1).astype(bool) + both = ~empty & ~((src & 1).astype(bool)) + dest[empty] = src[empty] + dest[both] |= src[both] + + +def dilate(array, dilation=10, invert=True): + """ + Dilate a binary array by a specified nummber of pixels using a + disk-like radial dilation. + + By default, invalid (e.g. False or 0) values are dilated. This is + suitable for applications such as cloud masking (e.g. creating a + buffer around cloudy or shadowed pixels). This functionality can + be reversed by specifying `invert=False`. + + Parameters + ---------- + array : array + The binary array to dilate. + dilation : int, optional + An optional integer specifying the number of pixels to dilate + by. Defaults to 10, which will dilate `array` by 10 pixels. + invert : bool, optional + An optional boolean specifying whether to invert the binary + array prior to dilation. The default is True, which dilates the + invalid values in the array (e.g. False or 0 values). + + Returns + ------- + array + An array of the same shape as `array`, with valid data pixels + dilated by the number of pixels specified by `dilation`. + """ + + y, x = np.ogrid[ + -dilation : (dilation + 1), + -dilation : (dilation + 1), + ] + + # disk-like radial dilation + kernel = (x * x) + (y * y) <= (dilation + 0.5) ** 2 + + # If invert=True, invert True values to False etc + if invert: + array = ~array + + return ~binary_dilation( + array.astype(bool), structure=kernel.reshape((1,) + kernel.shape) + ) + + +def _select_along_axis(values, idx, axis): + other_ind = np.ix_(*[np.arange(s) for s in idx.shape]) + sl = other_ind[:axis] + (idx,) + other_ind[axis:] + return values[sl] + + +def first(array: xr.DataArray, dim: str, index_name: str = None) -> xr.DataArray: + """ + Finds the first occuring non-null value along the given dimension. + + Parameters + ---------- + array : xr.DataArray + The array to search. + dim : str + The name of the dimension to reduce by finding the first non-null value. + + Returns + ------- + reduced : xr.DataArray + An array of the first non-null values. + The `dim` dimension will be removed, and replaced with a coord of the + same name, containing the value of that dimension where the last value + was found. + """ + axis = array.get_axis_num(dim) + idx_first = np.argmax(~pd.isnull(array), axis=axis) + reduced = array.reduce(_select_along_axis, idx=idx_first, axis=axis) + reduced[dim] = array[dim].isel({dim: xr.DataArray(idx_first, dims=reduced.dims)}) + if index_name is not None: + reduced[index_name] = xr.DataArray(idx_first, dims=reduced.dims) + return reduced + + +def last(array: xr.DataArray, dim: str, index_name: str = None) -> xr.DataArray: + """ + Finds the last occuring non-null value along the given dimension. + + Parameters + ---------- + array : xr.DataArray + The array to search. + dim : str + The name of the dimension to reduce by finding the last non-null value. + index_name : str, optional + If given, the name of a coordinate to be added containing the index + of where on the dimension the nearest value was found. + + Returns + ------- + reduced : xr.DataArray + An array of the last non-null values. + The `dim` dimension will be removed, and replaced with a coord of the + same name, containing the value of that dimension where the last value + was found. + """ + axis = array.get_axis_num(dim) + rev = (slice(None),) * axis + (slice(None, None, -1),) + idx_last = -1 - np.argmax(~pd.isnull(array)[rev], axis=axis) + reduced = array.reduce(_select_along_axis, idx=idx_last, axis=axis) + reduced[dim] = array[dim].isel({dim: xr.DataArray(idx_last, dims=reduced.dims)}) + if index_name is not None: + reduced[index_name] = xr.DataArray(idx_last, dims=reduced.dims) + return reduced + + +def nearest( + array: xr.DataArray, dim: str, target, index_name: str = None +) -> xr.DataArray: + """ + Finds the nearest values to a target label along the given dimension, for + all other dimensions. + + E.g. For a DataArray with dimensions ('time', 'x', 'y') + + nearest_array = nearest(array, 'time', '2017-03-12') + + will return an array with the dimensions ('x', 'y'), with non-null values + found closest for each (x, y) pixel to that location along the time + dimension. + + The returned array will include the 'time' coordinate for each x,y pixel + that the nearest value was found. + + Parameters + ---------- + array : xr.DataArray + The array to search. + dim : str + The name of the dimension to look for the target label. + target : same type as array[dim] + The value to look up along the given dimension. + index_name : str, optional + If given, the name of a coordinate to be added containing the index + of where on the dimension the nearest value was found. + + Returns + ------- + nearest_array : xr.DataArray + An array of the nearest non-null values to the target label. + The `dim` dimension will be removed, and replaced with a coord of the + same name, containing the value of that dimension closest to the + given target label. + """ + before_target = slice(None, target) + after_target = slice(target, None) + + da_before = array.sel({dim: before_target}) + da_after = array.sel({dim: after_target}) + + da_before = last(da_before, dim, index_name) if da_before[dim].shape[0] else None + da_after = first(da_after, dim, index_name) if da_after[dim].shape[0] else None + + if da_before is None and da_after is not None: + return da_after + if da_after is None and da_before is not None: + return da_before + + target = array[dim].dtype.type(target) + is_before_closer = abs(target - da_before[dim]) < abs(target - da_after[dim]) + nearest_array = xr.where(is_before_closer, da_before, da_after) + nearest_array[dim] = xr.where(is_before_closer, da_before[dim], da_after[dim]) + if index_name is not None: + nearest_array[index_name] = xr.where( + is_before_closer, da_before[index_name], da_after[index_name] + ) + return nearest_array + +def parallel_apply(ds, dim, func, *args): + """ + Applies a custom function in parallel along the dimension of an + xarray.Dataset or xarray.DataArray. + + The function can be any function that can be applied to an + individual xarray.Dataset or xarray.DataArray (e.g. data for a + single timestep). The function should also return data in + xarray.Dataset or xarray.DataArray format. + + This function is useful as a simple method for parallising code + that cannot easily be parallised using Dask. + + Parameters + ---------- + ds : xarray.Dataset or xarray.DataArray + xarray data with a dimension `dim` to apply the custom function + along. + dim : string + The dimension along which the custom function will be applied. + func : function + The function that will be applied in parallel to each array + along dimension `dim`. The first argument passed to this + function should be the array along `dim`. + *args : + Any number of arguments that will be passed to `func`. + + Returns + ------- + xarray.Dataset + A concatenated dataset containing an output for each array + along the input `dim` dimension. + """ + + from concurrent.futures import ProcessPoolExecutor + from tqdm import tqdm + from itertools import repeat + + with ProcessPoolExecutor() as executor: + + # Apply func in parallel + groups = [group for (i, group) in ds.groupby(dim)] + to_iterate = (groups, *(repeat(i, len(groups)) for i in args)) + out_list = list(tqdm(executor.map(func, *to_iterate), total=len(groups))) + + # Combine to match the original dataset + return xr.concat(out_list, dim=ds[dim]) + + +def pan_sharpen_brovey(band_1, band_2, band_3, pan_band): + """ + Brovey pan sharpening on surface reflectance input using numexpr + and return three xarrays. + Parameters + ---------- + band_1, band_2, band_3 : xarray.DataArray or numpy.array + Three input multispectral bands, either as xarray.DataArrays or + numpy.arrays. These bands should have already been resampled to + the spatial resolution of the panchromatic band. + pan_band : xarray.DataArray or numpy.array + A panchromatic band corresponding to the above multispectral + bands that will be used to pan-sharpen the data. + Returns + ------- + band_1_sharpen, band_2_sharpen, band_3_sharpen : numpy.arrays + Three numpy arrays equivelent to `band_1`, `band_2` and `band_3` + pan-sharpened to the spatial resolution of `pan_band`. + """ + # Calculate total + exp = 'band_1 + band_2 + band_3' + total = numexpr.evaluate(exp) + + # Perform Brovey Transform in form of: band/total*panchromatic + exp = 'a/b*c' + band_1_sharpen = numexpr.evaluate(exp, local_dict={'a': band_1, + 'b': total, + 'c': pan_band}) + band_2_sharpen = numexpr.evaluate(exp, local_dict={'a': band_2, + 'b': total, + 'c': pan_band}) + band_3_sharpen = numexpr.evaluate(exp, local_dict={'a': band_3, + 'b': total, + 'c': pan_band}) + + return band_1_sharpen, band_2_sharpen, band_3_sharpen + +def load_s1_by_orbits(dc,query): + ''' + Function to query and load ascending and descending Sentinel-1 data + and add a variable to denote acquisition orbits + + Parameters: + dc: connected datacube + query: a query dictionary to define spatial extent, measurements, time range and spatial resolution + + Returns: + Queried dataset with variable 'is_ascending' added to denote orbit path + + ''' + # load ascending data + print('\nQuerying and loading Sentinel-1 ascending data...') + ds_s1_ascending=load_ard(dc=dc,products=['s1_rtc'],resampling='bilinear', + dtype='native',sat_orbit_state='ascending',**query) + # add an variable denoting data source + ds_s1_ascending['is_ascending']=xr.DataArray(np.ones(len(ds_s1_ascending.time)), + dims=('time'),coords={'time': ds_s1_ascending.time}) + + # load descending data + print('\nQuerying and loading Sentinel-1 descending data...') + ds_s1_descending=load_ard(dc=dc,products=['s1_rtc'],resampling='bilinear', + dtype='native',sat_orbit_state='descending',**query) + # add an variable denoting data source + ds_s1_descending['is_ascending']=xr.DataArray(np.zeros(len(ds_s1_descending.time)), + dims=('time'),coords={'time': ds_s1_descending.time}) + + # merge datasets together + ds_s1=xr.concat([ds_s1_ascending,ds_s1_descending],dim='time').sortby('time') + + return ds_s1 + +def filter_obs_by_orbit(ds_s1): + ''' + Function to impliment per-pixel filtering of Sentinel-1 observations + to keep only observations from the orbit (ascending/descending) with higher frequency over time. + + Each of the Sentinel-1 observations was acquired from either a descending or ascending orbit, + which has impacts on the local incidence angle and backscattering value. + Here we do the filtering to minimise the effects of inconsistent looking angle and obit direction for each individual pixel. + + Parameters: + ds_s1: xarray.Dataset + Time-series observations of Sentinel-1 data, + with two required variables: 'is_ascending' denoting orbit path and 'mask' to identify acquisition exent + + Returns: + ds_s1_filtered: xarray.Dataset + Filtered dataset + ''' + + print('\nFiltering Sentinel-1 product by orbit...') + cnt_ascending=((ds_s1["is_ascending"]==1)&(ds_s1['mask']!=0)).sum(dim='time') + cnt_descending=((ds_s1["is_ascending"]==0)&(ds_s1['mask']!=0)).sum(dim='time') + + ds_s1_filtered=ds_s1.where(((cnt_ascending>=cnt_descending)&(ds_s1["is_ascending"]==1))| + ((cnt_ascending=0 + thresholded_ds = thresholded_ds.where(~nodata) + # use 20% ~ 80% wet frequency to identify potential coastal zone + coastal_mask=(thresholded_ds.mean(dim='time') >= 0.2)&(thresholded_ds.mean(dim='time') <= 0.8) + # buffering + print('\nApplying buffering of {} Sentinel-2 pixels (parameter buffer_pixels)...'.format(buffer_pixels)) + coastal_mask=xr.apply_ufunc(binary_dilation,coastal_mask.compute(),disk(buffer_pixels)) + return coastal_mask + +def choose_product(ds_ls,ds_s2,ds_s1,ds_ls_s2,time_step,**kwargs): + ''' + Rule-based guide on choosing the best availabel dataset in a given time step and optionally within a coastal zone mask + + Parameters: + ds_ls: xarray.Dataset + Time series Landsat data + ds_s2: xarray.Dataset + Time series Sentinel-2 data + ds_s1: xarray.Dataset + Time series Sentinel-1 data + ds_ls_s2: xarray.Dataset or None + Time series of combined Landsat and Sentinel-2 data. + time_step: string + Time step for temporal composition + **kwargs: A set of optional parameters including: + thresh_n_valid: integer + Threhold of minimum average number of valid observations within each time step + thresh_freq: float + Threshold of minimum frequency of valid observations within each time step + buffer_pixels: integer + Number of pixels to buffer coastal zone + coastal_masking: Boolean + whether to calculate a coastal zone mask and restrict the comparison of the products within the mask + Returns: + Xarray.Dataset of the best product + String of the best product name: 'ls', 's2', 's1' or 'ls_s2' + ''' + + # check if optional parameters are defined otherwise set default values + thresh_n_valid=10 if "thresh_n_valid" not in kwargs else kwargs["thresh_n_valid"] + thresh_freq=0.2 if "thresh_freq" not in kwargs else kwargs["thresh_freq"] + buffer_pixels=100 if "buffer_pixels" not in kwargs else kwargs["buffer_pixels"] + print('\nThreshold number of valid observations (parameter thresh_n_valid): {}'.format(thresh_n_valid)) + print('\nThreshold frequency of valid observations (parameter thresh_freq): {}'.format(thresh_freq)) + + # create mask if requested + coastal_masking=False if "coastal_masking" not in kwargs else kwargs["coastal_masking"] + if coastal_masking==True: + # calculate index + ds_s2 = calculate_indices(ds_s2, index='MNDWI', satellite_mission='s2') + mask=create_coastal_mask(ds_s2['MNDWI'],buffer_pixels) + else: + print('\nNo coastal masking required, using all pixels within the selected region...') + mask=None + + # calculate mean number and fraction of clear observations within each timestep and the mask + print('\nCalculating number and frequency of valid observations...') + + n_valid_obs_s2,freq_valid_s2=get_mean_number_freq_valid_obs(ds_s2['green'],mask,time_step) + print('\nSentinel-2: Average number and frequency of valid observations: {:.0f} and {:.2f}'.format(n_valid_obs_s2.mean().values,freq_valid_s2.mean().values)) + + n_valid_obs_ls,freq_valid_ls=get_mean_number_freq_valid_obs(ds_ls['green'],mask,time_step) + print('\nLandsat: Average number and frequency of valid observations: {:.0f} and {:.2f}'.format(n_valid_obs_ls.mean().values,freq_valid_ls.mean().values)) + +# n_valid_obs_s1,freq_valid_s1=get_mean_number_freq_valid_obs(ds_s1['vh'],mask,time_step) # dont need this as sentinel-1 will only be chosen when optical datasets are not sufficient + if not ds_ls_s2 is None: + n_valid_obs_ls_s2,freq_valid_ls_s2=get_mean_number_freq_valid_obs(ds_ls_s2['green'],mask,time_step) + print('\nCombined Landsat and Sentinel-2 product: Average number and frequency of valid observations: {:.0f} and {:.2f}'.format(n_valid_obs_ls_s2.mean().values,freq_valid_ls_s2.mean().values)) + + # apply decision rules + print('\nApplying rules to choose product...') + + # if Sentinel-2 meets requirements + if ((n_valid_obs_s2>=thresh_n_valid).all()) and ((freq_valid_s2>=thresh_freq).all()): + print('\nSentinel-2 product has met the minimum required average number and frequency of valid observations within all time periods') + # if combined product is available, choose combined product if it has both higher number and frequency + if not ds_ls_s2 is None: + if ((n_valid_obs_ls_s2>n_valid_obs_s2).all()) and ((freq_valid_ls_s2>freq_valid_s2).all()): + ds_selected, product_name=ds_ls_s2,'ls_s2' + print('\nChoosing combined Landsat and Sentinel-2 product as it has both higher number and frequency of valid observations within all time periods') + else: + ds_selected, product_name=ds_s2,'s2' + print('\nChoosing Sentinel-2 product as neither Landsat or the combined product meets both requirements or is significantly better than Sentinel-2') + # if combined product is unavailable, choose Landsat if it has both higher number and frequency + elif ((n_valid_obs_ls>=n_valid_obs_s2).all()) and ((freq_valid_ls>=freq_valid_s2).all()): + ds_selected, product_name=ds_ls,'ls' + print('\nChoosing Landsat product as it has both higher average number and frequency of valid observations within all time periods') + # otherwise choose Sentinel-2 + else: + ds_selected, product_name=ds_s2,'s2' + print('\nChoosing Sentinel-2 product as Landsat product does not meet both requirements or is not significantly better than Sentinel-2') + # if Sentinel-2 doesn't meet both requirements,but Landsat does, either choose Landsat or combined product if available + elif ((n_valid_obs_ls>=thresh_n_valid).all()) and ((freq_valid_ls>=thresh_freq).all()): + print('\nSentinel-2 does not meet the minimum required average number and frequency of valid observations within all time periods, but Landsat does') + if not ds_ls_s2 is None: + ds_selected, product_name=ds_ls_s2,'ls_s2' + print('\nChoosing combined Landsat and Sentinel-2 product as it has both higher number and frequency of valid observations within all time periods') + else: + ds_selected, product_name=ds_ls,'ls' + print('\nChoosing Landsat product') + # if neither Sentinel-2 or Landsat meet both requirements, choose combined product if it meets requirements + elif not ds_ls_s2 is None: + print('\nNeither Sentinel-2 or Landsat meets the minimum required average number and frequency of valid observations within all time periods') + # but the combined product meet requirements + if ((n_valid_obs_ls_s2>=thresh_n_valid).all()) and ((freq_valid_ls_s2>=thresh_freq).all()): + ds_selected, product_name=ds_ls_s2,'ls_s2' + print('\nChoosing combined Landsat and Sentinel-2 product as it meets the minimum required average number and frequency of valid observations within all time periods') + else: + ds_selected, product_name=ds_s1,'s1' + print('\nChoosing Sentinel-1 as no other products available that meet the requirements') + # otherwise choose Sentinel-1 + else: + print('\nNeither Sentinel-2 or Landsat meets the minimum required average number and frequency of valid observations within all time periods') + ds_selected, product_name=ds_s1,'s1' + print('\nChoosing Sentinel-1 product as no other products available that meet the requirements') + + print('\nBest available product: ',product_name) + return ds_selected, product_name + +def load_combined_ls_s2(dc,query): + '''function to query and load combined Landsat and Sentinel-2 data + + Parameters: + dc: connected datacube + query: a query dictionary to define spatial extent, time range, measurements and spatial resolution for both datasets + + Returns: + ds_combined: Combined data as xarray.Dataset + ''' + print('Querying and loading combined Landsat and Sentinel-2 products...') + # Load available Landsat data resampled to Sentinel-2 resolution + ds_ls = load_ard(dc=dc, products=['ls8_sr', 'ls9_sr'],align=(10, 10), + resampling='bilinear',**query) + + # add an variable denoting data source (for future analysis) + is_ls=xr.DataArray(np.ones(len(ds_ls.time)),dims=('time'),coords={'time': ds_ls.time}) + ds_ls['is_ls'] = is_ls + + # Load Sentinel-2 data + ds_s2 = load_ard(dc=dc,products=['s2_l2a'],resampling='bilinear', + align=(10, 10),mask_filters=[("opening", 2), ("dilation", 5)],**query) + # add an variable denoting data source (for future analysis) + is_ls=xr.DataArray(np.zeros(len(ds_s2.time)),dims=('time'),coords={'time': ds_s2.time}) + ds_s2['is_ls'] = is_ls + + # merge two datasets together + ds_combined=xr.concat([ds_ls,ds_s2],dim='time').sortby('time') + + return ds_combined + +def load_best_available_ds(dc, lat_range, lon_range, time_range, time_step, **kwargs): + ''' + Function to query, load and compare different products, select and return the best available product + + Parameters: + dc: connected datacube + lat_range: range of latitudes in tuple or list + lon_range: range of longitude in tuple or list + time_range: range of time to query the data in tuple or list + time_step: string, pre-defined time step for temporal aggregation, e.g. '1Y' + **kwargs: A set of optional parameters on data query or comparison between products which may include: + combine_ls_s2: A boolean value indicating whether to include merged/stacked Landsat and Sentinel-2 products as an option. Default to False. + set_resolution: integer of spatial resolution in metres to query all products + coastal_masking: A boolean value indicating whether to calculate a mask + and restrict the comparison of the products within the masked zone. + set_product: Set this to only query and load a pre-selected product, 'ls','s2','ls_s2' or 's1' + i.e. no other products will be queried or compared. + thresh_n_valid: Threhold of minimum average number of valid observations within each time step, integer + thresh_freq: Threshold of minimum frequency of valid observations within each time step, float between 0~1 + buffer_pixels: Number of pixels to buffer coastal zone, integer + + Returns: + ds_selected: selected product as xarray.Dataset + product_name: name of selected product in string format, i.e. 'ls','s2','ls_s2','s1' + ''' + # parse input time range to accommodate queries before and after 2017 + min_time=min(parser.parse(time_range_i,default=datetime(1987,1,1,0,0)) + for time_range_i in time_range) + if min_time ds.lat.max(): +# warnings.warn( +# "Lats must be in range {} .. {}. Got: {}".format( +# ds.lat.min().values, ds.lat.max().values, lat +# ) +# ) +# if min(lon) < ds.lon.min() or max(lon) > ds.lon.max(): +# warnings.warn( +# "Lons must be in range {} .. {}. Got: {}".format( +# ds.lon.min().values, ds.lon.max().values, lon +# ) +# ) +# # Find existing coords between min&max +# lats = ds.lat[np.logical_and(ds.lat >= min(lat), ds.lat <= max(lat))].values +# # If there was nothing between, just plan to grab closest +# if len(lats) == 0: +# lats = np.unique(ds.lat.sel(lat=np.array(lat), method="nearest")) +# lons = ds.lon[np.logical_and(ds.lon >= min(lon), ds.lon <= max(lon))].values +# if len(lons) == 0: +# lons = np.unique(ds.lon.sel(lon=np.array(lon), method="nearest")) +# # crop and keep attrs +# output = ds.sel(lat=lats, lon=lons) +# output.attrs = ds.attrs +# for var in output.data_vars: +# output[var].attrs = ds[var].attrs +# return output + + +# def era5_area_nearest(ds, lat, lon): +# """ +# Crop a dataset containing EAR5 variables to a location. +# The output spatial grid is snapped to the nearest input grid points. + +# Parameters +# ---------- +# ds : xarray dataset +# A dataset containing ERA5 variables of interest. + +# lat: tuple or list +# Latitude range for query. + +# lon: tuple or list +# Longitude range for query. + +# Returns +# ------- +# An xarray dataset containing ERA5 variables for the selected location. + +# """ + +# if min(lon) < 0: +# # re-order along longitude to go from -180 to 180 +# ds = ds.assign_coords({"lon": (((ds.lon + 180) % 360) - 180)}) +# ds = ds.reindex({"lon": np.sort(ds.lon)}) + +# # find the nearest lat lon boundary points +# test = ds.sel(lat=lat, lon=lon, method="nearest") +# # define the lat/lon grid +# lat_range = slice(test.lat.max().values, test.lat.min().values) +# lon_range = slice(test.lon.min().values, test.lon.max().values) +# # crop and keep attrs +# output = ds.sel(lat=lat_range, lon=lon_range) +# output.attrs = ds.attrs +# for var in output.data_vars: +# output[var].attrs = ds[var].attrs +# return output + + +# def load_era5_netcdf(var, lat, lon, time, grid="nearest", **kwargs): +# """ +# Returns a ERA5 variable for a selected location and time window. + +# Parameters +# ---------- +# var : string +# Name of the ERA5 climate variable to download, e.g "air_temperature_at_2_metres" + +# lat: tuple or list +# Latitude range for query. + +# lon: tuple or list +# Longitude range for query. + +# time: tuple or list +# Time range for query. + +# grid: string +# Option for output spatial gridding. +# The default is 'nearest', for which output spatial grid is snapped to the nearest ERA5 input grid points. +# Alternatively, output spatial grid will either include input grid points within lat/lon boundaries or the nearest point if none is within the search location. + +# Returns +# ------- +# An xarray dataset containing the variable for the selected location and time window. + +# """ + +# ds = get_era5_daily(var, time[0], time[1], **kwargs) +# if grid == "nearest": +# return era5_area_nearest(ds, lat, lon).compute() +# else: +# return era5_area_crop(ds, lat, lon).compute() diff --git a/deafrica_tools/load_isda.py b/deafrica_tools/load_isda.py new file mode 100644 index 0000000..2e9cca1 --- /dev/null +++ b/deafrica_tools/load_isda.py @@ -0,0 +1,92 @@ +""" +Functions to retrieve iSDAsoil data. +""" + +import numpy as np +import pandas as pd +import xarray as xr +import matplotlib.pyplot as plt +import matplotlib.patches as mpatches +import rasterio as rio +from pyproj import Transformer +import matplotlib.pyplot as plt +import os +import numpy as np + +from urllib.parse import urlparse +import boto3 +from pystac import stac_io, Catalog + +#this function allows us to directly query the data on s3, adapted from iSDA tutorial https://github.com/iSDA-Africa/isdasoil-tutorial/blob/main/iSDAsoil-tutorial.ipynb +def my_read_method(uri): + parsed = urlparse(uri) + if parsed.scheme == 's3': + bucket = parsed.netloc + key = parsed.path[1:] + s3 = boto3.resource('s3') + obj = s3.Object(bucket, key) + return obj.get()['Body'].read().decode('utf-8') + else: + return stac_io.default_read_text_method(uri) + +stac_io.read_text_method = my_read_method + +catalog = Catalog.from_file("https://isdasoil.s3.amazonaws.com/catalog.json") + +assets = {} + +for root, catalogs, items in catalog.walk(): + for item in items: + str(f"Type: {item.get_parent().title}") + # save all items to a dictionary as we go along + assets[item.id] = item + for asset in item.assets.values(): + if asset.roles == ['data']: + str(f"Title: {asset.title}") + str(f"Description: {asset.description}") + str(f"URL: {asset.href}") + str("------------") + +# define load_isda() function + +def load_isda(var, lat, lon): + """ + Download and return iSDA variable with number of bands corresponding to number of iSDA layers. + Parameters + ---------- + var : string + Name of the iSDA variable to download, e.g "ph" + lat: tuple or list + Latitude range for query. + lon: tuple or list + Longitude range for query. + """ + + bands = assets[var].assets["image"].extra_fields.get('eo:bands') + bands = [val['description'] for val in bands] + + if len(np.unique(bands)) > 1: + + ds = xr.open_dataset(assets[var].assets["image"].href, engine="rasterio").rio.clip_box( + minx=lon[0], + miny=lat[0], + maxx=lon[1], + maxy=lat[1], + crs="EPSG:4326", + ) + + ds_layered = ds.drop_dims('band') + for x in np.unique(ds.band): + ds_layered[bands[x-1]] = ds.sel(band=x).to_array(dim='band').squeeze() + + else: + + ds_layered = xr.open_dataset(assets[var].assets["image"].href, engine="rasterio").rio.clip_box( + minx=lon[0], + miny=lat[0], + maxx=lon[1], + maxy=lat[1], + crs="EPSG:4326", + ).squeeze() + + return ds_layered diff --git a/deafrica_tools/load_soil_moisture.py b/deafrica_tools/load_soil_moisture.py new file mode 100644 index 0000000..4fb1a58 --- /dev/null +++ b/deafrica_tools/load_soil_moisture.py @@ -0,0 +1,39 @@ +import xarray as xr +import numpy as np + +# function to load soil moisture data + +def load_soil_moisture(lat, lon, time, product = 'surface', grid = 'nearest'): + product_baseurl = 'https://dapds00.nci.org.au/thredds/dodsC/ub8/global/GRAFS/' + assert product in ['surface', 'rootzone'], 'product parameter must be surface or root-zone' + # lat, lon grid + if grid == 'nearest': + # select lat/lon range from data; snap to nearest grid + lat_range, lon_range = None, None + else: + # define a grid that covers the entire area of interest + lat_range = np.arange(np.max(np.ceil(np.array(lat)*10.+0.5)/10.-0.05), np.min(np.floor(np.array(lat)*10.-0.5)/10.+0.05)-0.05, -0.1) + lon_range = np.arange(np.min(np.floor(np.array(lon)*10.-0.5)/10.+0.05), np.max(np.ceil(np.array(lon)*10.+0.5)/10.-0.05)+0.05, 0.1) + # split time window into years + day_range = np.array(time).astype("M8[D]") + year_range = np.array(time).astype("M8[Y]") + if product == 'surface': + product_name = 'GRAFS_TopSoilRelativeWetness_' + else: product_name = 'GRAFS_RootzoneSoilWaterIndex_' + datasets = [] + for year in np.arange(year_range[0], year_range[1]+1, np.timedelta64(1, 'Y')): + start = np.max([day_range[0], year.astype("M8[D]")]) + end = np.min([day_range[1], (year+1).astype("M8[D]")-1]) + product_url = product_baseurl + product_name +'%s.nc'%str(year) + print(product_url) + # data is loaded lazily through OPeNDAP + ds = xr.open_dataset(product_url) + if lat_range is None: + # select lat/lon range from data if not specified; snap to nearest grid + test = ds.sel(lat=list(lat), lon=list(lon), method='nearest') + lat_range = slice(test.lat.values[0], test.lat.values[1]) + lon_range = slice(test.lon.values[0], test.lon.values[1]) + # slice before return + ds = ds.sel(lat=lat_range, lon=lon_range, time=slice(start, end)).compute() + datasets.append(ds) + return xr.merge(datasets) \ No newline at end of file diff --git a/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.mo b/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.mo new file mode 100644 index 0000000..6b7f6ac Binary files /dev/null and b/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.mo differ diff --git a/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.po b/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.po new file mode 100644 index 0000000..06d9e91 --- /dev/null +++ b/deafrica_tools/locales/fr/LC_MESSAGES/deafrica_tools.po @@ -0,0 +1,117 @@ +msgid "" +msgstr "" +"MIME-Version: 1.0\n" +"Content-Type: text/plain; charset=UTF-8\n" +"Content-Transfer-Encoding: 8bit\n" +"X-Generator: POEditor.com\n" +"Project-Id-Version: deafrica_tools\n" +"Language: fr\n" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:83 +msgid "None" +msgstr "Aucun" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:84 +msgid "ESRI World Imagery" +msgstr "Imagerie mondiale ESRI" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:85 +msgid "Sentinel-2 Geomedian" +msgstr "Sentinel-2 Geomedian" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:86 +msgid "Water Observations from Space" +msgstr "Observations de l'eau depuis l'espace-WOfS" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:99 +msgid "Wetlands Insight Tool" +msgstr "Outil d'analyse des zones humides" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:100 +msgid "Select parameters and AOI" +msgstr "Sélectionner les paramètres et la zone d'intérêt" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:124 +msgid "Total polygon area" +msgstr "Superficie totale du polygone" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:128 +msgid "Area falls within recommended limit" +msgstr "La zone se situe dans la limite recommandée" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:131 +msgid "Area is too large, please update your polygon" +msgstr "La zone est trop grande, veuillez réduire votre polygone" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:151 +msgid "Map Overlays" +msgstr "Superpositions de cartes" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:176 +msgid "Run" +msgstr "Exécuter" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:183 +msgid "Map Overlay:" +msgstr "Carte superposée :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:185 +msgid "Start Date:" +msgstr "Date de début :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:187 +msgid "End Date:" +msgstr "Date de fin :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:189 +msgid "Minimum Good Data:" +msgstr "Minimum de bonnes données :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:191 +msgid "Resampling Frequency:" +msgstr "Fréquence de rééchantillonnage :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:193 +msgid "Output CSV:" +msgstr "Sortie CSV :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:195 +msgid "Output Plot:" +msgstr "Tracé de sortie :" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:308 +msgid "Progress" +msgstr "Progrès" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:326 +msgid "WIT complete" +msgstr "WIT achevée" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:328 +msgid "No polygon selected" +msgstr "Aucun polygone sélectionné" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:365 +msgid "open water" +msgstr "eau libre" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:366 +msgid "wet" +msgstr "humide" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:367 +msgid "green veg" +msgstr "végétation verts" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:368 +msgid "dry veg" +msgstr "végétation seche" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:369 +msgid "bare soil" +msgstr "sol nu" + +#: Tools/deafrica_tools/app/wetlandsinsighttool.py:382 +msgid "Percentage Fractional Cover, Wetness, and Water" +msgstr "Pourcentage de couverture fractionnée, humidité et eau" + diff --git a/deafrica_tools/plotting.py b/deafrica_tools/plotting.py new file mode 100644 index 0000000..d554a4c --- /dev/null +++ b/deafrica_tools/plotting.py @@ -0,0 +1,1272 @@ +""" +Functions for plotting Digital Earth Africa data. +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import math +import folium +import ipywidgets +import branca +import numpy as np +import geopandas as gpd +import matplotlib as mpl +import matplotlib.patheffects as PathEffects +import matplotlib.pyplot as plt +import matplotlib.animation as animation +from datetime import datetime +import matplotlib.cm as cm +from matplotlib import colors as mcolours +from pyproj import Transformer +from IPython.display import display +from matplotlib.colors import ListedColormap +from mpl_toolkits.axes_grid1.inset_locator import inset_axes +from mpl_toolkits.axes_grid1 import make_axes_locatable +from ipyleaflet import Map, Marker, Popup, GeoJSON, basemaps, Choropleth +from skimage import exposure +from branca.colormap import linear +from odc.ui import image_aspect + +from matplotlib.animation import FuncAnimation +import pandas as pd +from pathlib import Path +from shapely.geometry import box +from skimage.exposure import rescale_intensity +from tqdm.auto import tqdm +import warnings + + +def rgb( + ds, + bands=["red", "green", "blue"], + index=None, + index_dim="time", + robust=True, + percentile_stretch=None, + col_wrap=4, + size=6, + aspect=None, + savefig_path=None, + savefig_kwargs={}, + **kwargs, +): + + """ + Takes an xarray dataset and plots RGB images using three imagery + bands (e.g ['red', 'green', 'blue']). The `index` + parameter allows easily selecting individual or multiple images for + RGB plotting. Images can be saved to file by specifying an output + path using `savefig_path`. + This function was designed to work as an easier-to-use wrapper + around xarray's `.plot.imshow()` functionality. + + Last modified: April 2021 + + Parameters + ---------- + ds : xarray Dataset + A two-dimensional or multi-dimensional array to plot as an RGB + image. If the array has more than two dimensions (e.g. multiple + observations along a 'time' dimension), either use `index` to + select one (`index=0`) or multiple observations + (`index=[0, 1]`), or create a custom faceted plot using e.g. + `col="time"`. + bands : list of strings, optional + A list of three strings giving the band names to plot. Defaults + to '['red', 'green', 'blue']'. If the dataset does not contain + bands named `'red', 'green', 'blue'`, then `bands` must be + specified. + index : integer or list of integers, optional + `index` can be used to select one (`index=0`) or multiple + observations (`index=[0, 1]`) from the input dataset for + plotting. If multiple images are requested these will be plotted + as a faceted plot. + index_dim : string, optional + The dimension along which observations should be plotted if + multiple observations are requested using `index`. Defaults to + `time`. + robust : bool, optional + Produces an enhanced image where the colormap range is computed + with 2nd and 98th percentiles instead of the extreme values. + Defaults to True. + percentile_stretch : tuple of floats + An tuple of two floats (between 0.00 and 1.00) that can be used + to clip the colormap range to manually specified percentiles to + get more control over the brightness and contrast of the image. + The default is None; '(0.02, 0.98)' is equivelent to + `robust=True`. If this parameter is used, `robust` will have no + effect. + col_wrap : integer, optional + The number of columns allowed in faceted plots. Defaults to 4. + size : integer, optional + The height (in inches) of each plot. Defaults to 6. + aspect : integer, optional + Aspect ratio of each facet in the plot, so that aspect * size + gives width of each facet in inches. Defaults to None, which + will calculate the aspect based on the x and y dimensions of + the input data. + savefig_path : string, optional + Path to export image file for the RGB plot. Defaults to None, + which does not export an image file. + savefig_kwargs : dict, optional + A dict of keyword arguments to pass to + `matplotlib.pyplot.savefig` when exporting an image file. For + all available options, see: + https://matplotlib.org/api/_as_gen/matplotlib.pyplot.savefig.html + **kwargs : optional + Additional keyword arguments to pass to `xarray.plot.imshow()`. + For example, the function can be used to plot into an existing + matplotlib axes object by passing an `ax` keyword argument. + For more options, see: + http://xarray.pydata.org/en/stable/generated/xarray.plot.imshow.html + Returns + ------- + An RGB plot of one or multiple observations, and optionally an image + file written to file. + """ + + # If bands are not in the dataset + ds_vars = list(ds.data_vars) + if set(bands).issubset(ds_vars) == False: + raise ValueError( + "rgb() bands do not match band names in dataset. " + "Note the default rgb() bands are ['red', 'green', 'blue']." + ) + + # If ax is supplied via kwargs, ignore aspect and size + if "ax" in kwargs: + + # Create empty aspect size kwarg that will be passed to imshow + aspect_size_kwarg = {} + else: + # Compute image aspect + if not aspect: + aspect = image_aspect(ds) + + # Populate aspect size kwarg with aspect and size data + aspect_size_kwarg = {"aspect": aspect, "size": size} + + # If no value is supplied for `index` (the default), plot using default + # values and arguments passed via `**kwargs` + if index is None: + + # Select bands and convert to DataArray + da = ds[bands].to_array() + + # If percentile_stretch == True, clip plotting to percentile vmin, vmax + if percentile_stretch: + vmin, vmax = da.compute().quantile(percentile_stretch).values + kwargs.update({"vmin": vmin, "vmax": vmax}) + + # If there are more than three dimensions and the index dimension == 1, + # squeeze this dimension out to remove it + if (len(ds.dims) > 2) and ("col" not in kwargs) and (len(da[index_dim]) == 1): + + da = da.squeeze(dim=index_dim) + + # If there are more than three dimensions and the index dimension + # is longer than 1, raise exception to tell user to use 'col'/`index` + elif (len(ds.dims) > 2) and ("col" not in kwargs) and (len(da[index_dim]) > 1): + + raise Exception( + f"The input dataset `ds` has more than two dimensions: " + "{list(ds.dims.keys())}. Please select a single observation " + "using e.g. `index=0`, or enable faceted plotting by adding " + 'the arguments e.g. `col="time", col_wrap=4` to the function ' + "call" + ) + da = da.compute() + img = da.plot.imshow( + robust=robust, col_wrap=col_wrap, **aspect_size_kwarg, **kwargs + ) + + # If values provided for `index`, extract corresponding observations and + # plot as either single image or facet plot + else: + + # If a float is supplied instead of an integer index, raise exception + if isinstance(index, float): + raise Exception( + f"Please supply `index` as either an integer or a list of " "integers" + ) + + # If col argument is supplied as well as `index`, raise exception + if "col" in kwargs: + raise Exception( + f"Cannot supply both `index` and `col`; please remove one and " + "try again" + ) + + # Convert index to generic type list so that number of indices supplied + # can be computed + index = index if isinstance(index, list) else [index] + + # Select bands and observations and convert to DataArray + da = ds[bands].isel(**{index_dim: index}).to_array().compute() + + # If percentile_stretch == True, clip plotting to percentile vmin, vmax + if percentile_stretch: + vmin, vmax = da.compute().quantile(percentile_stretch).values + kwargs.update({"vmin": vmin, "vmax": vmax}) + + # If multiple index values are supplied, plot as a faceted plot + if len(index) > 1: + + img = da.plot.imshow( + robust=robust, + col=index_dim, + col_wrap=col_wrap, + **aspect_size_kwarg, + **kwargs, + ) + + # If only one index is supplied, squeeze out index_dim and plot as a + # single panel + else: + + img = da.squeeze(dim=index_dim).plot.imshow( + robust=robust, **aspect_size_kwarg, **kwargs + ) + + # If an export path is provided, save image to file. Individual and + # faceted plots have a different API (figure vs fig) so we get around this + # using a try statement: + if savefig_path: + + print(f"Exporting image to {savefig_path}") + + try: + img.fig.savefig(savefig_path, **savefig_kwargs) + except: + img.figure.savefig(savefig_path, **savefig_kwargs) + + +def display_map(x, y, crs="EPSG:4326", margin=-0.5, zoom_bias=0): + """ + Given a set of x and y coordinates, this function generates an + interactive map with a bounded rectangle overlayed on Google Maps + imagery. + + Last modified: September 2019 + + Modified from function written by Otto Wagner available here: + https://github.com/ceos-seo/data_cube_utilities/tree/master/data_cube_utilities + + Parameters + ---------- + x : (float, float) + A tuple of x coordinates in (min, max) format. + y : (float, float) + A tuple of y coordinates in (min, max) format. + crs : string, optional + A string giving the EPSG CRS code of the supplied coordinates. + The default is 'EPSG:4326'. + margin : float + A numeric value giving the number of degrees lat-long to pad + the edges of the rectangular overlay polygon. A larger value + results more space between the edge of the plot and the sides + of the polygon. Defaults to -0.5. + zoom_bias : float or int + A numeric value allowing you to increase or decrease the zoom + level by one step. Defaults to 0; set to greater than 0 to zoom + in, and less than 0 to zoom out. + Returns + ------- + folium.Map : A map centered on the supplied coordinate bounds. A + rectangle is drawn on this map detailing the perimeter of the x, y + bounds. A zoom level is calculated such that the resulting + viewport is the closest it can possibly get to the centered + bounding rectangle without clipping it. + """ + + # Convert each corner coordinates to lat-lon + all_x = (x[0], x[1], x[0], x[1]) + all_y = (y[0], y[0], y[1], y[1]) + transformer = Transformer.from_crs(crs, "EPSG:4326") + all_longitude, all_latitude = transformer.transform(all_x, all_y) + + # Calculate zoom level based on coordinates + lat_zoom_level = ( + _degree_to_zoom_level(min(all_latitude), max(all_latitude), margin=margin) + + zoom_bias + ) + lon_zoom_level = ( + _degree_to_zoom_level(min(all_longitude), max(all_longitude), margin=margin) + + zoom_bias + ) + zoom_level = min(lat_zoom_level, lon_zoom_level) + + # Identify centre point for plotting + center = [np.mean(all_latitude), np.mean(all_longitude)] + + # Create map + interactive_map = folium.Map( + location=center, + zoom_start=zoom_level, + tiles="http://mt1.google.com/vt/lyrs=y&z={z}&x={x}&y={y}", + attr="Google", + ) + + # Create bounding box coordinates to overlay on map + line_segments = [ + (all_latitude[0], all_longitude[0]), + (all_latitude[1], all_longitude[1]), + (all_latitude[3], all_longitude[3]), + (all_latitude[2], all_longitude[2]), + (all_latitude[0], all_longitude[0]), + ] + + # Add bounding box as an overlay + interactive_map.add_child( + folium.features.PolyLine(locations=line_segments, color="red", opacity=0.8) + ) + + # Add clickable lat-lon popup box + interactive_map.add_child(folium.features.LatLngPopup()) + + return interactive_map + + +def map_shapefile( + gdf, + attribute, + continuous=False, + cmap="viridis", + basemap=basemaps.Esri.WorldImagery, + default_zoom=None, + hover_col=True, + **style_kwargs, +): + """ + Plots a geopandas GeoDataFrame over an interactive ipyleaflet + basemap, with features coloured based on attribute column values. + Optionally, can be set up to print selected data from features in + the GeoDataFrame. + + Last modified: February 2020 + + Parameters + ---------- + gdf : geopandas.GeoDataFrame + A GeoDataFrame containing the spatial features to be plotted + over the basemap. + attribute: string, required + An required string giving the name of any column in the + GeoDataFrame you wish to have coloured on the choropleth. + continuous: boolean, optional + Whether to plot data as a categorical or continuous variable. + Defaults to remapping the attribute which is suitable for + categorical data. For continuous data set `continuous` to True. + cmap : string, optional + A string giving the name of a `matplotlib.cm` colormap that will + be used to style the features in the GeoDataFrame. Features will + be coloured based on the selected attribute. Defaults to the + 'viridis' colormap. + basemap : ipyleaflet.basemaps object, optional + An optional `ipyleaflet.basemaps` object used as the basemap for + the interactive plot. Defaults to `basemaps.Esri.WorldImagery`. + default_zoom : int, optional + An optional integer giving a default zoom level for the + interactive ipyleaflet plot. Defaults to None, which infers + the zoom level from the extent of the data. + hover_col : boolean or str, optional + If True (the default), the function will print values from the + GeoDataFrame's `attribute` column above the interactive map when + a user hovers over the features in the map. Alternatively, a + custom shapefile field can be specified by supplying a string + giving the name of the field to print. Set to False to prevent + any attributes from being printed. + **style_kwargs : + Optional keyword arguments to pass to the `style` paramemter of + the `ipyleaflet.Choropleth` function. This can be used to + control the appearance of the shapefile, for example 'stroke' + and 'weight' (controlling line width), 'fillOpacity' (polygon + transparency) and 'dashArray' (whether to plot lines/outlines + with dashes). For more information: + https://ipyleaflet.readthedocs.io/en/latest/api_reference/choropleth.html + """ + + def on_hover(event, id, properties): + with dbg: + text = properties.get(hover_col, "???") + lbl.value = f"{hover_col}: {text}" + + # Verify that attribute exists in shapefile + if attribute not in gdf.columns: + raise ValueError( + f"The `attribute` {attribute} does not exist " + f"in the geopandas.GeoDataFrame. " + f"Valid attributes include {gdf.columns.values}." + ) + + # If hover_col is True, use 'attribute' as the default hover attribute. + # Otherwise, hover_col will use the supplied attribute field name + if hover_col and (hover_col is True): + hover_col = attribute + + # If a custom string if supplied to hover_col, check this exists + elif hover_col and (type(hover_col) == str): + if hover_col not in gdf.columns: + raise ValueError( + f"The `hover_col` field {hover_col} does " + f"not exist in the geopandas.GeoDataFrame. " + f"Valid attributes include " + f"{gdf.columns.values}." + ) + + # Convert to WGS 84 and GeoJSON format + gdf_wgs84 = gdf.to_crs(epsg=4326) + data_geojson = gdf_wgs84.__geo_interface__ + + # If continuous is False, remap categorical classes for visualisation + if not continuous: + + # Zip classes data together to make a dictionary + classes_uni = list(gdf[attribute].unique()) + classes_clean = list(range(0, len(classes_uni))) + classes_dict = dict(zip(classes_uni, classes_clean)) + + # Get values to colour by as a list + classes = gdf[attribute].map(classes_dict).tolist() + + # If continuous is True then do not remap + else: + + # Get values to colour by as a list + classes = gdf[attribute].tolist() + + # Create the dictionary to colour map by + keys = gdf.index + id_class_dict = dict(zip(keys.astype(str), classes)) + + # Get centroid to focus map on + lon1, lat1, lon2, lat2 = gdf_wgs84.total_bounds + lon = (lon1 + lon2) / 2 + lat = (lat1 + lat2) / 2 + + if default_zoom is None: + + # Calculate default zoom from latitude of features + default_zoom = _degree_to_zoom_level(lat1, lat2, margin=-0.5) + + # Plot map + m = Map( + center=(lat, lon), + zoom=default_zoom, + basemap=basemap, + layout=dict(width="800px", height="600px"), + ) + + # Define default plotting parameters for the choropleth map. + # The nested dict structure sets default values which can be + # overwritten/customised by `choropleth_kwargs` values + style_kwargs = dict({"fillOpacity": 0.8}, **style_kwargs) + + # Get `branca.colormap` object from matplotlib string + cm_cmap = cm.get_cmap(cmap, 30) + colormap = branca.colormap.LinearColormap( + [cm_cmap(i) for i in np.linspace(0, 1, 30)] + ) + + # Create the choropleth + choropleth = Choropleth( + geo_data=data_geojson, + choro_data=id_class_dict, + colormap=colormap, + style=style_kwargs, + ) + + # If the vector data contains line features, they will not be + # be coloured by default. To resolve this, we need to manually copy + # across the 'fillColor' attribute to the 'color' attribute for each + # feature, then plot the data as a GeoJSON layer rather than the + # choropleth layer that we use for polygon data. + linefeatures = any( + x in ["LineString", "MultiLineString"] for x in gdf.geometry.type.values + ) + if linefeatures: + + # Copy colour from fill to line edge colour + for i in keys: + choropleth.data["features"][i]["properties"]["style"][ + "color" + ] = choropleth.data["features"][i]["properties"]["style"]["fillColor"] + + # Add GeoJSON layer to map + feature_layer = GeoJSON(data=choropleth.data, style=style_kwargs) + m.add_layer(feature_layer) + + else: + + # Add Choropleth layer to map + m.add_layer(choropleth) + + # If a column is specified by `hover_col`, print data from the + # hovered feature above the map + if hover_col and not linefeatures: + + # Use cholopleth object if data is polygon + lbl = ipywidgets.Label() + dbg = ipywidgets.Output() + choropleth.on_hover(on_hover) + display(lbl) + + else: + + lbl = ipywidgets.Label() + dbg = ipywidgets.Output() + feature_layer.on_hover(on_hover) + display(lbl) + + # Display the map + display(m) + + +def xr_animation( + ds, + bands=None, + output_path="animation.mp4", + width_pixels=500, + interval=100, + percentile_stretch=(0.02, 0.98), + image_proc_funcs=None, + show_gdf=None, + show_date="%d %b %Y", + show_text=None, + show_colorbar=True, + gdf_kwargs={}, + annotation_kwargs={}, + imshow_kwargs={}, + colorbar_kwargs={}, + limit=None, +): + """ + Takes an `xarray` timeseries and animates the data as either a + three-band (e.g. true or false colour) or single-band animation, + allowing changes in the landscape to be compared across time. + + Animations can be customised to include text and date annotations + or use specific combinations of input bands. Vector data can be + overlaid and animated on top of imagery, and custom image + processing functions can be applied to each frame. + Supports .mp4 (ideal for Twitter/social media) and .gif (ideal + for all purposes, but can have large file sizes) format files. + + Last modified: October 2020 + + Parameters + ---------- + ds : xarray.Dataset + An xarray dataset with multiple time steps (i.e. multiple + observations along the `time` dimension). + bands : list of strings + An list of either one or three band names to be plotted, + all of which must exist in `ds`. + output_path : str, optional + A string giving the output location and filename of the + resulting animation. File extensions of '.mp4' and '.gif' are + accepted. Defaults to 'animation.mp4'. + width_pixels : int, optional + An integer defining the output width in pixels for the + resulting animation. The height of the animation is set + automatically based on the dimensions/ratio of the input + xarray dataset. Defaults to 500 pixels wide. + interval : int, optional + An integer defining the milliseconds between each animation + frame used to control the speed of the output animation. Higher + values result in a slower animation. Defaults to 100 + milliseconds between each frame. + percentile_stretch : tuple of floats, optional + An optional tuple of two floats that can be used to clip one or + three-band arrays by percentiles to produce a more vibrant, + visually attractive image that is not affected by outliers/ + extreme values. The default is `(0.02, 0.98)` which is + equivalent to xarray's `robust=True`. This parameter is ignored + completely if `vmin` and `vmax` are provided as kwargs to + `imshow_kwargs`. + image_proc_funcs : list of funcs, optional + An optional list containing functions that will be applied to + each animation frame (timestep) prior to animating. This can + include image processing functions such as increasing contrast, + unsharp masking, saturation etc. The function should take AND + return a `numpy.ndarray` with shape [y, x, bands]. If your + function has parameters, you can pass in custom values using + a lambda function: + `image_proc_funcs=[lambda x: custom_func(x, param1=10)]`. + show_gdf: geopandas.GeoDataFrame, optional + Vector data (e.g. ESRI shapefiles or GeoJSON) can be optionally + plotted over the top of imagery by supplying a + `geopandas.GeoDataFrame` object. To customise colours used to + plot the vector features, create a new column in the + GeoDataFrame called 'colors' specifying the colour used to plot + each feature: e.g. `gdf['colors'] = 'red'`. + To plot vector features at specific moments in time during the + animation, create new 'start_time' and/or 'end_time' columns in + the GeoDataFrame that define the time range used to plot each + feature. Dates can be provided in any string format that can be + converted using the `pandas.to_datetime()`. e.g. + `gdf['end_time'] = ['2001', '2005-01', '2009-01-01']` + show_date : string or bool, optional + An optional string or bool that defines how (or if) to plot + date annotations for each animation frame. Defaults to + '%d %b %Y'; can be customised to any format understood by + strftime (https://strftime.org/). Set to False to remove date + annotations completely. + show_text : str or list of strings, optional + An optional string or list of strings with a length equal to + the number of timesteps in `ds`. This can be used to display a + static text annotation (using a string), or a dynamic title + (using a list) that displays different text for each timestep. + By default, no text annotation will be plotted. + show_colorbar : bool, optional + An optional boolean indicating whether to include a colourbar + for single-band animations. Defaults to True. + gdf_kwargs : dict, optional + An optional dictionary of keyword arguments to customise the + appearance of a `geopandas.GeoDataFrame` supplied to + `show_gdf`. Keyword arguments are passed to `GeoSeries.plot` + (see http://geopandas.org/reference.html#geopandas.GeoSeries.plot). + For example: `gdf_kwargs = {'linewidth': 2}`. + annotation_kwargs : dict, optional + An optional dict of keyword arguments for controlling the + appearance of text annotations. Keyword arguments are passed + to `plt.annotate` from `matplotlib`. + (see https://matplotlib.org/api/_as_gen/matplotlib.pyplot.annotate.html + for options). For example, + `annotation_kwargs={'fontsize':20, 'color':'red', 'family':'serif'}`. + imshow_kwargs : dict, optional + An optional dict of keyword arguments for controlling the + appearance of arrays passed to `matplotlib`'s `plt.imshow` + (see https://matplotlib.org/api/_as_gen/matplotlib.pyplot.imshow.html + for options). For example, a green colour scheme and custom + stretch could be specified using: + `onebandplot_kwargs={'cmap':'Greens`, 'vmin':0.2, 'vmax':0.9}`. + (some parameters like 'cmap' will only have an effect for + single-band animations, not three-band RGB animations). + colorbar_kwargs : dict, optional + An optional dict of keyword arguments used to control the + appearance of the colourbar. Keyword arguments are passed to + `matplotlib.pyplot.tick_params` + (see https://matplotlib.org/api/_as_gen/matplotlib.pyplot.tick_params.html + for options). This can be used to customise the colourbar + ticks, e.g. changing tick label colour depending on the + background of the animation: + `colorbar_kwargs={'colors': 'black'}`. + limit: int, optional + An optional integer specifying how many animation frames to + render (e.g. `limit=50` will render the first 50 frames). This + can be useful for quickly testing animations without rendering + the entire time-series. + """ + + def _start_end_times(gdf, ds): + """ + Converts 'start_time' and 'end_time' columns in a + `geopandas.GeoDataFrame` to datetime objects to allow vector + features to be plotted at specific moments in time during an + animation, and sets default values based on the first + and last time in `ds` if this information is missing from the + dataset. + """ + + # Make copy of gdf so we do not modify original data + gdf = gdf.copy() + + # Get min and max times from input dataset + minmax_times = pd.to_datetime(ds.time.isel(time=[0, -1]).values) + + # Update both `start_time` and `end_time` columns + for time_col, time_val in zip(["start_time", "end_time"], minmax_times): + + # Add time_col if it does not exist + if time_col not in gdf: + gdf[time_col] = np.nan + + # Convert values to datetimes and fill gaps with relevant time value + gdf[time_col] = pd.to_datetime(gdf[time_col], errors="ignore") + gdf[time_col] = gdf[time_col].fillna(time_val) + + return gdf + + def _add_colorbar(fig, ax, vmin, vmax, imshow_defaults, colorbar_defaults): + """ + Adds a new colorbar axis to the animation with custom minimum + and maximum values and styling. + """ + + # Create new axis object for colorbar + cax = fig.add_axes([0.02, 0.02, 0.96, 0.03]) + + # Initialise color bar using plot min and max values + img = ax.imshow(np.array([[vmin, vmax]]), **imshow_defaults) + fig.colorbar( + img, cax=cax, orientation="horizontal", ticks=np.linspace(vmin, vmax, 2) + ) + + # Fine-tune appearance of colorbar + cax.xaxis.set_ticks_position("top") + cax.tick_params(axis="x", **colorbar_defaults) + cax.get_xticklabels()[0].set_horizontalalignment("left") + cax.get_xticklabels()[-1].set_horizontalalignment("right") + + def _frame_annotation(times, show_date, show_text): + """ + Creates a custom annotation for the top-right of the animation + by converting a `xarray.DataArray` of times into strings, and + combining this with a custom text annotation. Handles cases + where `show_date=False/None`, `show_text=False/None`, or where + `show_text` is a list of strings. + """ + + # Test if show_text is supplied as a list + is_sequence = isinstance(show_text, (list, tuple, np.ndarray)) + + # Raise exception if it is shorter than number of dates + if is_sequence and (len(show_text) == 1): + show_text, is_sequence = show_text[0], False + elif is_sequence and (len(show_text) < len(times)): + raise ValueError( + f"Annotations supplied via `show_text` must have " + f"either a length of 1, or a length >= the number " + f"of timesteps in `ds` (n={len(times)})" + ) + + times_list = ( + times.dt.strftime(show_date).values if show_date else [None] * len(times) + ) + text_list = show_text if is_sequence else [show_text] * len(times) + annotation_list = [ + "\n".join([str(i) for i in (a, b) if i]) + for a, b in zip(times_list, text_list) + ] + + return annotation_list + + def _update_frames( + i, + ax, + extent, + annotation_text, + gdf, + gdf_defaults, + annotation_defaults, + imshow_defaults, + ): + """ + Animation called by `matplotlib.animation.FuncAnimation` to + animate each frame in the animation. Plots array and any text + annotations, as well as a temporal subset of `gdf` data based + on the times specified in 'start_time' and 'end_time' columns. + """ + + # Clear previous frame to optimise render speed and plot imagery + ax.clear() + ax.imshow( + array[i, ...].clip(0.0, 1.0), + extent=extent, + vmin=0.0, + vmax=1.0, + **imshow_defaults, + ) + + # Add annotation text + ax.annotate(annotation_text[i], **annotation_defaults) + + # Add geodataframe annotation + if show_gdf is not None: + + # Obtain start and end times to filter geodataframe features + time_i = ds.time.isel(time=i).values + + # Subset geodataframe using start and end dates + gdf_subset = show_gdf.loc[ + (show_gdf.start_time <= time_i) & (show_gdf.end_time >= time_i) + ] + + if len(gdf_subset.index) > 0: + + # Set color to geodataframe field if supplied + if ("color" in gdf_subset) and ("color" not in gdf_kwargs): + gdf_defaults.update({"color": gdf_subset["color"].tolist()}) + + gdf_subset.plot(ax=ax, **gdf_defaults) + + # Remove axes to show imagery only + ax.axis("off") + + # Update progress bar + progress_bar.update(1) + + # Test if bands have been supplied, or convert to list to allow + # iteration if a single band is provided as a string + if bands is None: + raise ValueError( + f"Please use the `bands` parameter to supply " + f"a list of one or three bands that exist as " + f"variables in `ds`, e.g. {list(ds.data_vars)}" + ) + elif isinstance(bands, str): + bands = [bands] + + # Test if bands exist in dataset + missing_bands = [b for b in bands if b not in ds.data_vars] + if missing_bands: + raise ValueError( + f"Band(s) {missing_bands} do not exist as " + f"variables in `ds` {list(ds.data_vars)}" + ) + + # Test if time dimension exists in dataset + if "time" not in ds.dims: + raise ValueError( + f"`ds` does not contain a 'time' dimension " + f"required for generating an animation" + ) + + # Set default parameters + outline = [PathEffects.withStroke(linewidth=2.5, foreground="black")] + annotation_defaults = { + "xy": (1, 1), + "xycoords": "axes fraction", + "xytext": (-5, -5), + "textcoords": "offset points", + "horizontalalignment": "right", + "verticalalignment": "top", + "fontsize": 20, + "color": "white", + "path_effects": outline, + } + imshow_defaults = {"cmap": "magma", "interpolation": "nearest"} + colorbar_defaults = {"colors": "white", "labelsize": 12, "length": 0} + gdf_defaults = {"linewidth": 1.5} + + # Update defaults with kwargs + annotation_defaults.update(annotation_kwargs) + imshow_defaults.update(imshow_kwargs) + colorbar_defaults.update(colorbar_kwargs) + gdf_defaults.update(gdf_kwargs) + + # Get info on dataset dimensions + height, width = ds.geobox.shape + scale = width_pixels / width + left, bottom, right, top = ds.geobox.extent.boundingbox + + # Prepare annotations + annotation_list = _frame_annotation(ds.time, show_date, show_text) + + # Prepare geodataframe + if show_gdf is not None: + show_gdf = show_gdf.to_crs(ds.geobox.crs) + show_gdf = gpd.clip(show_gdf, mask=box(left, bottom, right, top)) + show_gdf = _start_end_times(show_gdf, ds) + + # Convert data to 4D numpy array of shape [time, y, x, bands] + ds = ds[bands].to_array().transpose(..., "variable")[0:limit, ...] + array = ds.astype(np.float32).values + + # Optionally apply image processing along axis 0 (e.g. to each timestep) + bar_format = ( + "{l_bar}{bar}| {n_fmt}/{total_fmt} ({remaining_s:.1f} " + "seconds remaining at {rate_fmt}{postfix})" + ) + if image_proc_funcs: + print("Applying custom image processing functions") + for i, array_i in tqdm( + enumerate(array), + total=len(ds.time), + leave=False, + bar_format=bar_format, + unit=" frames", + ): + for func in image_proc_funcs: + array_i = func(array_i) + array[i, ...] = array_i + + # Clip to percentiles and rescale between 0.0 and 1.0 for plotting + vmin, vmax = np.quantile(array[np.isfinite(array)], q=percentile_stretch) + + # Replace with vmin and vmax if present in `imshow_defaults` + if "vmin" in imshow_defaults: + vmin = imshow_defaults.pop("vmin") + if "vmax" in imshow_defaults: + vmax = imshow_defaults.pop("vmax") + + # Rescale between 0 and 1 + array = rescale_intensity(array, in_range=(vmin, vmax), out_range=(0.0, 1.0)) + array = np.squeeze(array) # remove final axis if only one band + + # Set up figure + fig, ax = plt.subplots() + fig.set_size_inches(width * scale / 72, height * scale / 72, forward=True) + fig.subplots_adjust(left=0, bottom=0, right=1, top=1, wspace=0, hspace=0) + + # Optionally add colorbar + if show_colorbar & (len(bands) == 1): + _add_colorbar(fig, ax, vmin, vmax, imshow_defaults, colorbar_defaults) + + # Animate + print(f"Exporting animation to {output_path}") + anim = FuncAnimation( + fig=fig, + func=_update_frames, + fargs=( + ax, # axis to plot into + [left, right, bottom, top], # imshow extent + annotation_list, # list of text annotations + show_gdf, # geodataframe to plot over imagery + gdf_defaults, # any kwargs used to plot gdf + annotation_defaults, # kwargs for annotations + imshow_defaults, + ), # kwargs for imshow + frames=len(ds.time), + interval=interval, + repeat=False, + ) + + # Set up progress bar + progress_bar = tqdm(total=len(ds.time), unit=" frames", bar_format=bar_format) + + # Export animation to file + if Path(output_path).suffix == ".gif": + anim.save(output_path, writer="pillow") + else: + anim.save(output_path, dpi=72) + + # Update progress bar to fix progress bar moving past end + if progress_bar.n != len(ds.time): + progress_bar.n = len(ds.time) + progress_bar.last_print_n = len(ds.time) + + +def _degree_to_zoom_level(l1, l2, margin=0.0): + """ + Helper function to set zoom level for `display_map` + """ + + degree = abs(l1 - l2) * (1 + margin) + zoom_level_int = 0 + if degree != 0: + zoom_level_float = math.log(360 / degree) / math.log(2) + zoom_level_int = int(zoom_level_float) + else: + zoom_level_int = 18 + return zoom_level_int + + +def plot_wofs(wofs, legend=True, **plot_kwargs): + """Plot a water observation bit flag image. + + Parameters + ---------- + wofs : xr.DataArray + A DataArray containing water observation bit flags. + legend : bool + Whether to plot a legend. Default True. + plot_kwargs : dict + Keyword arguments passed on to DataArray.plot. + + Returns + ------- + plot + """ + cmap = mcolours.ListedColormap( + [ + np.array([150, 150, 110]) / 255, # dry - 0 + np.array([0, 0, 0]) / 255, # nodata, - 1 + np.array([119, 104, 87]) / 255, # terrain - 16 + np.array([89, 88, 86]) / 255, # cloud_shadow - 32 + np.array([216, 215, 214]) / 255, # cloud - 64 + np.array([242, 220, 180]) / 255, # cloudy terrain - 80 + np.array([79, 129, 189]) / 255, # water - 128 + np.array([51, 82, 119]) / 255, # shady water - 160 + np.array([186, 211, 242]) / 255, # cloudy water - 192 + ] + ) + bounds = [ + 0, + 1, + 16, + 32, + 64, + 80, + 128, + 160, + 192, + ] + norm = mcolours.BoundaryNorm(np.array(bounds) - 0.1, cmap.N) + cblabels = [ + "dry", + "nodata", + "terrain", + "cloud shadow", + "cloud", + "cloudy terrain", + "water", + "shady water", + "cloudy water", + ] + + try: + im = wofs.plot.imshow(cmap=cmap, norm=norm, add_colorbar=legend, **plot_kwargs) + except AttributeError: + im = wofs.plot(cmap=cmap, norm=norm, add_colorbar=legend, **plot_kwargs) + + if legend: + try: + cb = im.colorbar + except AttributeError: + cb = im.cbar + ticks = cb.get_ticks() + cb.set_ticks(ticks + np.diff(ticks, append=193) / 2) + cb.set_ticklabels(cblabels) + return im + + +def plot_lulc(lulc, product=None, legend=True, **plot_kwargs): + """Plot a LULC image. + + Parameters + ---------- + lulc : xr.DataArray + A DataArray containing LULC bit flags. + product : str + 'ESA', 'IO', 'CGLS', or 'CCI', 'ESRI' + legend : bool + Whether to plot a legend. Default True. + plot_kwargs : dict + Keyword arguments passed on to DataArray.plot. + + Returns + ------- + plot + """ + + if "ESRI" in product: + # this is for the orignal ESRI/IO 10 class product for 2020 + try: + cmap = mcolours.ListedColormap( + [ + np.array([0, 0, 0]) / 255, + np.array([65, 155, 223]) / 255, + np.array([57, 125, 73]) / 255, + np.array([136, 176, 83]) / 255, + np.array([122, 135, 198]) / 255, + np.array([228, 150, 53]) / 255, + np.array([223, 195, 90]) / 255, + np.array([196, 40, 27]) / 255, + np.array([165, 155, 143]) / 255, + np.array([168, 235, 255]) / 255, + np.array([97, 97, 97]) / 255, + ] + ) + bounds = range(0, 12) + norm = mcolours.BoundaryNorm(np.array(bounds), cmap.N) + cblabels = [ + "no data", + "water", + "trees", + "grass", + "flooded vegetation", + "crops", + "scrub/shrub", + "built area", + "bare ground", + "snow/ice", + "clouds", + ] + except: + AttributeError + + if "IO" in product: + # this is for the ESRI/IO 9 class multiyear product; same color as preview + try: + cmap = mcolours.ListedColormap( + [ + np.array([0, 0, 0]) / 255, + np.array([65, 155, 223]) / 255, + np.array([57, 125, 73]) / 255, + np.array([122, 135, 198]) / 255, + np.array([228, 150, 53]) / 255, + np.array([196, 40, 27]) / 255, + np.array([165, 155, 143]) / 255, + np.array([168, 235, 255]) / 255, + np.array([97, 97, 97]) / 255, + np.array([227, 226, 195]) / 255, + ] + ) + bounds = [-0.5, 0.5, 1.5, 3, 4.5, 6, 7.5, 8.5, 9.5, 10.5, 11.5] + norm = mcolours.BoundaryNorm(np.array(bounds), cmap.N) + cblabels = [ + "no data", + "water", + "trees", + "flooded vegetation", + "crops", + "built area", + "bare ground", + "snow/ice", + "clouds", + "rangeland", + ] + ticks = list(np.mean((bounds[i+1], val)) for i, val in enumerate(bounds[:-1])) + except: + AttributeError + + if "ESA" in product: + try: + cmap = mcolours.ListedColormap( + [ + np.array([0, 0, 0]) / 255, + np.array([0, 100, 0]) / 255, + np.array([255, 187, 34]) / 255, + np.array([255, 255, 76]) / 255, + np.array([240, 150, 255]) / 255, + np.array([250, 0, 0]) / 255, + np.array([180, 180, 180]) / 255, + np.array([240, 240, 240]) / 255, + np.array([0, 100, 200]) / 255, + np.array([0, 150, 160]) / 255, + np.array([0, 207, 117]) / 255, + np.array([250, 230, 160]) / 255, + ] + ) + bounds = [-5, 5, 15, 25, 35, 45, 55, 65, 75, 85, 92, 98, 105] + norm = mcolours.BoundaryNorm(np.array(bounds), cmap.N) + cblabels = [ + "no data", + "tree cover", + "shrubland", + "grassland", + "cropland", + "built up", + "bare/sparse vegetation", + "snow and ice", + "permanent water bodies", + "herbaceous wetland", + "mangroves", + "moss and lichen", + ] + except: + AttributeError + + if "CGLS" in product: + try: + labels = {0: {'color': '#282828', 'flag': 'unknown'}, + 20: {'color': '#FFBB22', 'flag': 'shrubs'}, + 30: {'color': '#FFFF4C', 'flag': 'herbaceous_vegetation'}, + 40: {'color': '#F096FF', 'flag': 'cultivated_and_managed_vegetation_or_agriculture'}, + 50: {'color': '#FA0000', 'flag': 'urban_or_built_up'}, + 60: {'color': '#B4B4B4', 'flag': 'bare_or_sparse_vegetation'}, + 70: {'color': '#F0F0F0', 'flag': 'snow_and_ice'}, + 80: {'color': '#0032C8', 'flag': 'permanent_water_bodies'}, + 90: {'color': '#0096A0', 'flag': 'herbaceous_wetland'}, + 100: {'color': '#FAE6A0', 'flag': 'moss_and_lichen'}, + 111: {'color': '#58481F', 'flag': 'closed_forest_evergreen_needle_leaf'}, + 112: {'color': '#009900', 'flag': 'closed_forest_evergreen_broad_leaf'}, + 113: {'color': '#70663E', 'flag': 'closed_forest_deciduous_needle_leaf'}, + 114: {'color': '#00CC00', 'flag': 'closed_forest_deciduous_broad_leaf'}, + 115: {'color': '#4E751F', 'flag': 'closed_forest_mixed'}, + 116: {'color': '#007800', 'flag': 'closed_forest_not_matching_any_of_the_other_definitions'}, + 121: {'color': '#666000', 'flag': 'open_forest_evergreen_needle_leaf'}, + 122: {'color': '#8DB400', 'flag': 'open_forest_evergreen_broad_leaf'}, + 123: {'color': '#8D7400', 'flag': 'open_forest_deciduous_needle_leaf'}, + 124: {'color': '#A0DC00', 'flag': 'open_forest_deciduous_broad_leaf'}, + 125: {'color': '#929900', 'flag': 'open_forest_mixed'}, + 126: {'color': '#648C00', 'flag': 'open_forest_not_matching_any_of_the_other_definitions'}, + 200: {'color': '#000080', 'flag': 'oceans_seas'}} + + colors = [label['color'] for label in labels.values()] + cmap = ListedColormap([label['color'] for label in labels.values()]) + norm = mcolours.BoundaryNorm(list(labels.keys())+[201], cmap.N+1, extend='max') + ticks = list(np.mean((list(list(labels.keys())+[201])[i+1], val)) for i, val in enumerate(list(labels.keys()))) + cblabels=[label['flag'] for label in labels.values()] + + except: + AttributeError + if 'CCI' in product: + try: + labels = {0: {'color': '#282828', 'flag': 'no data'}, + 10: {'color': '#EBEB34', 'flag': 'cropland, rainfed'}, + 11: {'color': '#D9EB34', 'flag': 'cropland, rainfed, herbaceous cover'}, + 12: {'color': '#EBDF34', 'flag': 'cropland, rainfed, tree or shrub cover'}, + 20: {'color': '#34EBE2', 'flag': 'cropland, irrigated or post-flooding'}, + 30: {'color': '#EBBD34', 'flag': 'mosaic cropland/natural vegetation'}, + 40: {'color': '#eba534', 'flag': 'mosaic natural vegetation/cropland'}, + 50: {'color': '#34eb46', 'flag': 'tree cover, broadleaved, evergreen, closed to open'}, + 60: {'color': '#21750e', 'flag': 'tree cover, broadleaved, deciduous, closed to open'}, + 61: {'color': '#449432', 'flag': 'tree cover, broadleaved, deciduous, closed'}, + 62: {'color': '#5da64c', 'flag': 'tree cover, broadleaved, deciduous, open'}, + 70: {'color': '#16470b', 'flag': 'tree cover, needleleaved, evergreen, closed to open'}, + 71: {'color': '#237012', 'flag': 'tree cover, needleleaved, evergreen, closed'}, + 72: {'color': '#237012', 'flag': 'tree cover, needleleaved, evergreen, open'}, + 80: {'color': '#31a317', 'flag': 'tree cover, needleleaved, deciduous, closed to open'}, + 81: {'color': '#57ed34', 'flag': 'tree cover, needleleaved, deciduous, closed'}, + 82: {'color': '#81f765', 'flag': 'tree cover, needleleaved, deciduous, open'}, + 90: {'color': '#b6ed64', 'flag': 'tree cover, mixed leaf type'}, + 100: {'color': '#6f8f3f', 'flag': 'mosaic tree and shrub/herbaceous cover'}, + 110: {'color': '#ad950c', 'flag': 'mosaic herbaceous cover/tree and shrub'}, + 120: {'color': '#5e5209', 'flag': 'shrubland'}, + 121: {'color': '#292302', 'flag': 'shrubland, evergreen'}, + 122: {'color': '#a89008', 'flag': 'shrubland, deciduous'}, + 130: {'color': '#f7bf07', 'flag': 'grassland'}, + 140: {'color': '#f57feb', 'flag': 'lichens and mosses'}, + 150: {'color': '#f57feb', 'flag': 'sparse vegetation'}, + 151: {'color': '#fcf7a4', 'flag': 'sparse tree'}, + 152: {'color': '#d4cf87', 'flag': 'sparse shrub'}, + 153: {'color': '#b0aa54', 'flag': 'sparse herbaceous cover'}, + 160: {'color': '#159638', 'flag': 'tree cover, flooded, fresh or brakish water'}, + 170: {'color': '#22bf81', 'flag': 'tree cover, flooded, saline water'}, + 180: {'color': '#44eba9', 'flag': 'shrub or herbaceous cover, flooded, fresh/saline/brakish water'}, + 190: {'color': '#a3273c', 'flag': 'urban areas'}, + 200: {'color': '#fffbcc', 'flag': 'bare areas'}, + 201: {'color': '#b0afa4', 'flag': 'consolidated bare areas'}, + 202: {'color': '#d6d4b6', 'flag': 'unconsolidated bare areas'}, + 210: {'color': '#1A3EF0', 'flag': 'water bodies'}, + 220: {'color': '#ffffff', 'flag': 'permanent snow and ice'}} + + colors = [label['color'] for label in labels.values()] + cmap = ListedColormap([label['color'] for label in labels.values()]) + norm = mcolours.BoundaryNorm(list(labels.keys())+[221], cmap.N+1, extend='max') + ticks = list(np.mean((list(list(labels.keys())+[221])[i+1], val)) for i, val in enumerate(list(labels.keys()))) + cblabels=[label['flag'] for label in labels.values()] + + except: + AttributeError + + try: + im = lulc.plot.imshow(cmap=cmap, norm=norm, add_colorbar=legend, **plot_kwargs) + except AttributeError: + im = lulc.plot(cmap=cmap, norm=norm, add_colorbar=legend, **plot_kwargs) + + if legend: + try: + cb = im.colorbar + except AttributeError: + cb = im.cbar + + if "ESRI" in product: + cb.set_ticks(np.arange(0, 11, 1)+0.5) + cb.set_ticklabels(cblabels) + + if "IO" in product: + cb.set_ticks(ticks) + cb.set_ticklabels(cblabels) + + if "ESA" in product: + cb.set_ticks([0, 10, 20, 30, 40, 50, 60, 70, 80, 88.5, 95, 101.5]) + cb.set_ticklabels(cblabels) + + if "CGLS" in product: + cb.set_ticks(ticks) + cb.set_ticklabels(cblabels) + + if "CCI" in product: + cb.set_ticks(ticks) + cb.set_ticklabels(cblabels) + + return im diff --git a/deafrica_tools/spatial.py b/deafrica_tools/spatial.py new file mode 100644 index 0000000..de454f2 --- /dev/null +++ b/deafrica_tools/spatial.py @@ -0,0 +1,949 @@ +''' +Spatial analyses functions for Digital Earth Africa data. +''' + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os + +os.environ['USE_PYGEOS'] = '0' + +import multiprocessing as mp + +import dask +import fiona +import geopandas as gpd +import numpy as np +import odc.geo.xr # adds `.odc.x` attributes to our xarray objects. +import pandas as pd +import rasterio.features +import scipy.interpolate +import xarray as xr +from datacube.api.query import query_group_by +from datacube.model.utils import xr_apply +from datacube.utils.cog import write_cog +from datacube.utils.geometry import CRS, Geometry +from geopy.geocoders import Nominatim +from rasterstats import zonal_stats +from shapely.geometry import LineString, MultiLineString, mapping, shape +from skimage.measure import find_contours, label + + +def add_geobox(ds, crs=None): + """ + Ensure that an xarray DataArray has a GeoBox and .odc.* accessor + using `odc.geo`. + + If `ds` is missing a Coordinate Reference System (CRS), this can be + supplied using the `crs` param. + + Parameters + ---------- + ds : xarray.Dataset or xarray.DataArray + Input xarray object that needs to be checked for spatial + information. + crs : str, optional + Coordinate Reference System (CRS) information for the input `ds` + array. If `ds` already has a CRS, then `crs` is not required. + Default is None. + + Returns + ------- + xarray.Dataset or xarray.DataArray + The input xarray object with added `.odc.x` attributes to access + spatial information. + + """ + # If a CRS is not found, use custom provided CRS + if ds.odc.crs is None and crs is not None: + ds = ds.odc.assign_crs(crs) + elif ds.odc.crs is None and crs is None: + raise ValueError( + "Unable to determine `ds`'s coordinate " + "reference system (CRS). Please provide a " + "CRS using the `crs` parameter " + "(e.g. `crs='EPSG:3577'`)." + ) + + return ds + + +def xr_vectorize( + da, + attribute_col=None, + crs=None, + dtype="float32", + output_path=None, + verbose=True, + **rasterio_kwargs, +): + """ + Vectorises a raster ``xarray.DataArray`` into a vector + ``geopandas.GeoDataFrame``. + + Parameters + ---------- + da : xarray.DataArray + The input ``xarray.DataArray`` data to vectorise. + attribute_col : str, optional + Name of the attribute column in the resulting + ``geopandas.GeoDataFrame``. Values from ``da`` converted + to polygons will be assigned to this column. If None, + the column name will default to 'attribute'. + crs : str or CRS object, optional + If ``da``'s coordinate reference system (CRS) cannot be + determined, provide a CRS using this parameter. + (e.g. 'EPSG:3577'). + dtype : str, optional + Data type of must be one of int16, int32, uint8, uint16, + or float32 + output_path : string, optional + Provide an optional string file path to export the vectorised + data to file. Supports any vector file formats supported by + ``geopandas.GeoDataFrame.to_file()``. + verbose : bool, optional + Print debugging messages. Default True. + **rasterio_kwargs : + A set of keyword arguments to ``rasterio.features.shapes``. + Can include `mask` and `connectivity`. + + Returns + ------- + gdf : geopandas.GeoDataFrame + + """ + + # Add GeoBox and odc.* accessor to array using `odc-geo` + da = add_geobox(da, crs) + + # Run the vectorizing function + vectors = rasterio.features.shapes( + source=da.data.astype(dtype), transform=da.odc.transform, **rasterio_kwargs + ) + + # Convert the generator into a list + vectors = list(vectors) + + # Extract the polygon coordinates and values from the list + polygons = [polygon for polygon, value in vectors] + values = [value for polygon, value in vectors] + + # Convert polygon coordinates into polygon shapes + polygons = [shape(polygon) for polygon in polygons] + + # Create a geopandas dataframe populated with the polygon shapes + attribute_name = attribute_col if attribute_col is not None else "attribute" + gdf = gpd.GeoDataFrame( + data={attribute_name: values}, geometry=polygons, crs=da.odc.crs + ) + + # If a file path is supplied, export to file + if output_path is not None: + if verbose: + print(f"Exporting vector data to {output_path}") + gdf.to_file(output_path) + + return gdf + + +def xr_rasterize( + gdf, + da, + attribute_col=None, + crs=None, + name=None, + output_path=None, + verbose=True, + **rasterio_kwargs, +): + """ + Rasterizes a vector ``geopandas.GeoDataFrame`` into a + raster ``xarray.DataArray``. + + Parameters + ---------- + gdf : geopandas.GeoDataFrame + A ``geopandas.GeoDataFrame`` object containing the vector + data you want to rasterise. + da : xarray.DataArray or xarray.Dataset + The shape, coordinates, dimensions, and transform of this object + are used to define the array that ``gdf`` is rasterized into. + It effectively provides a spatial template. + attribute_col : string, optional + Name of the attribute column in ``gdf`` containing values for + each vector feature that will be rasterized. If None, the + output will be a boolean array of 1's and 0's. + crs : str or CRS object, optional + If ``da``'s coordinate reference system (CRS) cannot be + determined, provide a CRS using this parameter. + (e.g. 'EPSG:3577'). + name : str, optional + An optional name used for the output ``xarray.DataArray`. + output_path : string, optional + Provide an optional string file path to export the rasterized + data as a GeoTIFF file. + verbose : bool, optional + Print debugging messages. Default True. + **rasterio_kwargs : + A set of keyword arguments to ``rasterio.features.rasterize``. + Can include: 'all_touched', 'merge_alg', 'dtype'. + + Returns + ------- + da_rasterized : xarray.DataArray + The rasterized vector data. + """ + + # Add GeoBox and odc.* accessor to array using `odc-geo` + da = add_geobox(da, crs) + + # Reproject vector data to raster's CRS + gdf_reproj = gdf.to_crs(crs=da.odc.crs) + + # If an attribute column is specified, rasterise using vector + # attribute values. Otherwise, rasterise into a boolean array + if attribute_col is not None: + # Use the geometry and attributes from `gdf` to create an iterable + shapes = zip(gdf_reproj.geometry, gdf_reproj[attribute_col]) + else: + # Use geometry directly (will produce a boolean numpy array) + shapes = gdf_reproj.geometry + + # Rasterise shapes into a numpy array + im = rasterio.features.rasterize( + shapes=shapes, + out_shape=da.odc.geobox.shape, + transform=da.odc.geobox.transform, + **rasterio_kwargs, + ) + + # Convert numpy array to a full xarray.DataArray + # and set array name if supplied + da_rasterized = odc.geo.xr.wrap_xr(im=im, gbox=da.odc.geobox) + da_rasterized = da_rasterized.rename(name) + + # If a file path is supplied, export to file + if output_path is not None: + if verbose: + print(f"Exporting raster data to {output_path}") + write_cog(da_rasterized, output_path, overwrite=True) + + return da_rasterized + + +def subpixel_contours( + da, + z_values=[0.0], + crs=None, + attribute_df=None, + output_path=None, + min_vertices=2, + dim="time", + time_format="%Y-%m-%d", + errors="ignore", + verbose=True, +): + """ + Uses `skimage.measure.find_contours` to extract multiple z-value + contour lines from a two-dimensional array (e.g. multiple elevations + from a single DEM), or one z-value for each array along a specified + dimension of a multi-dimensional array (e.g. to map waterlines + across time by extracting a 0 NDWI contour from each individual + timestep in an xarray timeseries). + + Contours are returned as a geopandas.GeoDataFrame with one row per + z-value or one row per array along a specified dimension. The + `attribute_df` parameter can be used to pass custom attributes + to the output contour features. + + Last modified: May 2023 + + Parameters + ---------- + da : xarray DataArray + A two-dimensional or multi-dimensional array from which + contours are extracted. If a two-dimensional array is provided, + the analysis will run in 'single array, multiple z-values' mode + which allows you to specify multiple `z_values` to be extracted. + If a multi-dimensional array is provided, the analysis will run + in 'single z-value, multiple arrays' mode allowing you to + extract contours for each array along the dimension specified + by the `dim` parameter. + z_values : int, float or list of ints, floats + An individual z-value or list of multiple z-values to extract + from the array. If operating in 'single z-value, multiple + arrays' mode specify only a single z-value. + crs : string or CRS object, optional + If ``da``'s coordinate reference system (CRS) cannot be + determined, provide a CRS using this parameter. + (e.g. 'EPSG:3577'). + output_path : string, optional + The path and filename for the output shapefile. + attribute_df : pandas.Dataframe, optional + A pandas.Dataframe containing attributes to pass to the output + contour features. The dataframe must contain either the same + number of rows as supplied `z_values` (in 'multiple z-value, + single array' mode), or the same number of rows as the number + of arrays along the `dim` dimension ('single z-value, multiple + arrays mode'). + min_vertices : int, optional + The minimum number of vertices required for a contour to be + extracted. The default (and minimum) value is 2, which is the + smallest number required to produce a contour line (i.e. a start + and end point). Higher values remove smaller contours, + potentially removing noise from the output dataset. + dim : string, optional + The name of the dimension along which to extract contours when + operating in 'single z-value, multiple arrays' mode. The default + is 'time', which extracts contours for each array along the time + dimension. + time_format : string, optional + The format used to convert `numpy.datetime64` values to strings + if applied to data with a "time" dimension. Defaults to + "%Y-%m-%d". + errors : string, optional + If 'raise', then any failed contours will raise an exception. + If 'ignore' (the default), a list of failed contours will be + printed. If no contours are returned, an exception will always + be raised. + verbose : bool, optional + Print debugging messages. Default is True. + + Returns + ------- + output_gdf : geopandas geodataframe + A geopandas geodataframe object with one feature per z-value + ('single array, multiple z-values' mode), or one row per array + along the dimension specified by the `dim` parameter ('single + z-value, multiple arrays' mode). If `attribute_df` was + provided, these values will be included in the shapefile's + attribute table. + """ + + def _contours_to_multiline(da_i, z_value, min_vertices=2): + """ + Helper function to apply marching squares contour extraction + to an array and return a data as a shapely MultiLineString. + The `min_vertices` parameter allows you to drop small contours + with less than X vertices. + """ + + # Extracts contours from array, and converts each discrete + # contour into a Shapely LineString feature. If the function + # returns a KeyError, this may be due to an unresolved issue in + # scikit-image: https://github.com/scikit-image/scikit-image/issues/4830 + # A temporary workaround is to peturb the z-value by a tiny + # amount (1e-12) before using it to extract the contour. + try: + line_features = [ + LineString(i[:, [1, 0]]) + for i in find_contours(da_i.data, z_value) + if i.shape[0] >= min_vertices + ] + except KeyError: + line_features = [ + LineString(i[:, [1, 0]]) + for i in find_contours(da_i.data, z_value + 1e-12) + if i.shape[0] >= min_vertices + ] + + # Output resulting lines into a single combined MultiLineString + return MultiLineString(line_features) + + def _time_format(i, time_format): + """ + Converts numpy.datetime64 into formatted strings; + otherwise returns data as-is. + """ + if isinstance(i, np.datetime64): + ts = pd.to_datetime(str(i)) + i = ts.strftime(time_format) + return i + + # Verify input data is a xr.DataArray + if not isinstance(da, xr.DataArray): + raise ValueError( + "The input `da` is not an xarray.DataArray. " + "If you supplied an xarray.Dataset, pass in one " + "of its data variables using the syntax " + "`da=ds.`." + ) + + # Add GeoBox and odc.* accessor to array using `odc-geo` + da = add_geobox(da, crs) + + # If z_values is supplied is not a list, convert to list: + z_values = ( + z_values + if (isinstance(z_values, list) or isinstance(z_values, np.ndarray)) + else [z_values] + ) + + # If dask collection, load into memory + if dask.is_dask_collection(da): + if verbose: + print("Loading data into memory using Dask") + da = da.compute() + + # Test number of dimensions in supplied data array + if len(da.shape) == 2: + if verbose: + print("Operating in multiple z-value, single array mode") + dim = "z_value" + contour_arrays = { + _time_format(i, time_format): _contours_to_multiline(da, i, min_vertices) + for i in z_values + } + + else: + # Test if only a single z-value is given when operating in + # single z-value, multiple arrays mode + if verbose: + print("Operating in single z-value, multiple arrays mode") + if len(z_values) > 1: + raise ValueError( + "Please provide a single z-value when operating " + "in single z-value, multiple arrays mode" + ) + + contour_arrays = { + _time_format(i, time_format): _contours_to_multiline( + da_i, z_values[0], min_vertices + ) + for i, da_i in da.groupby(dim) + } + + # If attributes are provided, add the contour keys to that dataframe + if attribute_df is not None: + try: + attribute_df.insert(0, dim, contour_arrays.keys()) + + # If this fails, it is due to the applied attribute table not + # matching the structure of the loaded data + except ValueError: + if len(da.shape) == 2: + raise ValueError( + f"The provided `attribute_df` contains a different " + f"number of rows ({len(attribute_df.index)}) " + f"than the number of supplied `z_values` " + f"({len(z_values)})." + ) + else: + raise ValueError( + f"The provided `attribute_df` contains a different " + f"number of rows ({len(attribute_df.index)}) " + f"than the number of arrays along the '{dim}' " + f"dimension ({len(da[dim])})." + ) + + # Otherwise, use the contour keys as the only main attributes + else: + attribute_df = list(contour_arrays.keys()) + + # Convert output contours to a geopandas.GeoDataFrame + contours_gdf = gpd.GeoDataFrame( + data=attribute_df, geometry=list(contour_arrays.values()), crs=da.odc.crs + ) + + # Define affine and use to convert array coords to geographic coords. + # We need to add 0.5 x pixel size to the x and y to obtain the centre + # point of our pixels, rather than the top-left corner + affine = da.odc.geobox.transform + shapely_affine = [ + affine.a, + affine.b, + affine.d, + affine.e, + affine.xoff + affine.a / 2.0, + affine.yoff + affine.e / 2.0, + ] + contours_gdf["geometry"] = contours_gdf.affine_transform(shapely_affine) + + # Rename the data column to match the dimension + contours_gdf = contours_gdf.rename({0: dim}, axis=1) + + # Drop empty timesteps + empty_contours = contours_gdf.geometry.is_empty + failed = ", ".join(map(str, contours_gdf[empty_contours][dim].to_list())) + contours_gdf = contours_gdf[~empty_contours] + + # Raise exception if no data is returned, or if any contours fail + # when `errors='raise'. Otherwise, print failed contours + if empty_contours.all() and errors == "raise": + raise ValueError( + "Failed to generate any valid contours; verify that " + "values passed to `z_values` are valid and present " + "in `da`" + ) + elif empty_contours.all() and errors == "ignore": + if verbose: + print( + "Failed to generate any valid contours; verify that " + "values passed to `z_values` are valid and present " + "in `da`" + ) + elif empty_contours.any() and errors == "raise": + raise Exception(f"Failed to generate contours: {failed}") + elif empty_contours.any() and errors == "ignore": + if verbose: + print(f"Failed to generate contours: {failed}") + + # If asked to write out file, test if GeoJSON or ESRI Shapefile. If + # GeoJSON, convert to EPSG:4326 before exporting. + if output_path and output_path.endswith(".geojson"): + if verbose: + print(f"Writing contours to {output_path}") + contours_gdf.to_crs("EPSG:4326").to_file(filename=output_path) + + if output_path and output_path.endswith(".shp"): + if verbose: + print(f"Writing contours to {output_path}") + contours_gdf.to_file(filename=output_path) + + return contours_gdf + + +def interpolate_2d(ds, + x_coords, + y_coords, + z_coords, + method='linear', + factor=1, + verbose=False, + **kwargs): + + """ + This function takes points with X, Y and Z coordinates, and + interpolates Z-values across the extent of an existing xarray + dataset. This can be useful for producing smooth surfaces from point + data that can be compared directly against satellite data derived + from an OpenDataCube query. + + Supported interpolation methods include 'linear', 'nearest' and + 'cubic (using `scipy.interpolate.griddata`), and 'rbf' (using + `scipy.interpolate.Rbf`). + + Last modified: February 2020 + + Parameters + ---------- + ds : xarray DataArray or Dataset + A two-dimensional or multi-dimensional array from which x and y + dimensions will be copied and used for the area in which to + interpolate point data. + x_coords, y_coords : numpy array + Arrays containing X and Y coordinates for all points (e.g. + longitudes and latitudes). + z_coords : numpy array + An array containing Z coordinates for all points (e.g. + elevations). These are the values you wish to interpolate + between. + method : string, optional + The method used to interpolate between point values. This string + is either passed to `scipy.interpolate.griddata` (for 'linear', + 'nearest' and 'cubic' methods), or used to specify Radial Basis + Function interpolation using `scipy.interpolate.Rbf` ('rbf'). + Defaults to 'linear'. + factor : int, optional + An optional integer that can be used to subsample the spatial + interpolation extent to obtain faster interpolation times, then + up-sample this array back to the original dimensions of the + data as a final step. For example, setting `factor=10` will + interpolate data into a grid that has one tenth of the + resolution of `ds`. This approach will be significantly faster + than interpolating at full resolution, but will potentially + produce less accurate or reliable results. + verbose : bool, optional + Print debugging messages. Default False. + **kwargs : + Optional keyword arguments to pass to either + `scipy.interpolate.griddata` (if `method` is 'linear', 'nearest' + or 'cubic'), or `scipy.interpolate.Rbf` (is `method` is 'rbf'). + + Returns + ------- + interp_2d_array : xarray DataArray + An xarray DataArray containing with x and y coordinates copied + from `ds_array`, and Z-values interpolated from the points data. + """ + + # Extract xy and elev points + points_xy = np.vstack([x_coords, y_coords]).T + + # Extract x and y coordinates to interpolate into. + # If `factor` is greater than 1, the coordinates will be subsampled + # for faster run-times. If the last x or y value in the subsampled + # grid aren't the same as the last x or y values in the original + # full resolution grid, add the final full resolution grid value to + # ensure data is interpolated up to the very edge of the array + if ds.x[::factor][-1].item() == ds.x[-1].item(): + x_grid_coords = ds.x[::factor].values + else: + x_grid_coords = ds.x[::factor].values.tolist() + [ds.x[-1].item()] + + if ds.y[::factor][-1].item() == ds.y[-1].item(): + y_grid_coords = ds.y[::factor].values + else: + y_grid_coords = ds.y[::factor].values.tolist() + [ds.y[-1].item()] + + # Create grid to interpolate into + grid_y, grid_x = np.meshgrid(x_grid_coords, y_grid_coords) + + # Apply scipy.interpolate.griddata interpolation methods + if method in ('linear', 'nearest', 'cubic'): + + # Interpolate x, y and z values + interp_2d = scipy.interpolate.griddata(points=points_xy, + values=z_coords, + xi=(grid_y, grid_x), + method=method, + **kwargs) + + # Apply Radial Basis Function interpolation + elif method == 'rbf': + + # Interpolate x, y and z values + rbf = scipy.interpolate.Rbf(x_coords, y_coords, z_coords, **kwargs) + interp_2d = rbf(grid_y, grid_x) + + # Create xarray dataarray from the data and resample to ds coords + interp_2d_da = xr.DataArray(interp_2d, + coords=[y_grid_coords, x_grid_coords], + dims=['y', 'x']) + + # If factor is greater than 1, resample the interpolated array to + # match the input `ds` array + if factor > 1: + interp_2d_da = interp_2d_da.interp_like(ds) + + return interp_2d_da + + +def contours_to_arrays(gdf, col): + """ + This function converts a polyline shapefile into an array with three + columns giving the X, Y and Z coordinates of each vertex. This data + can then be used as an input to interpolation procedures (e.g. using + a function like `interpolate_2d`. + + Last modified: October 2021 + + Parameters + ---------- + gdf : Geopandas GeoDataFrame + A GeoPandas GeoDataFrame of lines to convert into point + coordinates. + col : str + A string giving the name of the GeoDataFrame field to use as + Z-values. + + Returns + ------- + A numpy array with three columns giving the X, Y and Z coordinates + of each vertex in the input GeoDataFrame. + + """ + + # Explode multi-part geometries into multiple single geometries. + gdf = gdf.explode(ignore_index=True) + + coords_zvals = [] + + for i in range(0, len(gdf)): + val = gdf.iloc[i][col] + + try: + coords = np.concatenate( + [np.vstack(x.coords.xy).T for x in gdf.iloc[i].geometry.geoms] + ) + except Exception: + coords = np.vstack(gdf.iloc[i].geometry.coords.xy).T + + coords_zvals.append( + np.column_stack((coords, np.full(np.shape(coords)[0], fill_value=val))) + ) + + return np.concatenate(coords_zvals) + + +def largest_region(bool_array, **kwargs): + + ''' + Takes a boolean array and identifies the largest contiguous region of + connected True values. This is returned as a new array with cells in + the largest region marked as True, and all other cells marked as False. + + Parameters + ---------- + bool_array : boolean array + A boolean array (numpy or xarray.DataArray) with True values for + the areas that will be inspected to find the largest group of + connected cells + **kwargs : + Optional keyword arguments to pass to `measure.label` + + Returns + ------- + largest_region : boolean array + A boolean array with cells in the largest region marked as True, + and all other cells marked as False. + + ''' + + # First, break boolean array into unique, discrete regions/blobs + blobs_labels = label(bool_array, background=0, **kwargs) + + # Count the size of each blob, excluding the background class (0) + ids, counts = np.unique(blobs_labels[blobs_labels > 0], + return_counts=True) + + # Identify the region ID of the largest blob + largest_region_id = ids[np.argmax(counts)] + + # Produce a boolean array where 1 == the largest region + largest_region = blobs_labels == largest_region_id + + return largest_region + + +def transform_geojson_wgs_to_epsg(geojson, EPSG): + """ + Takes a geojson dictionary and converts it from WGS84 (EPSG:4326) to desired EPSG + + Parameters + ---------- + geojson: dict + a geojson dictionary containing a 'geometry' key, in WGS84 coordinates + EPSG: int + numeric code for the EPSG coordinate referecnce system to transform into + + Returns + ------- + transformed_geojson: dict + a geojson dictionary containing a 'coordinates' key, in the desired CRS + + """ + gg = Geometry(geojson['geometry'], CRS('epsg:4326')) + gg = gg.to_crs(CRS(f'epsg:{EPSG}')) + return gg.__geo_interface__ + + +def zonal_stats_parallel(shp, + raster, + statistics, + out_shp, + ncpus, + **kwargs): + + """ + Summarizing raster datasets based on vector geometries in parallel. + Each cpu recieves an equal chunk of the dataset. + Utilizes the perrygeo/rasterstats package. + + Parameters + ---------- + shp : str + Path to shapefile that contains polygons over + which zonal statistics are calculated + raster: str + Path to the raster from which the statistics are calculated. + This can be a virtual raster (.vrt). + statistics: list + list of statistics to calculate. e.g. + ['min', 'max', 'median', 'majority', 'sum'] + out_shp: str + Path to export shapefile containing zonal statistics. + ncpus: int + number of cores to parallelize the operations over. + kwargs: + Any other keyword arguments to rasterstats.zonal_stats() + See https://github.com/perrygeo/python-rasterstats for + all options + + Returns + ------- + Exports a shapefile to disk containing the zonal statistics requested + + """ + + # yields n sized chunks from list l (used for splitting task to multiple processes) + def chunks(l, n): + for i in range(0, len(l), n): + yield l[i:i + n] + + # calculates zonal stats and adds results to a dictionary + def worker(z, raster, d): + z_stats = zonal_stats(z, raster, stats=statistics, **kwargs) + for i in range(0, len(z_stats)): + d[z[i]['id']] = z_stats[i] + + # write output polygon + def write_output(zones, out_shp, d): + # copy schema and crs from input and add new fields for each statistic + schema = zones.schema.copy() + crs = zones.crs + for stat in statistics: + schema['properties'][stat] = 'float' + + with fiona.open(out_shp, 'w', 'ESRI Shapefile', schema, crs) as output: + for elem in zones: + for stat in statistics: + elem['properties'][stat] = d[elem['id']][stat] + output.write({'properties': elem['properties'], 'geometry': mapping(shape(elem['geometry']))}) + + with fiona.open(shp) as zones: + jobs = [] + + # create manager dictionary (polygon ids=keys, stats=entries) + # where multiple processes can write without conflicts + man = mp.Manager() + d = man.dict() + + # split zone polygons into 'ncpus' chunks for parallel processing + # and call worker() for each + split = chunks(zones, len(zones)//ncpus) + for z in split: + p = mp.Process(target=worker, args=(z, raster, d)) + p.start() + jobs.append(p) + + # wait that all chunks are finished + [j.join() for j in jobs] + + write_output(zones, out_shp, d) + + +def reverse_geocode(coords, site_classes=None, state_classes=None): + """ + Takes a latitude and longitude coordinate, and performs a reverse + geocode to return a plain-text description of the location in the + form: + + Site, State + + E.g.: `reverse_geocode(coords=(-35.282163, 149.128835))` + + 'Canberra, Australian Capital Territory' + + Parameters + ---------- + coords : tuple of floats + A tuple of (latitude, longitude) coordinates used to perform + the reverse geocode. + site_classes : list of strings, optional + A list of strings used to define the site part of the plain + text location description. Because the contents of the geocoded + address can vary greatly depending on location, these strings + are tested against the address one by one until a match is made. + + Defaults to: + + ``['city', 'town', 'village', 'suburb', 'hamlet', 'county', 'municipality']`` + + state_classes : list of strings, optional + A list of strings used to define the state part of the plain + text location description. These strings are tested against the + address one by one until a match is made. Defaults to: + `['state', 'territory']`. + Returns + ------- + If a valid geocoded address is found, a plain text location + description will be returned: + + 'Site, State' + + If no valid address is found, formatted coordinates will be returned + instead: + + 'XX.XX S, XX.XX E' + """ + + # Run reverse geocode using coordinates + geocoder = Nominatim(user_agent='Digital Earth Africa') + out = geocoder.reverse(coords) + + # Create plain text-coords as fall-back + lat = f'{-coords[0]:.2f} S' if coords[0] < 0 else f'{coords[0]:.2f} N' + lon = f'{-coords[1]:.2f} W' if coords[1] < 0 else f'{coords[1]:.2f} E' + + try: + + # Get address from geocoded data + address = out.raw['address'] + + # Use site and state classes if supplied; else use defaults + default_site_classes = ['city', 'town', 'village', 'suburb', 'hamlet', + 'county', 'municipality'] + default_state_classes = ['state', 'territory'] + site_classes = site_classes if site_classes else default_site_classes + state_classes = state_classes if state_classes else default_state_classes + + # Return the first site or state class that exists in address dict + site = next((address[k] for k in site_classes if k in address), None) + state = next((address[k] for k in state_classes if k in address), None) + + # If site and state exist in the data, return this. + # Otherwise, return N/E/S/W coordinates. + if site and state: + + # Return as site, state formatted string + return f'{site}, {state}' + + else: + + # If no geocoding result, return N/E/S/W coordinates + print('No valid geocoded location; returning coordinates instead') + return f'{lat}, {lon}' + + except (KeyError, AttributeError): + + # If no geocoding result, return N/E/S/W coordinates + print('No valid geocoded location; returning coordinates instead') + return f'{lat}, {lon}' + + +def sun_angles(dc, query): + """ + For a given spatiotemporal query, calculate mean sun + azimuth and elevation for each satellite observation, and + return these as a new `xarray.Dataset` with 'sun_elevation' + and 'sun_azimuth' variables. + + Parameters: + ----------- + dc : datacube.Datacube object + Datacube instance used to load data. + query : dict + A dictionary containing query parameters used to identify + satellite observations and load metadata. + + Returns: + -------- + sun_angles_ds : xarray.Dataset + An `xarray.set` containing a 'sun_elevation' and + 'sun_azimuth' variables. + """ + # Identify satellite datasets and group outputs using the + # same approach used to group satellite imagery (i.e. solar day) + gb = query_group_by(**query) + datasets = dc.find_datasets(**query) + dataset_array = dc.group_datasets(datasets, gb) + + # Load and take the mean of metadata from each product + sun_azimuth = xr_apply( + dataset_array, + lambda t, dd: np.mean([d.metadata.eo_sun_azimuth for d in dd]), + dtype=float, + ) + sun_elevation = xr_apply( + dataset_array, + lambda t, dd: np.mean([d.metadata.eo_sun_elevation for d in dd]), + dtype=float, + ) + + # Combine into new xarray.Dataset + sun_angles_ds = xr.merge( + [sun_elevation.rename("sun_elevation"), sun_azimuth.rename("sun_azimuth")] + ) + + return sun_angles_ds diff --git a/deafrica_tools/temporal.py b/deafrica_tools/temporal.py new file mode 100644 index 0000000..6539451 --- /dev/null +++ b/deafrica_tools/temporal.py @@ -0,0 +1,576 @@ +""" +Functions for calculating per-pixel temporal summary statistics on a +timeseries stored in a xarray.DataArray. + +The key functions are: + +.. autosummary:: + :caption: Primary functions + :nosignatures: + :toctree: gen + + xr_phenology + temporal_statistics + +.. autosummary:: + :nosignatures: + :toctree: gen + +""" + +import sys +import dask +import numpy as np +import xarray as xr +import hdstats +from packaging import version +from datacube.utils.geometry import assign_crs + + +def allNaN_arg(da, dim, stat): + """ + Calculate da.argmax() or da.argmin() while handling + all-NaN slices. Fills all-NaN locations with an + float and then masks the offending cells. + + Parameters + ---------- + da : xarray.DataArray + dim : str + Dimension over which to calculate argmax, argmin e.g. 'time' + stat : str + The statistic to calculte, either 'min' for argmin() + or 'max' for .argmax() + + Returns + ------- + xarray.DataArray + """ + # generate a mask where entire axis along dimension is NaN + mask = da.isnull().all(dim) + + if stat == "max": + y = da.fillna(float(da.min() - 1)) + y = y.argmax(dim=dim, skipna=True).where(~mask) + return y + + if stat == "min": + y = da.fillna(float(da.max() + 1)) + y = y.argmin(dim=dim, skipna=True).where(~mask) + return y + + +def _vpos(da): + """ + vPOS = Value at peak of season + """ + return da.max("time") + + +def _pos(da): + """ + POS = DOY of peak of season + """ + return da.isel(time=da.argmax("time")).time.dt.dayofyear + + +def _trough(da): + """ + Trough = Minimum value + """ + return da.min("time") + + +def _aos(vpos, trough): + """ + AOS = Amplitude of season + """ + return vpos - trough + + +def _vsos(da, pos, method_sos="first"): + """ + vSOS = Value at the start of season + Params + ----- + da : xarray.DataArray + method_sos : str, + If 'first' then vSOS is estimated + as the first positive slope on the + greening side of the curve. If 'median', + then vSOS is estimated as the median value + of the postive slopes on the greening side + of the curve. + """ + # select timesteps before peak of season (AKA greening) + greenup = da.where(da.time < pos.time) + # find the first order slopes + green_deriv = greenup.differentiate("time") + # find where the first order slope is postive + pos_green_deriv = green_deriv.where(green_deriv > 0) + # positive slopes on greening side + pos_greenup = greenup.where(~np.isnan(pos_green_deriv)) + # find the median + median = pos_greenup.median("time") + # distance of values from median + distance = pos_greenup - median + + if method_sos == "first": + # find index (argmin) where distance is most negative + idx = allNaN_arg(distance, "time", "min").astype("int16") + + if method_sos == "median": + # find index (argmin) where distance is smallest absolute value + idx = allNaN_arg(np.fabs(distance), "time", "min").astype("int16") + + return pos_greenup.isel(time=idx) + + +def _sos(vsos): + """ + SOS = DOY for start of season + """ + return vsos.time.dt.dayofyear + + +def _veos(da, pos, method_eos="last"): + """ + vEOS = Value at the end of season + Params + ----- + method_eos : str + If 'last' then vEOS is estimated + as the last negative slope on the + senescing side of the curve. If 'median', + then vEOS is estimated as the 'median' value + of the negative slopes on the senescing + side of the curve. + """ + # select timesteps before peak of season (AKA greening) + senesce = da.where(da.time > pos.time) + # find the first order slopes + senesce_deriv = senesce.differentiate("time") + # find where the fst order slope is negative + neg_senesce_deriv = senesce_deriv.where(~np.isnan(senesce_deriv < 0)) + # negative slopes on senescing side + neg_senesce = senesce.where(neg_senesce_deriv) + # find medians + median = neg_senesce.median("time") + # distance to the median + distance = neg_senesce - median + + if method_eos == "last": + # index where last negative slope occurs + idx = allNaN_arg(distance, "time", "min").astype("int16") + + if method_eos == "median": + # index where median occurs + idx = allNaN_arg(np.fabs(distance), "time", "min").astype("int16") + + return neg_senesce.isel(time=idx) + + +def _eos(veos): + """ + EOS = DOY for end of seasonn + """ + return veos.time.dt.dayofyear + + +def _los(da, eos, sos): + """ + LOS = Length of season (in DOY) + """ + los = eos - sos + #handle negative values + los = xr.where( + los >= 0, + los, + da.time.dt.dayofyear.values[-1] + (eos.where(los < 0) - sos.where(los < 0)), + ) + + return los + + +def _rog(vpos, vsos, pos, sos): + """ + ROG = Rate of Greening (Days) + """ + return (vpos - vsos) / (pos - sos) + + +def _ros(veos, vpos, eos, pos): + """ + ROG = Rate of Senescing (Days) + """ + return (veos - vpos) / (eos - pos) + + +def xr_phenology( + da, + stats=[ + "SOS", + "POS", + "EOS", + "Trough", + "vSOS", + "vPOS", + "vEOS", + "LOS", + "AOS", + "ROG", + "ROS", + ], + method_sos="first", + method_eos="last", + verbose=True +): + """ + Obtain land surface phenology metrics from an + xarray.DataArray containing a timeseries of a + vegetation index like NDVI. + + last modified June 2020 + + Parameters + ---------- + da : xarray.DataArray + DataArray should contain a 2D or 3D time series of a + vegetation index like NDVI, EVI + stats : list + list of phenological statistics to return. Regardless of + the metrics returned, all statistics are calculated + due to inter-dependencies between metrics. + Options include: + + * `SOS` = DOY of start of season + * `POS` = DOY of peak of season + * `EOS` = DOY of end of season + * `vSOS` = Value at start of season + * `vPOS` = Value at peak of season + * `vEOS` = Value at end of season + * `Trough` = Minimum value of season + * `LOS` = Length of season (DOY) + * `AOS` = Amplitude of season (in value units) + * `ROG` = Rate of greening + * `ROS` = Rate of senescence + + method_sos : str + If 'first' then vSOS is estimated as the first positive + slope on the greening side of the curve. If 'median', + then vSOS is estimated as the median value of the postive + slopes on the greening side of the curve. + method_eos : str + If 'last' then vEOS is estimated as the last negative slope + on the senescing side of the curve. If 'median', then vEOS is + estimated as the 'median' value of the negative slopes on the + senescing side of the curve. + + Returns + ------- + xarray.Dataset + Dataset containing variables for the selected + phenology statistics + + """ + # Check inputs before running calculations + if dask.is_dask_collection(da): + if version.parse(xr.__version__) < version.parse("0.16.0"): + raise TypeError( + "Dask arrays are not currently supported by this function, " + + "run da.compute() before passing dataArray." + ) + stats_dtype = { + "SOS": np.int16, + "POS": np.int16, + "EOS": np.int16, + "Trough": np.float32, + "vSOS": np.float32, + "vPOS": np.float32, + "vEOS": np.float32, + "LOS": np.int16, + "AOS": np.float32, + "ROG": np.float32, + "ROS": np.float32, + } + da_template = da.isel(time=0).drop("time") + template = xr.Dataset( + { + var_name: da_template.astype(var_dtype) + for var_name, var_dtype in stats_dtype.items() + if var_name in stats + } + ) + da_all_time = da.chunk({"time": -1}) + + lazy_phenology = da_all_time.map_blocks( + xr_phenology, + kwargs=dict( + stats=stats, + method_sos=method_sos, + method_eos=method_eos, + ), + template=xr.Dataset(template), + ) + + try: + crs = da.geobox.crs + lazy_phenology = assign_crs(lazy_phenology, str(crs)) + except: + pass + + return lazy_phenology + + if method_sos not in ("median", "first"): + raise ValueError("method_sos should be either 'median' or 'first'") + + if method_eos not in ("median", "last"): + raise ValueError("method_eos should be either 'median' or 'last'") + + # If stats supplied is not a list, convert to list. + stats = stats if isinstance(stats, list) else [stats] + + # try to grab the crs info + try: + crs = da.geobox.crs + except: + pass + + # remove any remaining all-NaN pixels + mask = da.isnull().all("time") + da = da.where(~mask, other=0) + + # calculate the statistics + if verbose: + print(" Phenology...") + vpos = _vpos(da) + pos = _pos(da) + trough = _trough(da) + aos = _aos(vpos, trough) + vsos = _vsos(da, pos, method_sos=method_sos) + sos = _sos(vsos) + veos = _veos(da, pos, method_eos=method_eos) + eos = _eos(veos) + los = _los(da, eos, sos) + rog = _rog(vpos, vsos, pos, sos) + ros = _ros(veos, vpos, eos, pos) + + # Dictionary containing the statistics + stats_dict = { + "SOS": sos.astype(np.int16), + "EOS": eos.astype(np.int16), + "vSOS": vsos.astype(np.float32), + "vPOS": vpos.astype(np.float32), + "Trough": trough.astype(np.float32), + "POS": pos.astype(np.int16), + "vEOS": veos.astype(np.float32), + "LOS": los.astype(np.int16), + "AOS": aos.astype(np.float32), + "ROG": rog.astype(np.float32), + "ROS": ros.astype(np.float32), + } + + # intialise dataset with first statistic + ds = stats_dict[stats[0]].to_dataset(name=stats[0]) + + # add the other stats to the dataset + for stat in stats[1:]: + if verbose: + print(" " + stat) + stats_keep = stats_dict.get(stat) + ds[stat] = stats_dict[stat] + + try: + ds = assign_crs(ds, str(crs)) + except: + pass + + return ds.drop("time") + + +def temporal_statistics(da, stats): + """ + Calculate various generic summary statistics on any timeseries. + + This function uses the hdstats temporal library: + https://github.com/daleroberts/hdstats/blob/master/hdstats/ts.pyx + + last modified June 2020 + + Parameters + ---------- + da : xarray.DataArray + DataArray should contain a 3D time series. + stats : list + list of temporal statistics to calculate. + Options include: + + * 'discordance' = + * 'f_std' = std of discrete fourier transform coefficients, returns + three layers: f_std_n1, f_std_n2, f_std_n3 + * 'f_mean' = mean of discrete fourier transform coefficients, returns + three layers: f_mean_n1, f_mean_n2, f_mean_n3 + * 'f_median' = median of discrete fourier transform coefficients, returns + three layers: f_median_n1, f_median_n2, f_median_n3 + * 'mean_change' = mean of discrete difference along time dimension + * 'median_change' = median of discrete difference along time dimension + * 'abs_change' = mean of absolute discrete difference along time dimension + * 'complexity' = + * 'central_diff' = + * 'num_peaks' : The number of peaks in the timeseries, defined with a local + window of size 10. NOTE: This statistic is very slow + + Returns + ------- + xarray.Dataset + Dataset containing variables for the selected + temporal statistics + + """ + + # if dask arrays then map the blocks + if dask.is_dask_collection(da): + if version.parse(xr.__version__) < version.parse("0.16.0"): + raise TypeError( + "Dask arrays are only supported by this function if using, " + + "xarray v0.16, run da.compute() before passing dataArray." + ) + + # create a template that matches the final datasets dims & vars + arr = da.isel(time=0).drop("time") + + # deal with the case where fourier is first in the list + if stats[0] in ("f_std", "f_median", "f_mean"): + template = xr.zeros_like(arr).to_dataset(name=stats[0] + "_n1") + template[stats[0] + "_n2"] = xr.zeros_like(arr) + template[stats[0] + "_n3"] = xr.zeros_like(arr) + + for stat in stats[1:]: + if stat in ("f_std", "f_median", "f_mean"): + template[stat + "_n1"] = xr.zeros_like(arr) + template[stat + "_n2"] = xr.zeros_like(arr) + template[stat + "_n3"] = xr.zeros_like(arr) + else: + template[stat] = xr.zeros_like(arr) + else: + template = xr.zeros_like(arr).to_dataset(name=stats[0]) + + for stat in stats: + if stat in ("f_std", "f_median", "f_mean"): + template[stat + "_n1"] = xr.zeros_like(arr) + template[stat + "_n2"] = xr.zeros_like(arr) + template[stat + "_n3"] = xr.zeros_like(arr) + else: + template[stat] = xr.zeros_like(arr) + try: + template = template.drop("spatial_ref") + except: + pass + + # ensure the time chunk is set to -1 + da_all_time = da.chunk({"time": -1}) + + # apply function across chunks + lazy_ds = da_all_time.map_blocks( + temporal_statistics, kwargs={"stats": stats}, template=template + ) + + try: + crs = da.geobox.crs + lazy_ds = assign_crs(lazy_ds, str(crs)) + except: + pass + + return lazy_ds + + # If stats supplied is not a list, convert to list. + stats = stats if isinstance(stats, list) else [stats] + + # grab all the attributes of the xarray + x, y, time, attrs = da.x, da.y, da.time, da.attrs + + # deal with any all-NaN pixels by filling with 0's + mask = da.isnull().all("time") + da = da.where(~mask, other=0) + + # ensure dim order is correct for functions + da = da.transpose("y", "x", "time").values + + stats_dict = { + "discordance": lambda da: hdstats.discordance(da, n=10), + "f_std": lambda da: hdstats.fourier_std(da, n=3, step=5), + "f_mean": lambda da: hdstats.fourier_mean(da, n=3, step=5), + "f_median": lambda da: hdstats.fourier_median(da, n=3, step=5), + "mean_change": lambda da: hdstats.mean_change(da), + "median_change": lambda da: hdstats.median_change(da), + "abs_change": lambda da: hdstats.mean_abs_change(da), + "complexity": lambda da: hdstats.complexity(da), + "central_diff": lambda da: hdstats.mean_central_diff(da), + "num_peaks": lambda da: hdstats.number_peaks(da, 10), + } + + print(" Statistics:") + # if one of the fourier functions is first (or only) + # stat in the list then we need to deal with this + if stats[0] in ("f_std", "f_median", "f_mean"): + print(" " + stats[0]) + stat_func = stats_dict.get(str(stats[0])) + zz = stat_func(da) + n1 = zz[:, :, 0] + n2 = zz[:, :, 1] + n3 = zz[:, :, 2] + + # intialise dataset with first statistic + ds = xr.DataArray( + n1, attrs=attrs, coords={"x": x, "y": y}, dims=["y", "x"] + ).to_dataset(name=stats[0] + "_n1") + + # add other datasets + for i, j in zip([n2, n3], ["n2", "n3"]): + ds[stats[0] + "_" + j] = xr.DataArray( + i, attrs=attrs, coords={"x": x, "y": y}, dims=["y", "x"] + ) + else: + # simpler if first function isn't fourier transform + first_func = stats_dict.get(str(stats[0])) + print(" " + stats[0]) + ds = first_func(da) + + # convert back to xarray dataset + ds = xr.DataArray( + ds, attrs=attrs, coords={"x": x, "y": y}, dims=["y", "x"] + ).to_dataset(name=stats[0]) + + # loop through the other functions + for stat in stats[1:]: + print(" " + stat) + + # handle the fourier transform examples + if stat in ("f_std", "f_median", "f_mean"): + stat_func = stats_dict.get(str(stat)) + zz = stat_func(da) + n1 = zz[:, :, 0] + n2 = zz[:, :, 1] + n3 = zz[:, :, 2] + + for i, j in zip([n1, n2, n3], ["n1", "n2", "n3"]): + ds[stat + "_" + j] = xr.DataArray( + i, attrs=attrs, coords={"x": x, "y": y}, dims=["y", "x"] + ) + + else: + # Select a stats function from the dictionary + # and add to the dataset + stat_func = stats_dict.get(str(stat)) + ds[stat] = xr.DataArray( + stat_func(da), attrs=attrs, coords={"x": x, "y": y}, dims=["y", "x"] + ) + + # try to add back the geobox + try: + crs = da.geobox.crs + ds = assign_crs(ds, str(crs)) + except: + pass + + return ds diff --git a/deafrica_tools/untitled.txt b/deafrica_tools/untitled.txt new file mode 100644 index 0000000..e69de29 diff --git a/deafrica_tools/wetlands.py b/deafrica_tools/wetlands.py new file mode 100644 index 0000000..4d7377a --- /dev/null +++ b/deafrica_tools/wetlands.py @@ -0,0 +1,732 @@ +""" +Functions for working with the Wetlands Insight Tool (WIT) +""" + +# Import required packages + +# Force GeoPandas to use Shapely instead of PyGEOS +# In a future release, GeoPandas will switch to using Shapely by default. +import os +os.environ['USE_PYGEOS'] = '0' + +import warnings +import numpy as np +import pandas as pd +import geopandas as gpd +import seaborn as sns +import xarray as xr +import matplotlib.pyplot as plt +from skimage import exposure +import matplotlib.animation as animation +import matplotlib.patheffects as PathEffects +from mpl_toolkits.axes_grid1.inset_locator import inset_axes +from dask.distributed import progress + +import datacube +from datacube.utils import masking +from datacube.utils import geometry + +from deafrica_tools.bandindices import calculate_indices +from deafrica_tools.datahandling import load_ard, wofs_fuser +from deafrica_tools.spatial import xr_rasterize +from deafrica_tools.classification import HiddenPrints + + +def WIT_drill( + gdf, + time, + min_gooddata=0.85, + TCW_threshold=-0.035, + resample_frequency=None, + export_csv=None, + dask_chunks=None, + verbose=False, + verbose_progress=False, +): + """ + The Wetlands Insight Tool run onver an extent covered by a polygon. + This function loads FC, WOfS, and Landsat data, and calculates tasseled + cap wetness, in order to determine the dominant land cover class + within a polygon at each satellite observation. + + The output is a pandas dataframe containing a timeseries of the relative + fractions of each class at each time-step. This forms the input to produce + a stacked line-plot. + + Last modified: Oct 2021 + + Parameters + ---------- + gdf : geopandas.GeoDataFrame + The dataframe must only contain a single row, + containing the polygon you wish to interrograte. + time : tuple + a tuple containing the time range over which to run the WIT. + e.g. ('2015-01' , '2019-12') + min_gooddata : Float, optional + A number between 0 and 1 (e.g 0.8) indicating the minimum percentage + of good quality pixels required for a satellite observation to be loaded + and therefore included in the WIT plot. This number should, at a minimum, + be set to 0.80 to limit biases in the result if not resampling the time-series. + If resampling the data using the parameter `resample_frequency`, then + setting this number to 0 (or a low float number) is acceptable. + TCW_threshold : Int, optional + The tasseled cap wetness threshold, beyond which a pixel will be + considered 'wet'. Defaults to -0.035. + resample_frequency : str + Option for resampling time-series of input datasets. This option is useful + for either smoothing the WIT plot, or because the area of analysis is larger + than a scene width and therefore requires composites. Options include any + str accepted by `xarray.resample(time=)`. The resampling method used is .max() + export_csv : str, optional + To save the returned pandas dataframe as a .csv file, pass a + a location string (e.g. 'output/results.csv') + dask_chunks : dict, optional + To lazily load the datasets using dask, pass a dictionary containing + the dimensions over which to chunk e.g. {'time':-1, 'x':250, 'y':250}. + verbose: bool, optional + If true, print statements are putput detailing the progress of the tool. + verbose_progress: bool, optional + For use with Dask progress bar + + Returns + ------- + df : Pandas.Dataframe + A pandas dataframe containing the timeseries of relative fractions + of each land cover class (WOfs, FC, TCW) + + """ + # add geom to dc query dict + if isinstance(gdf, datacube.utils.geometry._base.Geometry): + gdf = gpd.GeoDataFrame({'col1':['name'],'geometry':gdf.geom}, crs=gdf.crs) + geom = geometry.Geometry(geom=gdf.iloc[0].geometry, crs=gdf.crs) + query = {"geopolygon": geom, "time": time} + + # Create a datacube instance + dc = datacube.Datacube(app="wetlands insight tool") + + # load landsat 5,7,8 data + warnings.filterwarnings("ignore") + + if verbose_progress: + print("Loading Landsat data") + ds_ls = load_ard( + dc=dc, + products=["ls8_sr", "ls7_sr", "ls5_sr"], + output_crs="epsg:6933", + min_gooddata=min_gooddata, + mask_filters=(['opening', 3], ['dilation', 3]), + measurements=["red", "green", "blue", "nir", "swir_1", "swir_2"], + dask_chunks=dask_chunks, + group_by="solar_day", + resolution=(-30, 30), + verbose=verbose, + **query, + ) + + # create polygon mask + mask = xr_rasterize(gdf.iloc[[0]], ds_ls) + ds_ls = ds_ls.where(mask) + + # calculate tasselled cap wetness within masked AOI + if verbose: + print("calculating tasseled cap wetness index ") + + with HiddenPrints(): #suppres the prints from this func + tcw = calculate_indices( + ds_ls, index=["TCW"], normalise=False, satellite_mission="ls", drop=True + ) + + if resample_frequency is not None: + if verbose: + print('Resampling TCW to '+ resample_frequency) + tcw = tcw.resample(time=resample_frequency).max() + + tcw = tcw.TCW >= TCW_threshold + tcw = tcw.where(mask, 0) + tcw = tcw.persist() + + if verbose: + print("Loading WOfS layers ") + + wofls = dc.load( + product="wofs_ls", + like=ds_ls, + fuse_func=wofs_fuser, + dask_chunks=dask_chunks, + collection_category="T1", + ) + + # boolean of wet/dry + wofls_wet = masking.make_mask(wofls.water, wet=True) + + if resample_frequency is not None: + if verbose: + print('Resampling WOfS to '+ resample_frequency) + wofls_wet = wofls_wet.resample(time=resample_frequency).max() + + # mask sure wofs matches other datasets + wofls_wet = wofls_wet.where(wofls_wet.time == tcw.time) + + # apply the polygon mask + wofls_wet = wofls_wet.where(mask) + + # load Fractional cover + if verbose: + print("Loading fractional Cover") + + # load fractional cover + fc_ds = dc.load( + product="fc_ls", + time=time, + dask_chunks=dask_chunks, + like=ds_ls, + measurements=["pv", "npv", "bs"], + collection_category="T1", + ) + + # use wofls mask to cloud mask FC + clear_and_dry = masking.make_mask(wofls, dry=True).water + fc_ds = fc_ds.where(clear_and_dry) + + if resample_frequency is not None: + if verbose: + print('Resampling FC to '+ resample_frequency) + fc_ds = fc_ds.resample(time=resample_frequency).max() + + # mask sure fc matches other datasets + fc_ds = fc_ds.where(fc_ds.time == tcw.time) + + # mask with polygon + fc_ds = fc_ds.where(mask) + + # mask with TC wetness + fc_ds_noTCW = fc_ds.where(tcw == False) + + if verbose: + print("Generating classification") + + # Cast the dataset to a dataarray + fc_ds_noTCW = fc_ds_noTCW.to_array(dim="variable", name="fc_ds_noTCW") + + # turn FC array into integer only as nanargmax doesn't + # seem to handle floats the way we want it to + fc_int = fc_ds_noTCW.astype("int8") + + # use nanargmax to get the index of the maximum value + BSPVNPV = fc_int.argmax(dim="variable") + + #int dytype remocves NaNs so we need to create mask again + FC_mask = np.isfinite(fc_ds_noTCW).all(dim="variable") + BSPVNPV = BSPVNPV.where(FC_mask) + + # Restack the Fractional cover dataset all together + # CAUTION:ARGMAX DEPENDS ON ORDER OF VARIABALES IN + # DATASET. NEED TO ADJUST BELOW DEPENDING ON ORDER OF FC VARIABLES + + FC_dominant = xr.Dataset( + { + "bs": (BSPVNPV == 2).where(FC_mask), + "pv": (BSPVNPV == 0).where(FC_mask), + "npv": (BSPVNPV == 1).where(FC_mask), + } + ) + + # pixel counts + pixels = mask.sum(dim=["x", "y"]) + + + if verbose_progress: + print("Computing wetness") + tcw_pixel_count = tcw.sum(dim=["x", "y"]).compute() + + if verbose_progress: + print("Computing green veg, dry veg, and bare soil") + FC_count = FC_dominant.sum(dim=["x", "y"]).compute() + + if verbose_progress: + print("Computing open water") + wofs_pixels = wofls_wet.sum(dim=["x", "y"]).compute() + + # count percentages + wofs_area_percent = (wofs_pixels / pixels) * 100 + tcw_area_percent = (tcw_pixel_count / pixels) * 100 + tcw_less_wofs = tcw_area_percent - wofs_area_percent # wet not wofs + + # Fractional cover pixel count method + # Get number of FC pixels, divide by total number of pixels per polygon + # Work out the number of nodata pixels in the data + BS_percent = (FC_count.bs / pixels) * 100 + PV_percent = (FC_count.pv / pixels) * 100 + NPV_percent = (FC_count.npv / pixels) * 100 + NoData_count = (( + 100 - wofs_area_percent - tcw_less_wofs - PV_percent - NPV_percent - BS_percent + ) / 100) * pixels + + # re-do percentages but now handling any no-data pixels within polygon + BS_percent = (FC_count.bs / (pixels - NoData_count)) * 100 + PV_percent = (FC_count.pv / (pixels - NoData_count)) * 100 + NPV_percent = (FC_count.npv / (pixels - NoData_count)) * 100 + wofs_area_percent = (wofs_pixels / (pixels - NoData_count)) * 100 + tcw_area_percent = (tcw_pixel_count / (pixels - NoData_count)) * 100 + tcw_less_wofs = tcw_area_percent - wofs_area_percent + + # Sometimes when we resample datastes, WOfS extent can be + # greater than the wetness extent, thus make negative values == zero + tcw_less_wofs = tcw_less_wofs.where(tcw_less_wofs>=0, 0) + + # start setup of dataframe by adding only one dataset + df = pd.DataFrame( + data=wofs_area_percent.data, + index=wofs_area_percent.time.values, + columns=["wofs_area_percent"], + ) + + # add data into pandas dataframe for export + df["wet_percent"] = tcw_less_wofs.data + df["green_veg_percent"] = PV_percent.data + df["dry_veg_percent"] = NPV_percent.data + df["bare_soil_percent"] = BS_percent.data + + # round numbers + df = df.round(2) + + # save the csv of the output data used to create the stacked plot for the polygon drill + if export_csv: + if verbose: + print("exporting csv: " + export_csv) + df.to_csv(export_csv, index_label="Datetime") + + return df + + +def animated_timeseries_WIT( + ds, + df, + output_path, + width_pixels=1000, + interval=200, + bands=["red", "green", "blue"], + percentile_stretch=(0.02, 0.98), + image_proc_func=None, + title=False, + show_date=True, + annotation_kwargs={}, + onebandplot_cbar=True, + onebandplot_kwargs={}, + shapefile_path=None, + shapefile_kwargs={}, + pandasplot_kwargs={}, + time_dim="time", + x_dim="x", + y_dim="y", +): + + ############### + # Setup steps # + ############### + + # Test if all dimensions exist in dataset + if time_dim in ds and x_dim in ds and y_dim in ds: + + # Test if there is one or three bands, and that all exist in both datasets: + if ((len(bands) == 3) | (len(bands) == 1)) & all( + [(b in ds.data_vars) for b in bands] + ): + + # Import xarrays as lists of three band numpy arrays + imagelist, vmin, vmax = _ds_to_arrraylist( + ds, + bands=bands, + time_dim=time_dim, + x_dim=x_dim, + y_dim=y_dim, + percentile_stretch=percentile_stretch, + image_proc_func=image_proc_func, + ) + + # Get time, x and y dimensions of dataset and calculate width vs height of plot + timesteps = len(ds[time_dim]) + width = len(ds[x_dim]) + height = len(ds[y_dim]) + width_ratio = float(width) / float(height) + height = 10.0 / width_ratio + + # If title is supplied as a string, multiply out to a list with one string per timestep. + # Otherwise, use supplied list for plot titles. + if isinstance(title, str) or isinstance(title, bool): + title_list = [title] * timesteps + else: + title_list = title + + # Set up annotation parameters that plt.imshow plotting for single band array images. + # The nested dict structure sets default values which can be overwritten/customised by the + # manually specified `onebandplot_kwargs` + onebandplot_kwargs = dict( + { + "cmap": "Greys", + "interpolation": "bilinear", + "vmin": vmin, + "vmax": vmax, + "tick_colour": "black", + "tick_fontsize": 11, + }, + **onebandplot_kwargs, + ) + + # Use pop to remove the two special tick kwargs from the onebandplot_kwargs dict, and save individually + onebandplot_tick_colour = onebandplot_kwargs.pop("tick_colour") + onebandplot_tick_fontsize = onebandplot_kwargs.pop("tick_fontsize") + + # Set up annotation parameters that control font etc. The nested dict structure sets default + # values which can be overwritten/customised by the manually specified `annotation_kwargs` + annotation_kwargs = dict( + { + "xy": (1, 1), + "xycoords": "axes fraction", + "xytext": (-5, -5), + "textcoords": "offset points", + "horizontalalignment": "right", + "verticalalignment": "top", + "fontsize": 15, + "color": "white", + "path_effects": [ + PathEffects.withStroke(linewidth=3, foreground="black") + ], + }, + **annotation_kwargs, + ) + + # Define default plotting parameters for the overlaying shapefile(s). The nested dict structure sets + # default values which can be overwritten/customised by the manually specified `shapefile_kwargs` + shapefile_kwargs = dict( + {"linewidth": 2, "edgecolor": "black", "facecolor": "#00000000"}, + **shapefile_kwargs, + ) + + # Define default plotting parameters for the right-hand line plot. The nested dict structure sets + # default values which can be overwritten/customised by the manually specified `pandasplot_kwargs` + pandasplot_kwargs = dict({}, **pandasplot_kwargs) + + ################### + # Initialise plot # + ################### + + # Set up figure + fig, (ax1, ax2) = plt.subplots( + ncols=2, gridspec_kw={"width_ratios": [1, 2]} + ) + fig.subplots_adjust(left=0, bottom=0, right=1, top=1, wspace=0.2, hspace=0) + fig.set_size_inches(10.0, height * 0.5, forward=True) + ax1.axis("off") + ax2.margins(x=0.01) + ax2.xaxis.label.set_visible(False) + + # Initialise axesimage objects to be updated during animation, setting extent from dims + extents = [ + float(ds[x_dim].min()), + float(ds[x_dim].max()), + float(ds[y_dim].min()), + float(ds[y_dim].max()), + ] + im = ax1.imshow(imagelist[0], extent=extents, **onebandplot_kwargs) + + # Initialise right panel and set y axis limits + # set up color palette + pal = [ + sns.xkcd_rgb["cobalt blue"], + sns.xkcd_rgb["neon blue"], + sns.xkcd_rgb["grass"], + sns.xkcd_rgb["beige"], + sns.xkcd_rgb["brown"], + ] + + # make a stacked area plot + ax2.stackplot( + df.index, + df.wofs_area_percent, + df.wet_percent, + df.green_veg_percent, + df.dry_veg_percent, + df.bare_soil_percent, + labels=["open water", "wet", "green veg", "dry veg", "bare soil"], + colors=pal, + alpha=0.6, + **pandasplot_kwargs, + ) + + ax2.legend(loc="lower left", framealpha=0.6) + + df1 = pd.DataFrame( + { + "wofs_area_percent": df.wofs_area_percent, + "wet_percent": df.wofs_area_percent + df.wet_percent, + "green_veg_percent": df.wofs_area_percent + + df.wet_percent + + df.green_veg_percent, + "dry_veg_percent": df.wofs_area_percent + + df.wet_percent + + df.green_veg_percent + + df.dry_veg_percent, + "bare_soil_percent": df.dry_veg_percent + + df.green_veg_percent + + df.wofs_area_percent + + df.wet_percent + + df.bare_soil_percent, + } + ) + df1 = df1.set_index(df.index) + + line_test = df1.plot( + ax=ax2, legend=False, color="black", **pandasplot_kwargs + ) + + # set axis limits to the min and max + ax2.set(xlim=(df.index[0], df.index[-1]), ylim=(0, 100)) + + # add a legend and a tight plot box + + ax2.set_title("Fractional Cover, Wetness, and Water") + + # Initialise annotation objects to be updated during animation + t = ax1.annotate("", **annotation_kwargs) + + ######################### + # Add optional overlays # + ######################### + + # Optionally add shapefile overlay(s) from either string path or list of string paths + if isinstance(shapefile_path, str): + + shapefile = gpd.read_file(shapefile_path) + shapefile.plot(**shapefile_kwargs, ax=ax1) + + elif isinstance(shapefile_path, list): + + # Iterate through list of string paths + for shapefile in shapefile_path: + + shapefile = gpd.read_file(shapefile) + shapefile.plot(**shapefile_kwargs, ax=ax1) + + # After adding shapefile, fix extents of plot + ax1.set_xlim(extents[0], extents[1]) + ax1.set_ylim(extents[2], extents[3]) + + # Optionally add colourbar for one band images + if (len(bands) == 1) & onebandplot_cbar: + _add_colourbar( + ax1, + im, + tick_fontsize=onebandplot_tick_fontsize, + tick_colour=onebandplot_tick_colour, + vmin=onebandplot_kwargs["vmin"], + vmax=onebandplot_kwargs["vmax"], + ) + + ######################################## + # Create function to update each frame # + ######################################## + + # Function to update figure + + def update_figure(frame_i): + + #################### + # Plot image panel # + #################### + + # If possible, extract dates from time dimension + try: + + # Get human-readable date info (e.g. "16 May 1990") + ts = ds[time_dim][{time_dim: frame_i}].dt + year = ts.year.item() + month = ts.month.item() + day = ts.day.item() + date_string = "{} {} {}".format( + day, calendar.month_abbr[month], year + ) + + except: + + date_string = ds[time_dim][{time_dim: frame_i}].values.item() + + # Create annotation string based on title and date specifications: + title = title_list[frame_i] + if title and show_date: + title_date = "{}\n{}".format(date_string, title) + elif title and not show_date: + title_date = "{}".format(title) + elif show_date and not title: + title_date = "{}".format(date_string) + else: + title_date = "" + + # Update left panel with annotation and image + im.set_array(imagelist[frame_i]) + t.set_text(title_date) + + ######################## + # Plot linegraph panel # + ######################## + + # Create list of artists to return + artist_list = [im, t] + + # Update right panel with temporal line subset, adding each new line into artist_list + for i, line in enumerate(line_test.lines): + + # Clip line data to current time, and get x and y values + y = df1[ + df1.index + <= datetime(year=year, month=month, day=day, hour=23, minute=59) + ].iloc[:, i] + x = df1[ + df1.index + <= datetime(year=year, month=month, day=day, hour=23, minute=59) + ].index + + # Plot lines after stripping NaNs (this produces continuous, unbroken lines) + line.set_data(x[y.notnull()], y[y.notnull()]) + artist_list.extend([line]) + + # Return the artists set + return artist_list + + # Nicely space subplots + fig.tight_layout() + + ############################## + # Generate and run animation # + ############################## + + # Generate animation + ani = animation.FuncAnimation( + fig=fig, + func=update_figure, + frames=timesteps, + interval=interval, + blit=True, + ) + + # Export as either MP4 or GIF + if output_path[-3:] == "mp4": + print(" Exporting animation to {}".format(output_path)) + ani.save(output_path, dpi=width_pixels / 10.0) + + elif output_path[-3:] == "wmv": + print(" Exporting animation to {}".format(output_path)) + ani.save( + output_path, + dpi=width_pixels / 10.0, + writer=animation.FFMpegFileWriter( + fps=1000 / interval, bitrate=4000, codec="wmv2" + ), + ) + + elif output_path[-3:] == "gif": + print(" Exporting animation to {}".format(output_path)) + ani.save(output_path, dpi=width_pixels / 10.0, writer="imagemagick") + + else: + print(" Output file type must be either .mp4, .wmv or .gif") + + else: + print( + "Please select either one or three bands that all exist in the input dataset" + ) + + else: + print( + "At least one x, y or time dimension does not exist in the input dataset. Please use the `time_dim`," + "`x_dim` or `y_dim` parameters to override the default dimension names used for plotting" + ) + + +# Define function to convert xarray dataset to list of one or three band numpy arrays + + +def _ds_to_arrraylist( + ds, bands, time_dim, x_dim, y_dim, percentile_stretch, image_proc_func=None +): + """ + Converts an xarray dataset to a list of numpy arrays for plt.imshow plotting + """ + + # Compute percents + p_low, p_high = ds[bands].to_array().quantile(percentile_stretch).values + + array_list = [] + for i, timestep in enumerate(ds[time_dim]): + + # Select single timestep from the data array + ds_i = ds[{time_dim: i}] + + # Get shape of array + x = len(ds[x_dim]) + y = len(ds[y_dim]) + + if len(bands) == 1: + + # Create new one band array + img_toshow = exposure.rescale_intensity( + ds_i[bands[0]].values, in_range=(p_low, p_high), out_range="image" + ) + + else: + + # Create new three band array + rawimg = np.zeros((y, x, 3), dtype=np.float32) + + # Add xarray bands into three dimensional numpy array + for band, colour in enumerate(bands): + + rawimg[:, :, band] = ds_i[colour].values + + # Stretch contrast using percentile values + img_toshow = exposure.rescale_intensity( + rawimg, in_range=(p_low, p_high), out_range=(0, 1.0) + ) + + # Optionally image processing + if image_proc_func: + + img_toshow = image_proc_func(img_toshow).clip(0, 1) + + array_list.append(img_toshow) + + return array_list, p_low, p_high + + +def _add_colourbar( + ax, im, vmin, vmax, cmap="Greys", tick_fontsize=15, tick_colour="black" +): + """ + Add a nicely formatted colourbar to an animation panel + """ + + # Add colourbar + axins2 = inset_axes(ax, width="97%", height="4%", loc=8, borderpad=1) + plt.gcf().colorbar( + im, cax=axins2, orientation="horizontal", ticks=np.linspace(vmin, vmax, 3) + ) + axins2.xaxis.set_ticks_position("top") + axins2.tick_params(axis="x", colors=tick_colour, labelsize=tick_fontsize) + + # Justify left and right labels to edge of plot + axins2.get_xticklabels()[0].set_horizontalalignment("left") + axins2.get_xticklabels()[-1].set_horizontalalignment("right") + labels = [item.get_text() for item in axins2.get_xticklabels()] + labels[0] = " " + labels[0] + labels[-1] = labels[-1] + " " + + +if __name__ == "__main__": + # print that we are running the testing + print("Testing..") + # import doctest to test our module for documentation + import doctest + + doctest.testmod() + print("Testing done") diff --git a/new_import.py b/new_import.py new file mode 100644 index 0000000..a49a430 --- /dev/null +++ b/new_import.py @@ -0,0 +1,320 @@ +import matplotlib.pyplot as plt + +# Common imports and settings +import os, sys +os.environ['USE_PYGEOS'] = '0' +from IPython.display import Markdown +import pandas as pd +pd.set_option("display.max_rows", None) +import xarray as xr + +# Datacube +import datacube +from datacube.utils.rio import configure_s3_access +from datacube.utils import masking +from datacube.utils.cog import write_cog +# https://github.com/GeoscienceAustralia/dea-notebooks/tree/develop/Tools +from dea_tools.plotting import display_map, rgb +from dea_tools.datahandling import mostcommon_crs + +# EASI defaults +easinotebooksrepo = '/home/jovyan/easi-notebooks' +if easinotebooksrepo not in sys.path: sys.path.append(easinotebooksrepo) +from easi_tools import EasiDefaults, xarray_object_size, notebook_utils, unset_cachingproxy +from easi_tools.load_s2l2a import load_s2l2a_with_offset +from dask.distributed import progress + +# Data tools +import numpy as np +from datetime import datetime + +# Datacube +from datacube.utils import masking # https://github.com/opendatacube/datacube-core/blob/develop/datacube/utils/masking.py +from odc.algo import enum_to_bool # https://github.com/opendatacube/odc-algo/blob/main/odc/algo/_masking.py +from odc.algo import xr_reproject # https://github.com/opendatacube/odc-algo/blob/main/odc/algo/_warp.py +from datacube.utils.geometry import GeoBox, box # https://github.com/opendatacube/datacube-core/blob/develop/datacube/utils/geometry/_base.py + +# Holoviews, Datashader and Bokeh +import hvplot.pandas +import hvplot.xarray +import holoviews as hv +import panel as pn +import colorcet as cc +import cartopy.crs as ccrs +from datashader import reductions +from holoviews import opts +from utils import load_data_geo +import rasterio +import rioxarray +# import geoviews as gv +# from holoviews.operation.datashader import rasterize +hv.extension('bokeh', logo=False) + +from deafrica_tools.bandindices import calculate_indices +from sklearn.ensemble import RandomForestClassifier +from sklearn.model_selection import train_test_split +from sklearn.metrics import accuracy_score, classification_report +from sklearn.preprocessing import LabelEncoder + +from sklearn.pipeline import Pipeline +from sklearn.ensemble import RandomForestClassifier +from sklearn.impute import SimpleImputer +from sklearn.preprocessing import StandardScaler +from sklearn.model_selection import GridSearchCV +from sklearn.model_selection import train_test_split +from sklearn.metrics import accuracy_score +from shapely.geometry import Point, Polygon +import geopandas as gpd +from pyproj import CRS +from matplotlib.colors import ListedColormap +from holoviews import opts +from datashader import reductions +from bokeh.models.tickers import FixedTicker +from rioxarray.merge import merge_arrays + +import joblib + + +def load_data(dc, date_range, longtitude_range, latitude_range): + product = 's2_l2a' + query = { + 'product': product, # Product name + 'x': longtitude_range, # "x" axis bounds + 'y': latitude_range, # "y" axis bounds + 'time': date_range, # Any parsable date strings + } + native_crs = notebook_utils.mostcommon_crs(dc, query) + print(f'Most common native CRS: {native_crs}') + measurements = ['blue', 'green', 'red', 'nir', 'scl'] + + load_params = { + 'measurements': measurements, # Selected measurement or alias names + 'output_crs': native_crs, # Target EPSG code + 'resolution': (-10, 10), # Target resolution + 'group_by': 'solar_day', # Scene grouping + 'dask_chunks': {'x': 2048, 'y': 2048}, # Dask chunks + } + data = load_s2l2a_with_offset( + dc, + query | load_params # Combine the two dicts that contain our search and load parameters + ) + return data + + +def mask_clean(data): + flag_name = 'scl' + flag_desc = masking.describe_variable_flags(data[flag_name]) # Pandas dataframe + display(flag_desc) + display(flag_desc.loc['qa'].values[1]) + # Create a "data quality" Mask layer + flags_def = flag_desc.loc['qa'].values[1] + good_pixel_flags = [flags_def[str(i)] for i in [2, 4, 5, 6]] # To pass strings to enum_to_bool() + + # enum_to_bool calculates the pixel-wise "or" of each set of pixels given by good_pixel_flags + # 1 = good data + # 0 = "bad" data + good_pixel_mask = enum_to_bool(data[flag_name], good_pixel_flags) + data_layer_names = [x for x in data.data_vars if x != 'scl'] + # Apply good pixel mask to blue, green, red and nir. + result = data[data_layer_names].where(good_pixel_mask).persist() + return result + + +def fill_nan(ndvi, time_split): + rs = [] + for times in time_split: + tmp = ndvi.sel(time=times) + fill_ds = tmp.sel(time=times).bfill(dim='time') + fill_ds = fill_ds.sel(time=times).ffill(dim='time') + rs.append(fill_ds) + merged_ndvi = xr.concat([i for i in rs], dim="time") + fill_m = merged_ndvi.bfill(dim="time") + fill_m = fill_m.ffill(dim="time") + return fill_m + + +def load_train_data(train_path): + train = load_data_geo(train_path) + return train + + +def load_sen1(name_vh, name_vv): + dsvv = rioxarray.open_rasterio(name_vv) + dsvh = rioxarray.open_rasterio(name_vh) + return dsvh, dsvv + + +def get_data_sen1_and_sen2(train, average_ndvi, dsvh, dsvv): + loaded_datasets = {} + for idx, point in train.iterrows(): + key = f"point_{idx + 1}" + try: + ndvi_data = average_ndvi.sel(x=point.geometry.x, y=point.geometry.y, method='nearest').values + vh_data = dsvh.sel(x=point.geometry.x, y=point.geometry.y, method='nearest').values + vv_data = dsvv.sel(x=point.geometry.x, y=point.geometry.y, method='nearest').values + loaded_datasets[key] = { + "data": np.concatenate((ndvi_data, vh_data, vv_data)), + "label": point.HT_code + } + except Exception as e: + # loaded_datasets[key] = None + print(e) + return loaded_datasets + + +def split_train_data(train, label_mapping, datasets): + label_encoder = LabelEncoder() + + # Fit and transform the labels + labels = train.Hientrang.values + numeric_labels = label_encoder.fit_transform([label_mapping[label] for label in labels]) + X = [] + x_new = [] + lb_new = [] + for k, v in datasets.items(): + X.append(v) + for i in range(len(X)): + if X[i] is not None: + x_new.append(X[i]["data"]) + lb_new.append(numeric_labels[i]) + X_train, X_temp, y_train, y_temp= train_test_split(x_new, lb_new, test_size=0.4, random_state=42) + X_val, X_test, y_val, y_test = train_test_split(X_temp, y_temp, test_size=0.5, random_state=42) + return X_train, X_val, X_test, y_train, y_val, y_test + + +def train_with_rf(X_train, X_val, y_train, y_val): + # Takes 1-2 minutes to complete + + # Tạo RandomForestClassifier mặc định để sử dụng làm mô hình ban đầu trong pipeline + base_model = RandomForestClassifier(random_state=42, n_jobs=-1) + + # Tạo pipeline + pipeline = Pipeline([ + # ('imputer', SimpleImputer(strategy='mean')), + ('scaler', StandardScaler()), + ('classifier', base_model), + ]) + # Thiết lập các tham số bạn muốn tối ưu hóa + param_grid = { + 'classifier__n_estimators': [100, 300, 500, 700, 1000], + 'classifier__max_depth': [6, 8, 10, 15, 20], + 'classifier__criterion': ['gini', 'entropy'], + } + + # Sử dụng GridSearchCV để tìm bộ tham số tốt nhất + grid_search = GridSearchCV(pipeline, param_grid, cv=5, scoring='accuracy', n_jobs=-1) + grid_search.fit(X_train, y_train) + + # In ra bộ tham số tốt nhất + best_params = grid_search.best_params_ + print("Best Parameters:", best_params) + + # Dự đoán trên tập kiểm tra + y_pred = grid_search.predict(X_val) + + # Đánh giá kết quả + accuracy = accuracy_score(y_val, y_pred) + print(f"Accuracy: {round(accuracy, 2)*100} %") + return grid_search + + +def save_model(name_file, grid_search): + dir_save_model = "model_train" + if not os.path.exists(dir_save_model): + os.mkdir(dir_save_model) + joblib.dump(grid_search, os.path.join(dir_save_model, name_file)) + print("Done!") + + +def predict(model, data_crs, ndvi, vh, vv): + data_predict = [] + for i in range(ndvi.shape[1]): + ndvi_tmp = ndvi.isel(y=i).values + vh_data = vh.sel(y=ndvi.y.values[i], method='nearest').values + vv_data = vv.sel(y=ndvi.y.values[i], method='nearest').values + all_tmp = np.concatenate((ndvi_tmp, vh_data, vv_data), axis=0) + data_predict.extend(all_tmp.T) + y_pred = model.predict(data_predict) + final_label = y_pred.reshape(ndvi.y.shape[0], ndvi.x.shape[0]) + + final_xarray_save = xr.DataArray(final_label, dims=("y", "x")) + final_xarray_save = final_xarray_save.rio.write_crs(data_crs) + + x_values = ndvi.x.values + y_values = ndvi.y.values + + data_array = xr.DataArray(final_xarray_save, + coords={'x': x_values, 'y': y_values}, + dims=['y', 'x']) + data_array = data_array.rio.write_crs(ndvi.rio.crs) + return data_array + + +def cut_according_shp(thuanhoa_path, average_ndvi, data_array): + gdf = gpd.read_file(thuanhoa_path) + gdf = gdf.to_crs(average_ndvi.rio.crs) + polygon_coords = list(gdf.geometry.values[0].exterior.coords) + polygon_coordinates = [(x, y) for x, y in polygon_coords] + + geometries = [ + { + 'type': 'Polygon', + 'coordinates': [polygon_coordinates] + } + ] + region_result = data_array.rio.clip(geometries, data_array.rio.crs, drop=False) + region_result = region_result.where(region_result >= 0, float('nan')) + return region_result + + +def compare(KD_path, KetQuaPhanLoaiDat, CODE_MAP, HT_MAP): + gdf = gpd.read_file(KD_path, crs="EPSG:9209") + polygon = gdf.geometry.values + label = gdf.tenchu.values + ouput_image = rioxarray.open_rasterio(KetQuaPhanLoaiDat) + code_tq = HT_MAP["TQ"]["data"][0] + code_pnn = HT_MAP["PNN"]["data"][0] + result = {} + for key, values in HT_MAP.items(): + print(f"process {key}") + array_list = [] + for i in range(len(polygon)): + po = polygon[i] + lb = label[i] + code_lb = CODE_MAP.get(lb, code_tq) + try: + qr = ouput_image.rio.clip([po], "EPSG:9209") + if code_lb in values["data"]: + if code_lb == code_pnn: + qr = qr.where((qr != float(code_pnn)), np.nan) + # qr = qr.where((qr != 3.0), np.nan) + elif code_lb == code_tq: + qr = qr.where((qr != float(code_pnn)), np.nan) + qr = qr.where((qr != 3.0), np.nan) + else: + qr = qr.where(qr != float(code_lb), np.nan) + else: + qr.values[:, :, :] = np.nan + array_list.append(qr) + except Exception as e: + pass + result.update({key: array_list}) + return result + + +def save_result(result, HT_MAP): + # cmap = ListedColormap(colors) + save_path = "ThuanHoa/KetQua" + if not os.path.exists(save_path): + os.mkdir(save_path) + + for k, v in result.items(): + rs = merge_arrays(v, nodata = np.nan) + rs.rio.to_raster(f"{save_path}/{k}.tif") + print(f"save {save_path}/{k}.tif") + # img = rs.plot(cmap=cmap, add_colorbar=False) + # cbar = plt.colorbar(img) + # cbar.ax.set_yticklabels(labels) + # plt.title(f'{HT_MAP[k]["name"]}') + # plt.axis('off') + # plt.show() \ No newline at end of file diff --git a/region/ST_region.dbf b/region/ST_region.dbf new file mode 100644 index 0000000..6782865 Binary files /dev/null and b/region/ST_region.dbf differ diff --git a/region/ST_region.prj b/region/ST_region.prj new file mode 100644 index 0000000..0ee7d78 --- /dev/null +++ b/region/ST_region.prj @@ -0,0 +1 @@ +GEOGCS["GCS_WGS_1984",DATUM["D_WGS_1984",SPHEROID["WGS_1984",6378137.0,298.257223563]],PRIMEM["Greenwich",0.0],UNIT["Degree",0.0174532925199433],AUTHORITY["EPSG",4326]] \ No newline at end of file diff --git a/region/ST_region.shp b/region/ST_region.shp new file mode 100644 index 0000000..bed95f1 Binary files /dev/null and b/region/ST_region.shp differ diff --git a/region/ST_region.shx b/region/ST_region.shx new file mode 100644 index 0000000..83e5a7b Binary files /dev/null and b/region/ST_region.shx differ diff --git a/train/ST_training data_updated_1130points.dbf b/train/ST_training data_updated_1130points.dbf new file mode 100644 index 0000000..85bfa7c Binary files /dev/null and b/train/ST_training data_updated_1130points.dbf differ diff --git a/train/ST_training data_updated_1130points.prj b/train/ST_training data_updated_1130points.prj new file mode 100644 index 0000000..0202b8e --- /dev/null +++ b/train/ST_training data_updated_1130points.prj @@ -0,0 +1 @@ +PROJCS["WGS_1984_UTM_Zone_48N",GEOGCS["GCS_WGS_1984",DATUM["D_WGS_1984",SPHEROID["WGS_1984",6378137,298.257223563]],PRIMEM["Greenwich",0],UNIT["Degree",0.017453292519943295]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105],PARAMETER["scale_factor",0.9996],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["Meter",1]] \ No newline at end of file diff --git a/train/ST_training data_updated_1130points.qml b/train/ST_training data_updated_1130points.qml new file mode 100644 index 0000000..6bda571 --- /dev/null +++ b/train/ST_training data_updated_1130points.qml @@ -0,0 +1,835 @@ + + + + 1 + 1 + 1 + 0 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + 0 + 0 + 1 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + 0 + + + 0 + generatedlayout + + + + + + + + + + + + + + + + + + + + + + + + + + + "No" + + 0 + diff --git a/train/ST_training data_updated_1130points.qpj b/train/ST_training data_updated_1130points.qpj new file mode 100644 index 0000000..18131bc --- /dev/null +++ b/train/ST_training data_updated_1130points.qpj @@ -0,0 +1 @@ +PROJCS["WGS 84 / UTM zone 48N",GEOGCS["WGS 84",DATUM["WGS_1984",SPHEROID["WGS 84",6378137,298.257223563,AUTHORITY["EPSG","7030"]],AUTHORITY["EPSG","6326"]],PRIMEM["Greenwich",0,AUTHORITY["EPSG","8901"]],UNIT["degree",0.0174532925199433,AUTHORITY["EPSG","9122"]],AUTHORITY["EPSG","4326"]],PROJECTION["Transverse_Mercator"],PARAMETER["latitude_of_origin",0],PARAMETER["central_meridian",105],PARAMETER["scale_factor",0.9996],PARAMETER["false_easting",500000],PARAMETER["false_northing",0],UNIT["metre",1,AUTHORITY["EPSG","9001"]],AXIS["Easting",EAST],AXIS["Northing",NORTH],AUTHORITY["EPSG","32648"]] diff --git a/train/ST_training data_updated_1130points.shp b/train/ST_training data_updated_1130points.shp new file mode 100644 index 0000000..050307b Binary files /dev/null and b/train/ST_training data_updated_1130points.shp differ diff --git a/train/ST_training data_updated_1130points.shx b/train/ST_training data_updated_1130points.shx new file mode 100644 index 0000000..18f235d Binary files /dev/null and b/train/ST_training data_updated_1130points.shx differ diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..57b36fc --- /dev/null +++ b/utils.py @@ -0,0 +1,6 @@ +import geopandas as gpd + + +def load_data_geo(path: str): + gdf = gpd.read_file(path) + return gdf \ No newline at end of file