diff --git a/icecube_data_reader/event_types.py b/icecube_data_reader/event_types.py index 1a24c3d..88f4da7 100644 --- a/icecube_data_reader/event_types.py +++ b/icecube_data_reader/event_types.py @@ -18,6 +18,9 @@ def __str__(self): def __eq__(self, other): return self.S == other.S + + def __int__(self): + return self.S @dataclass(eq=False) diff --git a/icecube_data_reader/events.py b/icecube_data_reader/events.py index d670806..8ba5378 100644 --- a/icecube_data_reader/events.py +++ b/icecube_data_reader/events.py @@ -30,7 +30,6 @@ EventType, Refrigerator, ) -from icecube_data_reader.lifetime import LifeTime from typing import Self import logging @@ -62,11 +61,11 @@ def coords(self): @property def ra(self): - return self._ra + return self.coords.ra @property def dec(self): - return self._dec + return self.coords.dec @property def unit_vectors(self): @@ -151,13 +150,35 @@ def scramble_ra(self, seed: int = 42) -> None: """ logger.warning("Scrambling RA. To revert this operation reload the events") - rng = np.default_rng(seed=seed) + rng = np.random.default_rng(seed=seed) ra = rng.random(self.ra.size) * 2 * np.pi * u.rad - self.ra = ra.to(u.deg) - self.coords = SkyCoord(ra=self.ra, dec=self.dec, frame="icrs") + self._coords = SkyCoord(ra=ra, dec=self.dec, frame="icrs") + + def scramble_mjd(self, seed: int = 42) -> None: + """Scrambles event mjd. + + :param seed: random seed, defaults to 42 + :type seed: int, optional + :returns: None + """ + + from icecube_data_reader.lifetime import IceTrackDR2LifeTime as LifeTime + + dm_set, inverse, counts = np.unique(self.int_types, return_inverse=True, return_counts=True) + new_mjd = np.zeros(self.N) + counts_dict = {} + for dm, count in zip(dm_set, counts): + counts_dict[Refrigerator.int2dm(dm)] = count + lifetime = LifeTime() + draw_mjd = lifetime.draw_event_mjd(counts_dict) + for c, dm in enumerate(dm_set): + new_mjd[inverse==c] = draw_mjd[Refrigerator.int2dm(dm)].mjd + + self._mjd = Time(new_mjd, format="mjd") + + + - def scramble_mjd(self, lifetime: LifeTime, seed: int = 42) -> None: - pass def select(self, mask: npt.NDArray[np.bool_]) -> None: """ diff --git a/icecube_data_reader/lifetime.py b/icecube_data_reader/lifetime.py index 20ff5cb..18268ca 100644 --- a/icecube_data_reader/lifetime.py +++ b/icecube_data_reader/lifetime.py @@ -34,6 +34,10 @@ class LifeTime(ABC): @property def data(self): return self._data + + @property + def dists(self): + return self._dists def lifetime_from_mjd( self, mjd_min: float, mjd_max: float, squeeze: bool = True @@ -61,7 +65,7 @@ def lifetime_from_mjd( # Query histograms for fraction of total lifetime in a season # multiply by total lifetime in a season to get appropriate value time = ( - self._dists[s].cdf(mjd_max) - self._dists[s].cdf(mjd_min) + self.dists[s].cdf(mjd_max) - self.dists[s].cdf(mjd_min) ) * self._lifetimes[s] # set atol to 1e-9 days, so we are below the time resolution of event mjd (1e-8 days) if squeeze and np.isclose(time.to_value(u.d), 0.0, atol=1e-9): @@ -106,6 +110,24 @@ def lifetime_from_season( output[s] = self._lifetimes[s] return output + + def draw_event_mjd( + self, event_numbers: dict[EventType, int] + ) -> dict[EventType, Time]: + """Draw event MJDs for scrambling the arrival times of events + TODO: add min/max mjd + + :param event_numbers: Dict of event type and requested numbers + :type event_numbers: dict[EventType, int] + :return: Dictionary of event types and new MJDs + :rtype: dict[EventType, Time] + """ + + out = {} + for dm, num in event_numbers.items(): + out[dm] = Time(self._dists[dm].rvs(size=num), format="mjd") + + return out class IceTrackDR2LifeTime(LifeTime): @@ -173,16 +195,3 @@ def __init__(self): # density=True for on_off to be treated as density self._dists[s] = stats.rv_histogram((on_off, bins), density=True) - def draw_event_mjd( - self, event_numbers: dict[EventType, int] - ) -> dict[EventType, Time]: - """Draw event MJDs for scrambling the arrival times of events - TODO: add min/max mjd - - :param event_numbers: Dict of event type and requested numbers - :type event_numbers: dict[EventType, int] - :return: Dictionary of event types and new MJDs - :rtype: dict[EventType, Time] - """ - - raise NotImplementedError()