Skip to content

Commit eca130c

Browse files
authored
Merge pull request #90 from PyAutoLabs/feature/cti-resurrection-phase4
CTI resurrection Phase 4 companions: factor-graph runtime fixes
2 parents fcf6619 + 4c001b6 commit eca130c

16 files changed

Lines changed: 127 additions & 71 deletions

File tree

autocti/aggregator/imaging_ci.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,6 +89,7 @@ def values_from(hdu: int) -> aa.Array2D:
8989
cosmic_ray_map=cosmic_ray_map,
9090
settings_dict=settings_dict,
9191
layout=layout,
92+
check_noise_map=False,
9293
)
9394

9495
dataset_list.append(dataset.apply_mask(mask=mask))

autocti/charge_injection/imaging/imaging.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,11 @@ def __init__(
2424
noise_scaling_map_dict: Optional[Dict] = None,
2525
fpr_value: Optional[float] = None,
2626
settings_dict: Optional[Dict] = None,
27+
check_noise_map: bool = True,
2728
):
28-
super().__init__(data=data, noise_map=noise_map)
29+
super().__init__(
30+
data=data, noise_map=noise_map, check_noise_map=check_noise_map
31+
)
2932

3033
self.data = self.data.native
3134
self.noise_map = self.noise_map.native
@@ -319,7 +322,9 @@ def output_to_fits(
319322
exception is raised.
320323
"""
321324
fitsable.output_to_fits(
322-
values=np.asarray(self.data.native), file_path=data_path, overwrite=overwrite
325+
values=np.asarray(self.data.native),
326+
file_path=data_path,
327+
overwrite=overwrite,
323328
)
324329

325330
if noise_map_path is not None:

autocti/charge_injection/model/plotter.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77
from autocti.charge_injection.plot import fit_ci_plots
88
from autocti.model.plotter import Plotter, plot_setting
99

10+
from autoarray import exc as aa_exc
11+
1012
from autocti import exc
1113

1214
logger = logging.getLogger(__name__)
@@ -85,7 +87,16 @@ def should_plot(name):
8587
output_format=self.fmt,
8688
title_prefix=self.title_prefix,
8789
)
88-
except (exc.PlottingException, exc.RegionException, TypeError, ValueError):
90+
except (
91+
exc.PlottingException,
92+
exc.RegionException,
93+
aa_exc.ArrayException,
94+
TypeError,
95+
ValueError,
96+
):
97+
# Trimmed datasets (e.g. via `apply_settings`) have layouts whose
98+
# extraction regions no longer match the array shape, making the
99+
# binned-FPR diagnostic ill-defined.
89100
logger.info(
90101
"VISUALIZATION - Could not visualize the ImagingCI binned data"
91102
)

autocti/charge_injection/model/result.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,17 +5,17 @@
55
class ResultImagingCI(ResultDataset):
66
@property
77
def max_log_likelihood_full_fit(self) -> FitImagingCI:
8-
return self.analysis.fit_via_instance_and_dataset_from(
8+
return self.analysis_unwrapped.fit_via_instance_and_dataset_from(
99
instance=self.instance,
10-
dataset=self.analysis.dataset_full,
10+
dataset=self.analysis_unwrapped.dataset_full,
1111
hyper_noise_scale=True,
1212
)
1313

1414
@property
1515
def max_log_likelihood_full_fit_no_hyper_scaling(self):
16-
return self.analysis.fit_via_instance_and_dataset_from(
16+
return self.analysis_unwrapped.fit_via_instance_and_dataset_from(
1717
instance=self.instance,
18-
dataset=self.analysis.dataset_full,
18+
dataset=self.analysis_unwrapped.dataset_full,
1919
hyper_noise_scale=False,
2020
)
2121

autocti/charge_injection/model/visualizer.py

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -165,12 +165,15 @@ def visualize_combined(
165165
paths: af.DirectoryPaths,
166166
instance: af.ModelInstance,
167167
during_analysis: bool,
168+
quick_update: bool = False,
168169
):
169170
if analyses is None:
170171
return
171172

172173
fit_list = [
173-
analysis.fit_via_instance_from(instance=instance) for analysis in analyses
174+
# The factor graph passes one instance per analysis factor.
175+
analysis.fit_via_instance_from(instance=instance_single)
176+
for analysis, instance_single in zip(analyses, instance)
174177
]
175178

176179
fpr_value_list = [fit.dataset.fpr_value for fit in fit_list]
@@ -180,7 +183,9 @@ def visualize_combined(
180183
fpr_value_list=fpr_value_list,
181184
)
182185

183-
region_list = analyses[0].region_list_from(model=instance)
186+
# The factor graph passes one instance per analysis factor; the region
187+
# list is derived from the first (the CTI model is shared across factors).
188+
region_list = analyses[0].region_list_from(model=instance[0])
184189

185190
visualizer = PlotterImagingCI(image_path=paths.image_path)
186191
visualizer.fit_combined(fit_list=fit_list, during_analysis=during_analysis)
@@ -193,9 +198,9 @@ def visualize_combined(
193198
if analyses[0].dataset_full is not None:
194199
fit_full_list = [
195200
analysis.fit_via_instance_and_dataset_from(
196-
instance=instance, dataset=analysis.dataset_full
201+
instance=instance_single, dataset=analysis.dataset_full
197202
)
198-
for analysis in analyses
203+
for analysis, instance_single in zip(analyses, instance)
199204
]
200205

201206
fit_full_list = analyses[0].in_ascending_fpr_order_from(

autocti/charge_injection/plot/fit_ci_plots.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -226,6 +226,9 @@ def subplot_fit_list(
226226
_pf = (lambda t: f"{title_prefix.rstrip()} {t}") if title_prefix else (lambda t: t)
227227

228228
n = len(fit_list)
229+
if n == 0:
230+
raise ValueError("An empty list was passed to a *_list plot function.")
231+
229232
cols = min(n, 3)
230233
rows = (n + cols - 1) // cols
231234

@@ -279,6 +282,9 @@ def subplot_fit_region_list(
279282
output_format = output_format[0]
280283

281284
n = len(fit_list)
285+
if n == 0:
286+
raise ValueError("An empty list was passed to a *_list plot function.")
287+
282288
cols = min(n, 3)
283289
rows = (n + cols - 1) // cols
284290

autocti/charge_injection/plot/imaging_ci_plots.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -328,6 +328,9 @@ def subplot_data_region_list(
328328
output_format = output_format[0]
329329

330330
n = len(dataset_list)
331+
if n == 0:
332+
raise ValueError("An empty list was passed to a *_list plot function.")
333+
331334
cols = min(n, 3)
332335
rows = (n + cols - 1) // cols
333336

autocti/dataset_1d/dataset_1d/dataset_1d.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -182,7 +182,9 @@ def output_to_fits(
182182
exception is raised.
183183
"""
184184
fitsable.output_to_fits(
185-
values=np.asarray(self.data.native), file_path=data_path, overwrite=overwrite
185+
values=np.asarray(self.data.native),
186+
file_path=data_path,
187+
overwrite=overwrite,
186188
)
187189
fitsable.output_to_fits(
188190
values=np.asarray(self.noise_map.native),

autocti/dataset_1d/model/visualizer.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -139,12 +139,15 @@ def visualize_combined(
139139
paths: af.DirectoryPaths,
140140
instance: af.ModelInstance,
141141
during_analysis: bool,
142+
quick_update: bool = False,
142143
):
143144
if analyses is None:
144145
return
145146

146147
fit_list = [
147-
analysis.fit_via_instance_from(instance=instance) for analysis in analyses
148+
# The factor graph passes one instance per analysis factor.
149+
analysis.fit_via_instance_from(instance=instance_single)
150+
for analysis, instance_single in zip(analyses, instance)
148151
]
149152

150153
fpr_value_list = [fit.dataset.fpr_value for fit in fit_list]
@@ -167,9 +170,9 @@ def visualize_combined(
167170
if analyses[0].dataset_full is not None:
168171
fit_full_list = [
169172
analysis.fit_via_instance_and_dataset_from(
170-
instance=instance, dataset=analysis.dataset_full
173+
instance=instance_single, dataset=analysis.dataset_full
171174
)
172-
for analysis in analyses
175+
for analysis, instance_single in zip(analyses, instance)
173176
]
174177

175178
fit_full_list = analyses[0].in_ascending_fpr_order_from(

autocti/dataset_1d/plot/dataset_1d_plots.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -141,6 +141,9 @@ def subplot_dataset_list(
141141
suffix = f"_{region}" if region is not None else ""
142142

143143
n = len(dataset_list)
144+
if n == 0:
145+
raise ValueError("An empty list was passed to a *_list plot function.")
146+
144147
cols = min(n, 3)
145148
rows = (n + cols - 1) // cols
146149

0 commit comments

Comments
 (0)