import pandas as pd import numpy as np import plotly.graph_objects as go import plotly.express as px import xarray as xr import solara from uuid import uuid4 from sklearn.neighbors import NearestNeighbors from scipy.stats import linregress import geopandas as gpd from constants import REGION_COLUMN from Datasets import DATASETS def get_plotly_theme(): return "plotly_dark" if solara.lab.theme.dark_effective else "plotly_white" def load_dataset(dataset_config): dataset = xr.open_dataset(dataset_config["average_path"], decode_timedelta=False) variables = list(dataset.data_vars) float_vars = [var for var in variables if dataset[var].dtype in ["float32", "float64"]] data = dataset.to_dataframe().reset_index().dropna() data = data.rename(columns={"latitude": "lat", "longitude": "lon"}) if len(data) >= 4: sample_indices = np.random.choice(len(data), size=4, replace=False) sample_coords = data[["lat", "lon"]].iloc[sample_indices].values coords = data[["lat", "lon"]].values knn = NearestNeighbors(n_neighbors=2).fit(coords) distances, _ = knn.kneighbors(sample_coords, n_neighbors=2) grid_spacing = distances[:, 1].mean() else: grid_spacing = 0.25 dataset.close() return data, float_vars, grid_spacing def calculate_map_zoom(lat_range, lon_range): range_max = max(lat_range, lon_range) return (1 if range_max > 100 else 2 if range_max > 30 else 4 if range_max > 15 else 5 if range_max > 5 else 7 if range_max > 1 else 9 if range_max > 0.1 else 10) def create_heatmap(data, dataset_name, value_column, grid_size, color_scale, map_style, opacity): center = {"lat": data["lat"].mean(), "lon": data["lon"].mean()} zoom = calculate_map_zoom(data["lat"].max() - data["lat"].min(), data["lon"].max() - data["lon"].min()) features = [ { "type": "Feature", "id": str(uuid4()), "geometry": { "type": "Polygon", "coordinates": [[ [row["lon"] - grid_size/2, row["lat"] - grid_size/2], [row["lon"] + grid_size/2, row["lat"] - grid_size/2], [row["lon"] + grid_size/2, row["lat"] + grid_size/2], [row["lon"] - grid_size/2, row["lat"] + grid_size/2], [row["lon"] - grid_size/2, row["lat"] - grid_size/2] ]] }, "properties": {"value": row[value_column], "lat": row["lat"], "lon": row["lon"]} } for _, row in data.iterrows() ] geojson = {"type": "FeatureCollection", "features": features} fig = go.Figure( go.Choroplethmap( geojson=geojson, locations=[f["id"] for f in features], z=[f["properties"]["value"] for f in features], colorscale=color_scale, marker_opacity=opacity, marker_line_width=0, colorbar=dict(title=dataset_name), customdata=np.stack(([f["properties"]["lat"] for f in features], [f["properties"]["lon"] for f in features]), axis=-1), hovertemplate="Lat: %{customdata[0]:.3f}
Lon: %{customdata[1]:.3f}
Value: %{z}" ) ) fig.update_layout( map=dict(style=map_style, center=center, zoom=zoom), height=900, width=1250, title=f"Heatmap - {dataset_name}", template=get_plotly_theme() ) return fig def create_timeseries_plot(file_path, dataset_name, lat, lon, plotted_line_color, spatial_mean_data=None, show_spatial_mean=False, from_year=None, to_year=None, use_timeframe=False, region="India", show_trend_line=False, trend_line_color="red", spatial_mean_color="yellow"): try: dataset = xr.open_dataset(file_path, decode_timedelta=False) variable = list(dataset.data_vars)[0] lat_idx = np.abs((dataset.get("lat", dataset.get("latitude")).values - lat)).argmin() lon_idx = np.abs((dataset.get("lon", dataset.get("longitude")).values - lon)).argmin() timeseries = dataset[variable].isel(lat=lat_idx, lon=lon_idx).to_dataframe().reset_index() years = timeseries["year"].values if len(years) == 0: raise ValueError("No years found in the dataset") start_year = int(min(years)) end_year = int(max(years)) if use_timeframe and from_year is not None and to_year is not None: timeseries = timeseries[(timeseries["year"] >= from_year) & (timeseries["year"] <= to_year)] start_year = from_year end_year = to_year x = timeseries["year"].values y = timeseries[variable].values # Fit parabolic trend to actual data if show_trend_line and len(x) > 2: x_norm = x - x.mean() coeffs = np.polyfit(x_norm, y, deg=2) poly = np.poly1d(coeffs) trend_line = poly(x_norm) midpoint_idx = len(x) // 2 slope = 2 * coeffs[0] * x_norm[midpoint_idx] + coeffs[1] _, _, r_value, p_value, _ = linregress(x, y) trend = "Increasing" if p_value < 0.05 and slope > 0 else "Decreasing" if p_value < 0.05 and slope < 0 else "No Significant Trend" else: trend_line = np.full_like(x, np.nan) slope = 0 p_value = 1.0 trend = "No Significant Trend" print(f"Creating timeseries plot: dataset_name={dataset_name}, start_year={start_year}, end_year={end_year}, lat={lat:.2f}, lon={lon:.2f}, region={region}") fig = px.line( timeseries, x="year", y=variable, title=f"{dataset_name}, {start_year}-{end_year} for Lat: {lat:.2f}, Lon: {lon:.2f}", labels={"year": "Year", variable: dataset_name}, color_discrete_sequence=[plotted_line_color], render_mode="svg" ) fig.data[0].name = dataset_name if show_trend_line and len(x) > 2: fig.add_scatter( x=x, y=trend_line, mode="lines", name=f"{dataset_name} Trend", line=dict(color=trend_line_color, dash="dash") ) if show_spatial_mean and spatial_mean_data is not None: spatial_mean_df = spatial_mean_data.get("timeseries", pd.DataFrame()) if not spatial_mean_df.empty: fig.add_scatter( x=spatial_mean_df["year"], y=spatial_mean_df[variable], mode="lines", name=f"{region} Spatial Mean", line=dict(color=spatial_mean_color, width=2) ) fig.add_annotation( x=timeseries["year"].max(), y=timeseries[variable].min(), text=f"Slope: {slope:.4f}
p-value: {p_value:.4f}
Trend: {trend}", showarrow=False, xanchor="right", yanchor="top", bgcolor="white" if not solara.lab.theme.dark_effective else "#333333", bordercolor="black", borderpad=4, opacity=0.8, font=dict(size=12, color="black" if not solara.lab.theme.dark_effective else "white") ) fig.update_layout( xaxis_title="Year", yaxis_title=dataset_name, height=550, width=650, showlegend=True, legend=dict( x=0.99, y=0.99, xanchor="right", yanchor="top", bgcolor="rgba(0, 0, 0, 0)" ), template=get_plotly_theme() ) dataset.close() return fig except Exception as e: print(f"Error creating timeseries plot: {e}") dataset.close() if 'dataset' in locals() else None return go.Figure() def get_region_bounds(geodataframe, region): geometry = geodataframe[geodataframe[REGION_COLUMN] == region].geometry.iloc[0] minx, miny, maxx, maxy = geometry.bounds return {"min_lon": minx, "max_lon": maxx, "min_lat": miny, "max_lat": maxy} def filter_region_data(data, geodataframe, region): if not region or region == "India": return data try: points = gpd.GeoDataFrame( data, geometry=gpd.points_from_xy(data["lon"], data["lat"]), crs=geodataframe.crs ) points.sindex region_geometry = geodataframe[geodataframe[REGION_COLUMN] == region].geometry.iloc[0] bbox = region_geometry.bounds candidates = points[points.geometry.bounds.minx >= bbox[0]] candidates = candidates[candidates.geometry.bounds.maxx <= bbox[2]] candidates = candidates[candidates.geometry.bounds.miny >= bbox[1]] candidates = candidates[candidates.geometry.bounds.maxy <= bbox[3]] filtered_points = candidates[candidates.geometry.within(region_geometry)] return filtered_points.drop(columns=["geometry"]) except Exception as e: print(f"Error filtering region data: {e}") return pd.DataFrame() def calculate_spatial_mean(dataset_config, geodataframe, region, from_year=None, to_year=None, use_timeframe=False): try: dataset = xr.open_dataset(dataset_config["timeseries_path"], decode_timedelta=False) variable = list(dataset.data_vars)[0] no_data_value = dataset_config["no_data"] data = dataset[variable].to_dataframe().reset_index() data = data.rename(columns={"latitude": "lat", "longitude": "lon"}) data = data[data[variable] != no_data_value] if use_timeframe and from_year is not None and to_year is not None: data = data[(data["year"] >= from_year) & (data["year"] <= to_year)] if region != "India" and geodataframe is not None: data = filter_region_data(data, geodataframe, region) spatial_means = data.groupby("year")[variable].mean().reset_index() overall_mean = spatial_means[variable].mean() dataset.close() return {"overall_mean": overall_mean, "timeseries": spatial_means} except Exception as e: print(f"Error calculating spatial mean: {e}") return {"overall_mean": None, "timeseries": pd.DataFrame()} def calculate_selected_spatial_mean(dataset_config, selected_coords, from_year=None, to_year=None, use_timeframe=False): try: if not selected_coords: return {"overall_mean": None, "timeseries": pd.DataFrame()} dataset = xr.open_dataset(dataset_config["timeseries_path"], decode_timedelta=False) variable = list(dataset.data_vars)[0] no_data_value = dataset_config["no_data"] data_list = [] for lat, lon in selected_coords: lat_idx = np.abs((dataset.get("lat", dataset.get("latitude")).values - lat)).argmin() lon_idx = np.abs((dataset.get("lon", dataset.get("longitude")).values - lon)).argmin() point_data = dataset[variable].isel(lat=lat_idx, lon=lon_idx).to_dataframe().reset_index() data_list.append(point_data) data = pd.concat(data_list, ignore_index=True) data = data[data[variable] != no_data_value] if use_timeframe and from_year is not None and to_year is not None: data = data[(data["year"] >= from_year) & (data["year"] <= to_year)] spatial_means = data.groupby("year")[variable].mean().reset_index() overall_mean = spatial_means[variable].mean() dataset.close() return {"overall_mean": overall_mean, "timeseries": spatial_means} except Exception as e: print(f"Error calculating selected spatial mean: {e}") return {"overall_mean": None, "timeseries": pd.DataFrame()} def create_selected_timeseries_plot(file_path, dataset_name, selected_coords, plotted_line_color, selected_spatial_mean_data=None, from_year=None, to_year=None, use_timeframe=False, show_trend_line=False, trend_line_color="red", spatial_mean_color="yellow"): try: if not selected_coords: return go.Figure() dataset = xr.open_dataset(file_path, decode_timedelta=False) variable = list(dataset.data_vars)[0] no_data_value = next(d for d in DATASETS if d["name"] == dataset_name)["no_data"] data_list = [] for lat, lon in selected_coords: lat_idx = np.abs((dataset.get("lat", dataset.get("latitude")).values - lat)).argmin() lon_idx = np.abs((dataset.get("lon", dataset.get("longitude")).values - lon)).argmin() point_data = dataset[variable].isel(lat=lat_idx, lon=lon_idx).to_dataframe().reset_index() data_list.append(point_data) data = pd.concat(data_list, ignore_index=True) data = data[data[variable] != no_data_value] timeseries = data.groupby("year")[variable].mean().reset_index() years = timeseries["year"].values if len(years) == 0: raise ValueError("No years found in the dataset") start_year = int(min(years)) end_year = int(max(years)) if use_timeframe and from_year is not None and to_year is not None: timeseries = timeseries[(timeseries["year"] >= from_year) & (timeseries["year"] <= to_year)] start_year = from_year end_year = to_year x = timeseries["year"].values y = timeseries[variable].values # Fit parabolic trend to actual data if show_trend_line and len(x) > 2: x_norm = x - x.mean() coeffs = np.polyfit(x_norm, y, deg=2) poly = np.poly1d(coeffs) trend_line = poly(x_norm) midpoint_idx = len(x) // 2 slope = 2 * coeffs[0] * x_norm[midpoint_idx] + coeffs[1] _, _, r_value, p_value, _ = linregress(x, y) trend = "Increasing" if p_value < 0.05 and slope > 0 else "Decreasing" if p_value < 0.05 and slope < 0 else "No Significant Trend" else: trend_line = np.full_like(x, np.nan) slope = 0 p_value = 1.0 trend = "No Significant Trend" fig = px.line( timeseries, x="year", y=variable, title=f"{dataset_name}, Spatial Mean of Selected Points, {start_year}-{end_year}", labels={"year": "Year", variable: dataset_name}, color_discrete_sequence=[plotted_line_color] ) fig.data[0].name = dataset_name if show_trend_line and len(x) > 2: fig.add_scatter( x=x, y=trend_line, mode="lines", name=f"{dataset_name} Trend", line=dict(color=trend_line_color, dash="dash") ) if selected_spatial_mean_data is not None: spatial_mean_df = selected_spatial_mean_data.get("timeseries", pd.DataFrame()) if not spatial_mean_df.empty: fig.add_scatter( x=spatial_mean_df["year"], y=spatial_mean_df[variable], mode="lines", name="Selected Points Spatial Mean", line=dict(color=spatial_mean_color, width=2) ) fig.add_annotation( x=timeseries["year"].max(), y=timeseries[variable].min(), text=f"Slope: {slope:.4f}
p-value: {p_value:.4f}
Trend: {trend}", showarrow=False, xanchor="right", yanchor="top", bgcolor="white" if not solara.lab.theme.dark_effective else "#333333", bordercolor="black", borderpad=4, opacity=0.8, font=dict(size=12, color="black" if not solara.lab.theme.dark_effective else "white") ) fig.update_layout( xaxis_title="Year", yaxis_title=f"Mean {dataset_name}", height=550, width=650, showlegend=True, legend=dict( x=0.99, y=0.99, xanchor="right", yanchor="top", bgcolor="rgba(0, 0, 0, 0)" ), template=get_plotly_theme() ) dataset.close() return fig except Exception as e: print(f"Error creating selected timeseries plot: {e}") dataset.close() if 'dataset' in locals() else None return go.Figure() def get_available_years(dataset_config): try: dataset = xr.open_dataset(dataset_config["timeseries_path"], decode_timedelta=False) years = dataset["year"].values.tolist() dataset.close() return sorted(years) except Exception as e: print(f"Error getting available years: {e}") return []