River_Network / src /data /loaders /station_elevations.py
ageraustine's picture
Upload folder using huggingface_hub
a74054f verified
Raw
History Blame Contribute Delete
5.82 kB
"""
Station elevations data loader for DEM-derived station altitude and metadata.
Single Responsibility: Load physical station coordinates, elevations, and geographic metadata.
"""
from pathlib import Path
from typing import List, Optional, Dict, Any
import pandas as pd
import matplotlib.pyplot as plt
from .base import BaseDataLoader
class StationElevationsLoader(BaseDataLoader):
"""
Loads DEM-derived station elevations and geospatial coordinates.
"""
def __init__(
self,
data_path: Path,
station_ids: Optional[List[str]] = None
):
"""
Initialize station elevations loader.
Args:
data_path: Path to station_elevations.csv file
station_ids: Optional list of hydrometric station codes to filter
"""
super().__init__(data_path)
self.station_ids = station_ids
def load(self) -> pd.DataFrame:
"""
Load station elevation data and enforce clean dtypes.
Returns:
DataFrame containing [station_code, station_name, latitude,
longitude, elevation_m, data_source]
"""
df = pd.read_csv(self.data_path)
# Strip whitespace from string columns
str_cols = df.select_dtypes(include=['object']).columns
for col in str_cols:
df[col] = df[col].astype(str).str.strip()
# Ensure numeric casting for coordinates and elevation
numeric_cols = ['latitude', 'longitude', 'elevation_m']
for col in numeric_cols:
if col in df.columns:
df[col] = pd.to_numeric(df[col], errors='coerce')
# Filter by station IDs if specified
if self.station_ids and "station_code" in df.columns:
df = df[df["station_code"].isin(self.station_ids)].reset_index(drop=True)
return df
def get_elevation_map(self) -> Dict[str, float]:
"""
Convenience method to return a mapping of station_code -> elevation_m.
Returns:
Dict mapping station codes to elevation in meters.
"""
df = self.load()
return dict(zip(df["station_code"], df["elevation_m"]))
def get_coordinates_df(self) -> pd.DataFrame:
"""
Get standardized coordinates table for graph spatial layout.
Returns:
DataFrame with [station_code, lat, lon, alt]
"""
df = self.load()
coords_df = df.rename(columns={
'latitude': 'lat',
'longitude': 'lon',
'elevation_m': 'alt'
})[['station_code', 'lat', 'lon', 'alt']]
return coords_df
def get_metadata(self) -> Dict[str, Any]:
"""Get station elevations metadata."""
meta = super().get_metadata()
meta.update({
"data_type": "station_elevations",
"source_organization": "EU-DEM 25m (Open Topo Data)"
})
return meta
def plot_elevation_bar(
self,
df: Optional[pd.DataFrame] = None,
figsize: tuple = (10, 6),
save_path: Optional[Path] = None,
) -> plt.Axes:
"""
Bar chart of station elevations, sorted low to high. Quick way to
see the elevation range and spot outlier stations.
Args:
df: Optional pre-loaded DataFrame. If not provided, data is loaded from disk.
figsize: Figure size in inches.
save_path: If provided, saves the figure to this path.
Returns:
The matplotlib Axes object.
"""
if df is None:
df = self.load()
plot_df = df.dropna(subset=["elevation_m"]).sort_values("elevation_m")
fig, ax = plt.subplots(figsize=figsize)
labels = plot_df.get("station_name", plot_df["station_code"])
ax.barh(labels.astype(str), plot_df["elevation_m"], color="sienna")
ax.set_title("Station Elevations")
ax.set_xlabel("Elevation (m)")
ax.set_ylabel("Station")
ax.grid(True, alpha=0.3, axis="x")
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches="tight")
print(f"Plot saved to {save_path}")
return ax
def plot_station_map(
self,
df: Optional[pd.DataFrame] = None,
figsize: tuple = (8, 8),
save_path: Optional[Path] = None,
) -> plt.Axes:
"""
Scatter map of station locations, colored by elevation.
Args:
df: Optional pre-loaded DataFrame. If not provided, data is loaded from disk.
figsize: Figure size in inches.
save_path: If provided, saves the figure to this path.
Returns:
The matplotlib Axes object.
"""
if df is None:
df = self.load()
plot_df = df.dropna(subset=["latitude", "longitude"])
fig, ax = plt.subplots(figsize=figsize)
scatter = ax.scatter(
plot_df["longitude"], plot_df["latitude"],
c=plot_df["elevation_m"], cmap="terrain", s=80,
edgecolor="black", linewidth=0.5,
)
for _, row in plot_df.iterrows():
label = row.get("station_name", row["station_code"])
ax.annotate(str(label), (row["longitude"], row["latitude"]),
fontsize=7, xytext=(3, 3), textcoords="offset points")
cbar = plt.colorbar(scatter, ax=ax)
cbar.set_label("Elevation (m)")
ax.set_title("Station Locations")
ax.set_xlabel("Longitude")
ax.set_ylabel("Latitude")
ax.set_aspect("equal")
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches="tight")
print(f"Plot saved to {save_path}")
return ax