From c35c143f9069e2091e668b015719c77c029f150e Mon Sep 17 00:00:00 2001 From: Mario Damiano Date: Tue, 15 Sep 2026 10:55:46 -0700 Subject: [PATCH] Refactor array indexing to use .item() and .size for improved clarity and performance in EXOCAT1 and tieredScheduler classes --- EXOSIMS/StarCatalog/EXOCAT1.py | 2 +- EXOSIMS/SurveySimulation/tieredScheduler.py | 77 +++++++++++---------- 2 files changed, 42 insertions(+), 37 deletions(-) diff --git a/EXOSIMS/StarCatalog/EXOCAT1.py b/EXOSIMS/StarCatalog/EXOCAT1.py index 62ed7fef..87fd30fd 100644 --- a/EXOSIMS/StarCatalog/EXOCAT1.py +++ b/EXOSIMS/StarCatalog/EXOCAT1.py @@ -149,7 +149,7 @@ def __init__(self, catalogpath=None, wdsfilepath=None, **specs): # find indices of wdsdata in catalog inds = np.zeros(len(wdsdat), dtype=int) for j, h in enumerate(wdsdat["HIP"].data): - inds[j] = np.where(HIPnums == h)[0] + inds[j] = np.where(HIPnums == h)[0].item() # add attributes to catalog wdsatts = ['Close_Sep(")', "Close(M2-M1)", 'Bright_Sep(")', "Bright(M2-M1)"] diff --git a/EXOSIMS/SurveySimulation/tieredScheduler.py b/EXOSIMS/SurveySimulation/tieredScheduler.py index 0949f548..5378c464 100644 --- a/EXOSIMS/SurveySimulation/tieredScheduler.py +++ b/EXOSIMS/SurveySimulation/tieredScheduler.py @@ -276,7 +276,7 @@ def __init__( fZ = self.occ_valfZmin[sInds] # fZ = ZL.fZ(Obs, TL, sInds, TK.currentTimeAbs.copy(), char_mode) # Walker previous version. - JEZ = TL.JEZ0[char_mode["hex"]][sInds][self.known_earths] + JEZ = TL.JEZ0[char_mode["hex"]][sInds] if SU.lucky_planets: phi = (1 / np.pi) * np.ones(len(SU.d)) dMag = deltaMag(SU.p, SU.Rp, SU.d, phi)[ @@ -330,7 +330,7 @@ def run_sim(self): spectroModes = list( filter(lambda mode: "spec" in mode["inst"]["name"], OS.observingModes) ) - if np.any(spectroModes): + if spectroModes: char_mode = spectroModes[0] # if no spectro mode, default char mode is first observing mode else: @@ -352,6 +352,7 @@ def run_sim(self): DRM, sInd, occ_sInd, t_det, sd, occ_sInds = self.next_target( sInd, occ_sInd, det_mode, char_mode ) + t_det = 0 * u.d if t_det is None else t_det # print(TK.currentTimeAbs.copy()) @@ -408,8 +409,8 @@ def run_sim(self): cnt += 1 # clean up revisit list when one occurs to prevent repeats - if np.any(self.starRevisit) and np.any( - np.where(self.starRevisit[:, 0] == float(sInd)) + if self.starRevisit.size and np.any( + self.starRevisit[:, 0] == float(sInd) ): s_revs = np.where(self.starRevisit[:, 0] == float(sInd))[0] t_revs = np.where( @@ -463,7 +464,7 @@ def run_sim(self): self.starVisits[sInd] += 1 # PERFORM DETECTION and populate revisit list attribute. # First store dMag, WA - if np.any(pInds): + if pInds.size: DRM["det_dMag"] = SU.dMag[pInds].tolist() DRM["det_WA"] = SU.WA[pInds].to("mas").value.tolist() ( @@ -474,8 +475,8 @@ def run_sim(self): det_SNR, FA, ) = self.observation_detection(sInd, t_det, det_mode) - if np.any(pInds): - DRM["det_JEZ"] = det_JEZ + if pInds.size: + DRM["det_JEZ"] = det_JEZ.to(self.JEZ_unit) if np.any(detected > 0): self.sInd_detcounts[sInd] += 1 @@ -553,8 +554,8 @@ def run_sim(self): # make sure we don't accidentally double characterize TK.advanceToAbsTime(TK.currentTimeAbs.copy() + 0.01 * u.d) assert char_intTime != 0, "Integration time can't be 0." - if np.any(occ_pInds): - DRM["char_JEZ"] = char_JEZ.tolist() + if occ_pInds.size: + DRM["char_JEZ"] = char_JEZ.to(self.JEZ_unit) DRM["char_dMag"] = SU.dMag[occ_pInds].tolist() DRM["char_WA"] = SU.WA[occ_pInds].to("mas").value.tolist() DRM["char_mode"] = dict(char_mode) @@ -568,7 +569,7 @@ def run_sim(self): TL, occ_sInd, char_fZ, - char_JEZ, + TL.JEZ0[char_mode["hex"]][occ_sInd], TL.int_WA[occ_sInd], char_mode, )[0] @@ -590,7 +591,7 @@ def run_sim(self): DRM["FA_char_status"] = characterized[-1] if FA else 0 DRM["FA_char_SNR"] = char_SNR[-1] if FA else 0.0 DRM["FA_char_JEZ"] = ( - self.lastDetected[sInd, 1][-1] / u.arcsec**2 + self.lastDetected[sInd, 1][-1].to(self.JEZ_unit) if FA else 0.0 * u.ph / u.s / u.m**2 / u.arcsec**2 ) @@ -641,7 +642,7 @@ def run_sim(self): elif ( goal_GAdiff > 1 * u.d and (self.occ_arrives - TK.currentTimeAbs.copy()) < -5 * u.d - and not np.any(occ_sInds) + and occ_sInds.size == 0 ): self.vprint( ( @@ -808,7 +809,7 @@ def promote_coro_targets(self, occ_sInds, sInds): promote_stars = sInds[ np.where(self.sInd_detcounts[sInds] >= self.n_det_min)[0] ] - if np.any(promote_stars): + if promote_stars.size: for sInd in promote_stars: pInds = np.where(SU.plan2star == sInd)[0] sp = SU.s[pInds] @@ -1005,7 +1006,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): sInds = self.revisitFilter(sInds, TK.currentTimeNorm.copy()) # revisit list, with time after start - if np.any(occ_sInds): + if occ_sInds.size: occ_tovisit[occ_sInds] = ( self.occ_starVisits[occ_sInds] == self.occ_starVisits[occ_sInds].min() @@ -1074,7 +1075,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): occ_earths = np.intersect1d( np.where(SU.plan2star == occ_star)[0], self.known_earths ).astype(int) - if np.any(occ_earths): + if occ_earths.size: fZ = ZL.fZ( Obs, TL, @@ -1082,7 +1083,9 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): occ_startTimes[occ_star], char_mode, ) - JEZ = SU.scale_JEZ(occ_star, char_mode) + JEZ = SU.scale_JEZ( + occ_star, char_mode, pInds=occ_earths + ) if SU.lucky_planets: phi = (1 / np.pi) * np.ones(len(SU.d)) dMag = deltaMag(SU.p, SU.Rp, SU.d, phi)[ @@ -1197,7 +1200,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): # 6.1 Filter off any stars visited by the occulter more than the # max number of times - if np.any(occ_sInds): + if occ_sInds.size: occ_sInds = occ_sInds[ (self.occ_starVisits[occ_sInds] < self.occ_max_visits) ] @@ -1217,7 +1220,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): # 7 Filter off cornograph stars with too-long inttimes if self.occ_arrives > TK.currentTimeAbs: available_time = self.occ_arrives - TK.currentTimeAbs.copy() - if np.any(sInds[intTimes[sInds] < available_time]): + if np.any(intTimes[sInds] < available_time): sInds = sInds[intTimes[sInds] < available_time] # 8 remove occ targets on ignore_stars list @@ -1227,7 +1230,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): t_det = 0 * u.d occ_sInd = old_occ_sInd - if np.any(sInds): + if sInds.size: # choose sInd of next target sInd = self.choose_next_telescope_target( old_sInd, sInds, intTimes[sInds] @@ -1238,7 +1241,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): # 8 Choose best target from remaining # if the starshade has arrived at its destination, or it is # the first observation - if np.any(occ_sInds): + if occ_sInds.size: if old_occ_sInd is None or ( (TK.currentTimeAbs.copy() + t_det) >= self.occ_arrives and self.ready_to_update @@ -1253,7 +1256,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): self.occ_slewTime = slewTimes[occ_sInd] self.occ_sd = sd[occ_sInd] self.ready_to_update = False - elif not np.any(sInds): + elif sInds.size == 0: TK.advanceToAbsTime(TK.currentTimeAbs.copy() + 1 * u.d) continue @@ -1263,7 +1266,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): if self.tot_det_int_cutoff < self.tot_dettime: sInds = np.array([]) - if np.any(sInds): + if sInds.size: # choose sInd of next target sInd = self.choose_next_telescope_target( old_sInd, sInds, intTimes[sInds] @@ -1274,7 +1277,7 @@ def next_target(self, old_sInd, old_occ_sInd, det_mode, char_mode): sInd = None # if no observable target, call the TimeKeeping.wait() method - if not np.any(sInds) and not np.any(occ_sInds): + if sInds.size == 0 and occ_sInds.size == 0: self.vprint( "No Observable Targets at currentTimeNorm = " + str(TK.currentTimeNorm.copy()) @@ -1358,7 +1361,7 @@ def choose_next_occulter_target(self, old_occ_sInd, occ_sInds, intTimes): A = A + self.coeffs[2] * (intTimes[occ_sInds] / OS.intCutoff) # add factor for unvisited ramp for deep dive stars - if np.any(top_sInds): + if top_sInds.size: # add factor for least visited deep dive stars f_uv = np.zeros(nStars) u1 = np.isin(occ_sInds, top_sInds) @@ -1521,12 +1524,12 @@ def calc_int_inflection(self, t_sInds, JEZ, startTime, WA, mode, ischar=False): num_points = 500 intTimes = np.logspace(-5, 2, num_points) * u.d sInds = np.arange(TL.nStars) - # don't use WA input because we don't know planet positions - # before characterization + # Cached population curves use reference values for every star. + JEZ = TL.JEZ0[mode["hex"]] WA = TL.int_WA curve = np.zeros([1, sInds.size, intTimes.size]) - Cpath = os.path.join(Comp.classpath, Comp.filename + ".fcomp") + Cpath = os.path.join(Comp.cachedir, Comp.filename + ".JEZ.fcomp") # if no preexisting curves exist, either load from file or calculate if self.curves is None: @@ -1545,7 +1548,7 @@ def calc_int_inflection(self, t_sInds, JEZ, startTime, WA, mode, ischar=False): curve[0, :, t_i] = Comp.comp_per_intTime( t, TL, sInds, fZ, JEZ, WA, mode ) - curves[mode["systName"]] = curve + curves[mode["hex"]] = curve with open(Cpath, "wb") as cfile: pickle.dump(curves, cfile) self.vprint("completeness curves stored in {}".format(Cpath)) @@ -1554,8 +1557,8 @@ def calc_int_inflection(self, t_sInds, JEZ, startTime, WA, mode, ischar=False): # if no curves for current mode if ( - mode["systName"] not in self.curves.keys() - or TL.nStars != self.curves[mode["systName"]].shape[1] + mode["hex"] not in self.curves.keys() + or TL.nStars != self.curves[mode["hex"]].shape[1] ): for t_i, t in enumerate(intTimes): fZ = ZL.fZ(Obs, TL, sInds, startTime, mode) @@ -1563,14 +1566,14 @@ def calc_int_inflection(self, t_sInds, JEZ, startTime, WA, mode, ischar=False): t, TL, sInds, fZ, JEZ, WA, mode ) - self.curves[mode["systName"]] = curve + self.curves[mode["hex"]] = curve with open(Cpath, "wb") as cfile: pickle.dump(self.curves, cfile) self.vprint("recalculated completeness curves stored in {}".format(Cpath)) int_times = np.zeros(len(t_sInds)) * u.d for i, sInd in enumerate(t_sInds): - c_v_t = self.curves[mode["systName"]][0, sInd, :] + c_v_t = self.curves[mode["hex"]][0, sInd, :] dcdt = np.diff(c_v_t) / np.diff(intTimes) # find the inflection point of the completeness graph @@ -1726,7 +1729,7 @@ def observation_characterization(self, sInd, mode): ) fZ = ZL.fZ(Obs, TL, sInd, startTime, mode) - JEZ = JEZs[tochar] + JEZ = JEZs WAp = TL.int_WA[sInd] * np.ones(len(tochar)) dMag = TL.int_dMag[sInd] * np.ones(len(tochar)) @@ -1740,8 +1743,8 @@ def observation_characterization(self, sInd, mode): else: e_dMag = SU.dMag e_WA = SU.WA - WAp[pinds_earthlike[tochar]] = e_WA[pIndsDet[pinds_earthlike]] - dMag[pinds_earthlike[tochar]] = e_dMag[pIndsDet[pinds_earthlike]] + WAp[pinds_earthlike] = e_WA[pIndsDet[pinds_earthlike]] + dMag[pinds_earthlike] = e_dMag[pIndsDet[pinds_earthlike]] intTimes = np.zeros(len(tochar)) * u.day if self.int_inflection: @@ -1751,7 +1754,9 @@ def observation_characterization(self, sInd, mode): [sInd], JEZ[i], startTime, j, mode, ischar=True )[0] else: - intTimes[tochar] = OS.calc_intTime(TL, sInd, fZ, JEZ, dMag, WAp, mode) + intTimes[tochar] = OS.calc_intTime( + TL, sInd, fZ, JEZ[tochar], dMag[tochar], WAp[tochar], mode + ) intTimes[~np.isfinite(intTimes)] = 0 * u.d # add a predetermined margin to the integration times