from enum import Enum
import numpy as np
from numpy.typing import NDArray
from scipy.interpolate import RegularGridInterpolator, griddata, interp1d
from scipy.ndimage import gaussian_filter, gaussian_filter1d
from .plotting_constants import DEFAULT_SMOOTHING_CONFIG, SmoothingConfig
[docs]
class DataType3D(Enum):
"""Whether 3D plot data is given on a grid or as scattered points."""
GridInput = "GridInput"
ScatterInput = "ScatterInput"
[docs]
def prepare_smooth_data_2D(
x_list: NDArray[np.float64],
y_list: NDArray[np.float64],
smoothing_config: SmoothingConfig,
) -> tuple[NDArray[np.float64], NDArray[np.float64]]:
"""
Interpolate and smooth (x, y) data for 2D plots.
Parameters
----------
x_list : NDArray[float64]
X coordinates
y_list : NDArray[float64]
Y coordinates
smoothing_config : SmoothingConfig
Smoothing and interpolation settings
Returns
-------
x_dense : NDArray[float64]
Smoothed/interpolated x coordinates
y_dense : NDArray[float64]
Smoothed/interpolated y coordinates
"""
if smoothing_config.interp_factor > 1:
x_dense = np.linspace(
x_list.min(), x_list.max(), len(x_list) * smoothing_config.interp_factor
)
f = interp1d(
x_list,
y_list,
kind=smoothing_config.interp_method,
bounds_error=False,
fill_value="extrapolate",
)
y_dense = f(x_dense).astype(np.float64)
else:
x_dense, y_dense = x_list, y_list
if smoothing_config.smoothing_sigma > 0:
y_dense = gaussian_filter1d(y_dense, sigma=smoothing_config.smoothing_sigma)
return x_dense, y_dense
[docs]
def validate_data(
x_list: NDArray[np.float64], y_list: NDArray[np.float64], z_list: NDArray[np.float64]
) -> DataType3D:
"""
Determine if the input data is grid-based or scatter-based.
Parameters
----------
x_list : NDArray[float64]
X coordinates
y_list : NDArray[float64]
Y coordinates
z_list : NDArray[float64]
Z values
Returns
-------
DataType3D
Whether data is GridInput or ScatterInput
Raises
------
ValueError
If data dimensions are invalid or mismatched
"""
z_ndim = np.ndim(z_list)
if z_ndim == 2:
if z_list.shape != (len(x_list), len(y_list)):
raise ValueError(
f"Grid data dimension mismatch: z_list.shape {z_list.shape} "
f"does not match (len(x_list), len(y_list)) = ({len(x_list)}, {len(y_list)})"
)
return DataType3D.GridInput
elif z_ndim == 1:
if not (len(x_list) == len(y_list) == len(z_list)):
raise ValueError(
f"Scatter data length mismatch: len(x_list)={len(x_list)}, "
f"len(y_list)={len(y_list)}, len(z_list)={len(z_list)}. "
f"All must be equal for scatter data."
)
return DataType3D.ScatterInput
else:
raise ValueError(f"Invalid z_list dimensions: expected 1D or 2D array, got {z_ndim}D")
[docs]
def prepare_smooth_data_3D(
x_list: NDArray[np.float64],
y_list: NDArray[np.float64],
z_list: NDArray[np.float64],
smoothing_config: SmoothingConfig = DEFAULT_SMOOTHING_CONFIG,
) -> tuple[NDArray[np.float64], NDArray[np.float64], NDArray[np.float64]]:
"""
Interpolate and smooth grid-based data for 3D plots.
Requires grid-based input data.
Parameters
----------
x_list : NDArray[float64]
1D array of x coordinates defining the grid x-axis
y_list : NDArray[float64]
1D array of y coordinates defining the grid y-axis
z_list : NDArray[float64]
2D array of z values with shape (len(x_list), len(y_list))
smoothing_config : SmoothingConfig, optional
Smoothing and interpolation settings
Returns
-------
x_interp : NDArray[float64]
Interpolated x coordinates
y_interp : NDArray[float64]
Interpolated y coordinates
z_interp : NDArray[float64]
Smoothed and interpolated z values
"""
x_interp = np.linspace(
np.min(x_list), np.max(x_list), len(x_list) * smoothing_config.interp_factor
)
y_interp = np.linspace(
np.min(y_list), np.max(y_list), len(y_list) * smoothing_config.interp_factor
)
x_grid, y_grid = np.meshgrid(x_interp, y_interp, indexing="ij")
interp_func = RegularGridInterpolator(
(x_list, y_list),
z_list,
method=smoothing_config.interp_method,
bounds_error=False,
)
points = np.vstack([x_grid.ravel(), y_grid.ravel()]).T
z_interp = interp_func(points).reshape(x_grid.shape)
if smoothing_config.smoothing_sigma > 0:
z_interp = gaussian_filter(z_interp, sigma=smoothing_config.smoothing_sigma)
return x_interp, y_interp, z_interp.T
[docs]
def prepare_smooth_data_3D_scatter(
x_list: NDArray[np.float64],
y_list: NDArray[np.float64],
z_list: NDArray[np.float64],
smoothing_config: SmoothingConfig = DEFAULT_SMOOTHING_CONFIG,
) -> tuple[NDArray[np.float64], NDArray[np.float64], NDArray[np.float64]]:
"""
Interpolate scattered 3D data onto a regular grid.
Requires scatter-based input data (1D x, y, z arrays).
Parameters
----------
x_list : NDArray[float64]
1D array of x coordinates
y_list : NDArray[float64]
1D array of y coordinates
z_list : NDArray[float64]
1D array of z values
smoothing_config : SmoothingConfig, optional
Smoothing and interpolation settings
Returns
-------
x_interp : NDArray[float64]
Regular grid x coordinates
y_interp : NDArray[float64]
Regular grid y coordinates
z_interp : NDArray[float64]
Interpolated z values on regular grid suitable for surface/contour plots
"""
num_points = len(x_list)
grid_resolution = int(np.sqrt(num_points) * smoothing_config.interp_factor)
grid_resolution = max(grid_resolution, 50)
x_interp = np.linspace(np.min(x_list), np.max(x_list), grid_resolution)
y_interp = np.linspace(np.min(y_list), np.max(y_list), grid_resolution)
x_grid, y_grid = np.meshgrid(x_interp, y_interp, indexing="ij")
points = np.column_stack([x_list, y_list])
try:
z_interp = griddata(points, z_list, (x_grid, y_grid), method="cubic", fill_value=np.nan)
except:
z_interp = griddata(points, z_list, (x_grid, y_grid), method="linear", fill_value=np.nan)
if smoothing_config.smoothing_sigma > 0:
mask = ~np.isnan(z_interp)
if np.any(mask):
z_interp_smooth = z_interp.copy()
z_interp_smooth[mask] = gaussian_filter(
z_interp[mask].reshape(-1), sigma=smoothing_config.smoothing_sigma
)
z_interp = z_interp_smooth
return x_interp, y_interp, z_interp.T