Source code for across.tools.visibility.joint_visibility

from typing import Any, Generic, TypeVar
from uuid import UUID

import numpy as np
import plotly.graph_objects as go
from pydantic import Field, model_validator

from across.tools.core.plotting import plot_joint_visibility_windows, plot_visibility_windows

from ..core.enums import ConstraintType
from ..core.schemas import (
    AstropyDateTime,
    AstropyTimeDelta,
)
from .base import Visibility

[docs] T = TypeVar("T", bound=Visibility)
[docs] class JointVisibility(Visibility, Generic[T]): """ Computes joint visibility windows between multiple instruments. This class takes a list of Visibility objects with identical timestamp grids and computes the intersection of their visibility periods. Parameters ---------- visibilities : list[T] List of Visibility or Visibility child objects with identical timestamp grids. instrument_ids : list[UUID] List of IDs of the instruments belonging to the Visibility objects. """ # Parameters
[docs] visibilities: list[T] = Field(default_factory=list, exclude=True)
[docs] instrument_ids: list[UUID] = Field(default_factory=list, exclude=True)
# Values derived from input parameters
[docs] step_size: AstropyTimeDelta = Field(default=None)
[docs] begin: AstropyDateTime = Field(default=None)
[docs] end: AstropyDateTime = Field(default=None)
[docs] observatory_name: str = Field(default="", exclude=True)
@model_validator(mode="before") @classmethod
[docs] def validate_parameters(cls, values: dict[str, Any]) -> dict[str, Any]: """ Validate and synchronize coordinate, begin/end, and step size values. This method ensures that all the input Visibility objects have the same coordinates, begin and end times, and step sizes. If any of these values differ between Visibilities, a ValueError is raised. When these equalities have been validated, this method runs the base Visibility validate_parameters method to validate other parameter values. """ # Check that all ra/dec values are the same within tolerance (~15 arcsec) tolerance = 15.0 / 3600.0 ras = np.array([(visibility.ra) for visibility in values["visibilities"]], dtype=float) decs = np.array([(visibility.dec) for visibility in values["visibilities"]], dtype=float) if not (np.allclose(ras, ras[0], atol=tolerance) and np.allclose(decs, decs[0], atol=tolerance)): raise ValueError( f"All input visibilities must have the same coordinate within {tolerance} degrees" ) values["ra"] = values["visibilities"][0].ra values["dec"] = values["visibilities"][0].dec # Check that all begin/end values are the same if ( not len(set([visibility.begin for visibility in values["visibilities"]])) == 1 or not len(set([visibility.end for visibility in values["visibilities"]])) == 1 ): raise ValueError("All begin and end times must be the same") values["begin"] = values["visibilities"][0].begin values["end"] = values["visibilities"][0].end # Check that all step sizes are the same if not len(set([visibility.step_size for visibility in values["visibilities"]])) == 1: raise ValueError("All step sizes must be the same") values["step_size"] = values["visibilities"][0].step_size return values
[docs] def _constraint(self, i: int) -> ConstraintType: """ For a given index, return the constraint from the first visibility that is actually constrained. """ # Safely handle out-of-bounds indices if i < 0 or i >= len(self.timestamp) if self.timestamp else True: return ConstraintType.WINDOW # Find the first visibility that is constrained at this index for vis in self.visibilities: if vis.inconstraint[i]: return vis._constraint(i) return ConstraintType.UNKNOWN
[docs] def _get_id(self, i: int) -> UUID: """ For a given index, find the observatory ID of the first visibility that is constrained. """ # Safely handle out-of-bounds indices if self.timestamp is None or i < 0 or i >= len(self.timestamp): return self.visibilities[0]._get_id(i) if self.visibilities else UUID(int=0) for vis in self.visibilities: if vis.inconstraint[i]: return vis._get_id(i) # Unknown constraint found, so just return the first visibility's observatory ID return self.visibilities[0]._get_id(i) if self.visibilities else UUID(int=0)
[docs] def _get_name(self, i: int) -> str: """ For a given index, get the name of the first instrument that is constrained. """ # Safely handle out-of-bounds indices by using first visibility's name if self.timestamp is None or i < len(self.timestamp): for vis in self.visibilities: if i >= 0 and i < len(vis.inconstraint) and vis.inconstraint[i]: return vis.observatory_name # Return first visibility's name as fallback if self.visibilities: return self.visibilities[0].observatory_name return ""
[docs] def prepare_data(self) -> None: """ Compute joint visibility by ANDing all inconstraint arrays. Raises ------ ValueError If visibilities list is empty or if visibilities have different timestamp grids. """ if not self.visibilities: raise ValueError("No visibilities provided for joint visibility calculation") # Compute joint visibility by ORing all inconstraint arrays self.inconstraint = np.any([vis.inconstraint for vis in self.visibilities], axis=0)
[docs] def _merge_computed_values(self) -> None: """ Abstract method to merge computed values from constraints. """ for visibility in self.visibilities: if visibility.computed_values is not None: self.computed_values.merge(visibility.computed_values)
[docs] def plot( self, fig: go.Figure | None = None, offset: int | float = 0, width: int = 700, height: int = 1000 ) -> go.Figure: """ Method to visualize joint visibility windows using plotly. Plots the individual instrument visibility windows and the regions of joint visibility on one figure. Calls the across-tools plotting core functionality and configures the plot layout to user specifications. Parameters ---------- fig : go.Figure, optional An existing plotly figure to add to, by default None offset : int | float, optional The x-axis offset to plot new visibility windows, by default 0 width: int, optional The width of the plot, in pixels. Defaults to 700. height: int, optional The height of the plot, in pixels. Defaults to 1000. Returns ------- go.Figure The plotly figure containing the footprint plot """ if fig is None: fig = go.Figure() tickvals = [] ticktext = [] for i, visibility in enumerate(self.visibilities): fig = plot_visibility_windows( visibility_windows=[window.model_dump() for window in visibility.visibility_windows], observatory_name=visibility.observatory_name, fig=fig, offset=offset + i + 1, ) tickvals.append(offset + i + 1) ticktext.append(visibility.observatory_name) min_extent = min(tickvals) max_extent = max(tickvals) fig = plot_joint_visibility_windows( visibility_windows=[window.model_dump() for window in self.visibility_windows], min_extent=min_extent, max_extent=max_extent, fig=fig, ) fig.update_layout( title="Visibility Windows", yaxis=dict( title="Time (UTC)", range=[self.end.to_datetime(), self.begin.to_datetime()], # descending time type="date", autorange=False, # don't resize ), xaxis=dict( title="Visibility Windows", tickvals=tickvals, ticktext=ticktext, ), width=width, height=height, ) return fig
[docs] def compute_joint_visibility( visibilities: list[T], instrument_ids: list[UUID], min_vis: int = 0, ) -> JointVisibility[T]: """ Compute joint visibility windows for any number of instrument Visibilities. Assumes that the visibilities are in the same order as instrument_ids. Parameters ---------- visibilities: list[T] List of Visibility objects or children objects of Visibility. instrument_ids: list[UUID] List of IDs of the instruments belonging to the Visibility objects. min_vis: int, optional Minimum visibility duration in seconds for a window to be included. Default is 0. Returns ------- list[VisibilityWindow] List of VisibilityWindows capturing joint visibility between all inputted instruments. """ joint_vis = JointVisibility( visibilities=visibilities, instrument_ids=instrument_ids, min_vis=min_vis, ) joint_vis.compute() return joint_vis