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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 26 additions & 27 deletions echopype/metrics/summary_statistics.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,22 +13,22 @@
import xarray as xr


def delta_z(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def delta_z(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Helper function to calculate widths between range samples (dz) for discretized integral.

Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
if range_label not in ds:
raise ValueError(f"{range_label} not in the input Dataset!")
dz = ds[range_label].diff(dim="range_sample")
if range_var not in ds:
raise ValueError(f"{range_var} not in the input Dataset!")
dz = ds[range_var].diff(dim=range_var)
return dz.where(dz != 0, other=np.nan)


Expand All @@ -48,70 +48,69 @@ def convert_to_linear(ds: xr.Dataset, Sv_label="Sv") -> xr.DataArray:
return 10 ** (ds[Sv_label] / 10)


def abundance(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def abundance(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Calculates the area-backscattering strength (Sa) [unit: dB re 1 m^2 m^-2].

This quantity is the integral of volumetric backscatter over range (``echo_range``).

Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
dz = delta_z(ds, range_label=range_label)
dz = delta_z(ds, range_var=range_var)
sv = convert_to_linear(ds, "Sv")
return 10 * np.log10((sv * dz).sum(dim="range_sample"))
return 10 * np.log10((sv * dz).sum(dim=range_var))


def center_of_mass(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def center_of_mass(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Calculates the mean backscatter location [unit: m].

This quantity is the weighted average of backscatter along range (``echo_range``).

Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
dz = delta_z(ds, range_label=range_label)
dz = delta_z(ds, range_var=range_var)
sv = convert_to_linear(ds, "Sv")
return (ds[range_label] * sv * dz).sum(dim="range_sample") / (sv * dz).sum(dim="range_sample")
return (ds[range_var] * sv * dz).sum(dim=range_var) / (sv * dz).sum(dim=range_var)


def dispersion(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def dispersion(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Calculates the inertia (I) [unit: m^-2].

This quantity measures dispersion or spread of backscatter from the center of mass.

Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
dz = delta_z(ds, range_label=range_label)
dz = delta_z(ds, range_var=range_var)
sv = convert_to_linear(ds, "Sv")
cm = center_of_mass(ds)
return ((ds[range_label] - cm) ** 2 * sv * dz).sum(dim="range_sample") / (sv * dz).sum(
dim="range_sample"
)
# cm = center_of_mass(ds)
cm = center_of_mass(ds, range_var=range_var)
return ((ds[range_var] - cm) ** 2 * sv * dz).sum(dim=range_var) / (sv * dz).sum(dim=range_var)


def evenness(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def evenness(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Calculates the equivalent area (EA) [unit: m].

This quantity represents the area that would be occupied if all datacells
Expand All @@ -120,19 +119,19 @@ def evenness(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
dz = delta_z(ds, range_label=range_label)
dz = delta_z(ds, range_var=range_var)
sv = convert_to_linear(ds, "Sv")
return ((sv * dz).sum(dim="range_sample")) ** 2 / (sv**2 * dz).sum(dim="range_sample")
return ((sv * dz).sum(dim=range_var)) ** 2 / (sv**2 * dz).sum(dim=range_var)


def aggregation(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
def aggregation(ds: xr.Dataset, range_var="echo_range") -> xr.DataArray:
"""Calculated the index of aggregation (IA) [unit: m^-1].

This quantity is reciprocal of the equivalent area.
Expand All @@ -141,11 +140,11 @@ def aggregation(ds: xr.Dataset, range_label="echo_range") -> xr.DataArray:
Parameters
----------
ds : xr.Dataset
range_label : str
range_var : str
Name of an xarray DataArray in ``ds`` containing ``echo_range`` information.

Returns
-------
xr.DataArray
"""
return 1 / evenness(ds, range_label=range_label)
return 1 / evenness(ds, range_var=range_var)
26 changes: 12 additions & 14 deletions echopype/tests/metrics/test_metrics_summary_statistics.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,16 +19,14 @@
# Utility Function


def create_test_ds(Sv, echo_range):
def create_test_ds(Sv, echo_range, range_var="echo_range"):
freq = [30]
time = pd.date_range("2021-08-28", periods=2)
reference_time = pd.Timestamp("2021-08-27") # noqa: F841
r_b = [0, 1, 2]

testDS = xr.Dataset(
data_vars=dict(
Sv=(["frequency", "ping_time", "range_sample"], Sv),
echo_range=(["frequency", "ping_time", "range_sample"], echo_range),
Sv=(["frequency", "ping_time", range_var], Sv),
),
coords={
'frequency': xr.DataArray(
Expand All @@ -42,10 +40,10 @@ def create_test_ds(Sv, echo_range):
name='ping_time',
dims=['ping_time'],
),
'range_sample': xr.DataArray(
r_b,
name='range_sample',
dims=['range_sample'],
range_var: xr.DataArray(
echo_range,
name=range_var,
dims=[range_var],
),
},
)
Expand All @@ -58,7 +56,7 @@ def create_test_ds(Sv, echo_range):
def test_abundance():
"""Compares summary_statistics.py calculation of abundance with verified outcomes"""
Sv = np.array([[[20, 40, 60], [50, 20, 30]]])
echo_range = np.array([[[1, 2, 3], [2, 3, 4]]])
echo_range = np.array([1, 2, 3])

ab_ds1 = create_test_ds(Sv, echo_range)
ab_ds1_SOL = np.array([[60.04321374, 30.41392685]])
Expand All @@ -70,9 +68,9 @@ def test_abundance():
def test_center_of_mass():
"""Compares summary_statistics.py calculation of center_of_mass with verified outcomes"""
Sv = np.array([[[20, 40, 60], [50, 20, 30]]])
echo_range = np.array([[[1, 2, 3], [2, 3, 4]]])
echo_range = np.array([1, 2, 3])
cm_ds1 = create_test_ds(Sv, echo_range)
cm_ds1_SOL = np.array([[2.99009901, 3.90909090]])
cm_ds1_SOL = np.array([[2.99009901, 2.90909090]])
assert np.allclose(
center_of_mass(cm_ds1), cm_ds1_SOL, rtol=1e-09
), 'Calculated output does not match expected output'
Expand All @@ -81,7 +79,7 @@ def test_center_of_mass():
def test_inertia():
"""Compares summary_statistics.py calculation of inertia with verified outcomes"""
Sv = np.array([[[20, 40, 60], [50, 20, 30]]])
echo_range = np.array([[[1, 2, 3], [2, 3, 4]]])
echo_range = np.array([1, 2, 3])
in_ds1 = create_test_ds(Sv, echo_range)
in_ds1_SOL = np.array([[0.00980296, 0.08264463]])
assert np.allclose(
Expand All @@ -92,7 +90,7 @@ def test_inertia():
def test_evenness():
"""Compares summary_statistics.py calculation of evenness with verified outcomes"""
Sv = np.array([[[20, 40, 60], [50, 20, 30]]])
echo_range = np.array([[[1, 2, 3], [2, 3, 4]]])
echo_range = np.array([1, 2, 3])
ev_ds1 = create_test_ds(Sv, echo_range)
ev_ds1_SOL = np.array([[1.019998, 1.198019802]])
assert np.allclose(
Expand All @@ -103,7 +101,7 @@ def test_evenness():
def test_aggregation():
"""Compares summary_statistics.py calculation of aggregation with verified outcomes"""
Sv = np.array([[[20, 40, 60], [50, 20, 30]]])
echo_range = np.array([[[1, 2, 3], [2, 3, 4]]])
echo_range = np.array([1, 2, 3])
ag_ds1 = create_test_ds(Sv, echo_range)
ag_ds1_SOL = np.array([[0.9803940792, 0.8347107438]])
assert np.allclose(
Expand Down
Loading