Skip to content

Commit 7ce410e

Browse files
Use more memory maps also for final results, and eliminate some at the end
1 parent e6fcc63 commit 7ce410e

1 file changed

Lines changed: 27 additions & 9 deletions

File tree

hendrics/efsearch.py

Lines changed: 27 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)