@@ -604,31 +604,42 @@ def transient_search(
604604 )
605605 if all_results is None :
606606 results_shape = (len (allvalues ), nave .size , results .shape [1 ])
607+ use_memmap = False
607608 if np .prod (results_shape ) > 1e7 :
608- log .info (
609- "Transient search results are very large. "
610- f"Using memmapped arrays (shape { results_shape } ) to reduce memory usage." ,
611- )
612609 import tempfile
613610 tmp_results = tempfile .NamedTemporaryFile (delete = True ).name + "_hen.npy"
614611 tmp_f = tempfile .NamedTemporaryFile (delete = True ).name + "_hen.npy"
612+ log .info (
613+ "Transient search results are very large. "
614+ f"Using memmapped arrays ({ tmp_results } ; "
615+ f"shape { results_shape } ) to reduce memory usage." ,
616+ )
615617 all_results = np .lib .format .open_memmap (
616618 tmp_results , mode = "w+" , dtype = results .dtype , shape = results_shape
617619 )
618620 all_freqs = np .lib .format .open_memmap (
619621 tmp_f , mode = "w+" , dtype = mean_f .dtype , shape = (len (allvalues ),)
620622 )
623+ use_memmap = True
621624 else :
622625 all_results = np .empty (results_shape , dtype = results .dtype )
623626 all_freqs = np .empty ((len (allvalues ),), dtype = mean_f .dtype )
624627
625628 all_results [ii ] = results
626629 all_freqs [ii ] = mean_f
627630
628- all_results = np .array (all_results )
629- all_freqs = np .array (all_freqs )
630-
631631 times = dt * np .arange (all_results .shape [2 ])
632+ final_results_shape = (nave .size , all_results .shape [2 ], all_results .shape [0 ])
633+ if use_memmap :
634+
635+ tmp_results_stats = tempfile .NamedTemporaryFile (delete = True ).name + "_hen.npy"
636+ all_results_stats = np .lib .format .open_memmap (
637+ tmp_results_stats , mode = "w+" , dtype = results .dtype , shape = final_results_shape
638+ )
639+ else :
640+ all_results_stats = np .empty (final_results_shape , dtype = results .dtype )
641+ for i in range (nave .size ):
642+ all_results_stats [i ] = all_results [:, i , :].T
632643
633644 results = TransientResults ()
634645 results .oversample = oversample
@@ -638,7 +649,12 @@ def transient_search(
638649 results .nave = nave
639650 results .freqs = all_freqs
640651 results .times = times
641- results .stats = np .array ([all_results [:, i , :].T for i in range (nave .size )])
652+ results .stats = all_results_stats
653+
654+ if use_memmap :
655+ os .remove (all_results .filename )
656+ os .remove (all_freqs .filename )
657+ del all_results , all_freqs
642658
643659 return results
644660
@@ -658,7 +674,9 @@ def plot_transient_search(results, gif_name=None):
658674 result_name = gif_name .replace (".gif" , ".csv" )
659675 max_stats_rows = []
660676 all_images = []
661- for i , (ima , nave ) in enumerate (zip (results .stats , results .nave )):
677+ import tqdm
678+ log .info ("Generating plots for transient search..." )
679+ for i , (ima , nave ) in tqdm .tqdm (enumerate (zip (results .stats , results .nave )), total = len (results .nave )):
662680 f = results .freqs
663681 t = results .times
664682 nprof = ima .shape [0 ]
0 commit comments