Skip to content
Merged
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
234 changes: 148 additions & 86 deletions src/echodataflow/flows/flows_viz_cloud.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,15 @@
from pathlib import Path
import datetime
import configparser
import tempfile

import pandas as pd
import numpy as np
import xarray as xr
import s3fs

import echoregions as er

from prefect import flow, get_run_logger

from echodataflow.utils.utils import get_slice_start_end_times
Expand All @@ -31,7 +34,9 @@ def flow_update_cache_MVBS(
logger = get_run_logger()

# Set end_time to current time - time_offset_seconds
end_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(seconds=time_offset_seconds)
end_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(
seconds=time_offset_seconds
)

logger.info(
"flow started with parameters:\n"
Expand Down Expand Up @@ -68,15 +73,13 @@ def flow_update_cache_MVBS(
# Accept legacy naive values and offset-qualified values in the same file
if not df_MVBS.empty:
for column in ["first_ping_time", "last_ping_time"]:
df_MVBS[column] = pd.to_datetime(
df_MVBS[column], format="mixed", utc=True
)
df_MVBS[column] = pd.to_datetime(df_MVBS[column], format="mixed", utc=True)

# Get MVBS files in the specified time range (only 1 slice)
MVBS_filenames = sorted(
df_MVBS[
(df_MVBS["last_ping_time"] >= start_time[0]) &
(df_MVBS["first_ping_time"] <= end_time[0])
(df_MVBS["last_ping_time"] >= start_time[0])
& (df_MVBS["first_ping_time"] <= end_time[0])
]["MVBS_filename"].tolist()
)
logger.info(
Expand All @@ -89,17 +92,15 @@ def flow_update_cache_MVBS(
return
else:
# Assmeble fs mapper for the MVBS files
MVBS_filenames = [
fs.get_mapper(str(Path(path_MVBS) / mvbsf)) for mvbsf in MVBS_filenames
]
MVBS_filenames = [fs.get_mapper(str(Path(path_MVBS) / mvbsf)) for mvbsf in MVBS_filenames]

# Combine and prepare MVBS dataset
ds_MVBS = xr.open_mfdataset(
MVBS_filenames,
parallel=True,
coords="minimal",
data_vars="minimal",
compat='override',
compat="override",
chunks={"channel": -1, "ping_time": -1, "depth": -1}, # load everything into 1 chunk
engine="zarr", # use zarr engine for reading
consolidated=False,
Expand All @@ -110,8 +111,10 @@ def flow_update_cache_MVBS(

# Add actual_range to allow using holoviz
ds_MVBS["Sv"] = ds_MVBS["Sv"].assign_attrs(
actual_range=(float(ds_MVBS["Sv"].min().compute()),
float(ds_MVBS["Sv"].max().compute()))
actual_range=(
float(ds_MVBS["Sv"].min().compute()),
float(ds_MVBS["Sv"].max().compute()),
)
)

# Remove chunk encoding to prevent saving issues
Expand All @@ -123,14 +126,120 @@ def flow_update_cache_MVBS(

# Save to cache
logger.info(f"Saving MVBS dataset to cache: {str(Path(path_cache) / file_MVBS_zarr)}")
ds_MVBS.chunk(
{"channel": -1, "ping_time": -1, "echo_range": -1}
).to_zarr(
ds_MVBS.chunk({"channel": -1, "ping_time": -1, "echo_range": -1}).to_zarr(
Path(path_cache) / file_MVBS_zarr, # cache is local
mode="w",
consolidated=True,
)


@flow()
def flow_update_cache_contours(
time_offset_seconds: float = 0.0,
slice_mins: int = 180,
path_cache: str = "PATH_TO_DATA_CACHE",
path_EVR: str = "PATH_TO_EVR_DATA_STORE",
cred_file: str = "PATH_TO_CREDENTIALS_FILE",
file_contours_csv: str = "latest_contours.csv",
):
logger = get_run_logger()

# Set end time
end_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(
seconds=time_offset_seconds
)

logger.info(
"flow started with parameters:\n"
f"- end_time: {end_time}\n"
f"- slice_mins: {slice_mins}\n"
)

# Compute slice time range
start_time, end_time = get_slice_start_end_times(
end_time=end_time,
slice_mins=slice_mins,
num_slices=1,
)

# Get cloud bucket
config = configparser.ConfigParser()
config.read(cred_file)
fs = s3fs.S3FileSystem(
key=config["osn_sdsc_hake"]["access_key_id"],
secret=config["osn_sdsc_hake"]["secret_access_key"],
client_kwargs={"endpoint_url": config["osn_sdsc_hake"]["endpoint"]},
)

# Find EVR files
evr_files = fs.glob(f"{path_EVR}/*.evr")
selected_evr = []
for evr_file in evr_files:
filename = Path(evr_file).stem
# Example filename:
# prediction_20260710T182000
timestamp = datetime.datetime.strptime(
filename.split("_")[-1],
"%Y%m%dT%H%M%S",
).replace(tzinfo=datetime.timezone.utc)
# Subset for time range
if start_time[0] <= timestamp <= end_time[0]:
selected_evr.append((timestamp, evr_file))
selected_evr.sort(key=lambda x: x[0])

logger.info(
f"Found {len(selected_evr)} EVR files in time range:\n"
+ "".join([f"- {Path(f).name} ({t})\n" for t, f in selected_evr])
)

if len(selected_evr) == 0:
logger.info("Contours cache not updated: no EVR files in specified time range")
else:
# Read EVRs and collect dataframes
contours_dfs = []

with tempfile.TemporaryDirectory() as tmp_dir:
tmp_path = Path(tmp_dir)

for _, evr_file in selected_evr:
logger.info(f"Downloading EVR: {evr_file}")

local_evr = tmp_path / Path(evr_file).name

fs.get(
evr_file,
str(local_evr),
)

logger.info(f"Reading local EVR: {local_evr}")

try:
regions = er.read_evr(str(local_evr))

contours_dfs.append(regions.data)

finally:
if local_evr.exists():
local_evr.unlink()
logger.info(f"Removed temporary EVR: {local_evr}")

# Merge all regions
df_contours = pd.concat(
contours_dfs,
ignore_index=True,
)

logger.info(f"Merged and saving {len(df_contours)} contour regions")

# Save CSV cache
output_file = Path(path_cache) / file_contours_csv

df_contours.to_csv(
output_file,
index=False,
)


def _prepare_sv_for_echogram(
ds: xr.Dataset,
var_name: str = "Sv_masked",
Expand All @@ -157,11 +266,7 @@ def _prepare_sv_for_echogram(
else:
vertical = ds["range_sample"]

reduce_dims = [
dim
for dim in ["channel", "ping_time"]
if dim in vertical.dims
]
reduce_dims = [dim for dim in ["channel", "ping_time"] if dim in vertical.dims]

if reduce_dims:
vertical = vertical.median(
Expand All @@ -183,19 +288,13 @@ def _prepare_sv_for_echogram(
else:
valid_length = vertical_values.size

plot_ds = plot_ds.isel(
range_sample=slice(0, valid_length)
)
plot_ds = plot_ds.isel(range_sample=slice(0, valid_length))

vertical_values = vertical_values[:valid_length]

if (
vertical_values.size == 0
or not np.isfinite(vertical_values).all()
):
if vertical_values.size == 0 or not np.isfinite(vertical_values).all():
raise ValueError(
"Could not construct a finite vertical coordinate "
"for Sv visualization."
"Could not construct a finite vertical coordinate " "for Sv visualization."
)

plot_ds = plot_ds.assign_coords(
Expand All @@ -205,27 +304,15 @@ def _prepare_sv_for_echogram(
)
)

plot_ds = plot_ds.swap_dims(
{"range_sample": "echo_range"}
)
plot_ds = plot_ds.swap_dims({"range_sample": "echo_range"})

vmin = float(
plot_ds["Sv"].min(skipna=True).compute()
)
vmax = float(
plot_ds["Sv"].max(skipna=True).compute()
)
vmin = float(plot_ds["Sv"].min(skipna=True).compute())
vmax = float(plot_ds["Sv"].max(skipna=True).compute())

plot_ds["Sv"] = plot_ds["Sv"].assign_attrs(
actual_range=(vmin, vmax)
)
plot_ds["Sv"] = plot_ds["Sv"].assign_attrs(actual_range=(vmin, vmax))

masked_vmin = float(
plot_ds["Sv_masked"].min(skipna=True).compute()
)
masked_vmax = float(
plot_ds["Sv_masked"].max(skipna=True).compute()
)
masked_vmin = float(plot_ds["Sv_masked"].min(skipna=True).compute())
masked_vmax = float(plot_ds["Sv_masked"].max(skipna=True).compute())

plot_ds["Sv_masked"] = plot_ds["Sv_masked"].assign_attrs(
actual_range=(masked_vmin, masked_vmax)
Expand All @@ -240,6 +327,7 @@ def _prepare_sv_for_echogram(

return plot_ds


@flow()
def flow_update_cache_CPS(
path_CPS: str,
Expand Down Expand Up @@ -271,15 +359,10 @@ def flow_update_cache_CPS(
)

if transect is None:
print(
"CPS cache not updated: "
"no completed transects found in processing ledger."
)
print("CPS cache not updated: " "no completed transects found in processing ledger.")
return

transect_number = str(
transect["transect_part"]
)
transect_number = str(transect["transect_part"])

transect_start = pd.to_datetime(
transect["start_time"],
Expand All @@ -302,9 +385,7 @@ def flow_update_cache_CPS(
)

cps_paths = [
path_CPS / filename
for filename in cps_filenames
if (path_CPS / filename).exists()
path_CPS / filename for filename in cps_filenames if (path_CPS / filename).exists()
]

if not cps_paths:
Expand Down Expand Up @@ -345,19 +426,13 @@ def flow_update_cache_CPS(
return_index=True,
)

ds_CPS = ds_CPS.isel(
ping_time=np.sort(unique_idx)
)
ds_CPS = ds_CPS.isel(ping_time=np.sort(unique_idx))

# Ledger timestamps are UTC-aware; xarray ping_time is
# represented as timezone-naive datetime64.
transect_start_naive = (
transect_start.tz_convert(None)
)
transect_start_naive = transect_start.tz_convert(None)

transect_end_naive = (
transect_end.tz_convert(None)
)
transect_end_naive = transect_end.tz_convert(None)

ds_CPS = ds_CPS.sel(
ping_time=slice(
Expand All @@ -368,8 +443,7 @@ def flow_update_cache_CPS(

if ds_CPS.sizes.get("ping_time", 0) == 0:
print(
f"CPS cache not updated: no pings remain "
f"after slicing transect {transect_number}."
f"CPS cache not updated: no pings remain " f"after slicing transect {transect_number}."
)
return

Expand All @@ -385,32 +459,20 @@ def flow_update_cache_CPS(
ds_CPS.attrs.update(
{
"transect_number": transect_number,
"transect_start": str(
transect_start_naive
),
"transect_end": str(
transect_end_naive
),
"source_cps_files": len(
cps_paths
),
"transect_start": str(transect_start_naive),
"transect_end": str(transect_end_naive),
"source_cps_files": len(cps_paths),
}
)

# -----------------------------------------------------
# Save visualization cache
# -----------------------------------------------------

output_path = (
path_cache
/ file_CPS_zarr
)
output_path = path_cache / file_CPS_zarr

print(
f"Saving transect "
f"{transect_number} "
f"to CPS visualization cache: "
f"{output_path}"
f"Saving transect " f"{transect_number} " f"to CPS visualization cache: " f"{output_path}"
)

ds_CPS.chunk(
Expand All @@ -423,4 +485,4 @@ def flow_update_cache_CPS(
output_path,
mode="w",
consolidated=True,
)
)
Loading
Loading