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
3 changes: 3 additions & 0 deletions icecube_data_reader/event_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
37 changes: 29 additions & 8 deletions icecube_data_reader/events.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@
EventType,
Refrigerator,
)
from icecube_data_reader.lifetime import LifeTime

from typing import Self
import logging
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:
"""
Expand Down
37 changes: 23 additions & 14 deletions icecube_data_reader/lifetime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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()