|
1 | 1 | from __future__ import annotations |
2 | 2 | from abc import ABC, abstractmethod |
3 | 3 | from functools import partial |
| 4 | +import logging |
4 | 5 | from typing import List, Optional, Generator |
5 | 6 |
|
6 | 7 | import autofit as af |
7 | 8 |
|
| 9 | +logger = logging.getLogger(__name__) |
| 10 | + |
8 | 11 |
|
9 | 12 | class AggBase(ABC): |
10 | 13 | def __init__(self, aggregator: af.Aggregator): |
@@ -82,13 +85,13 @@ def weights_above_gen_from(self, minimum_weight: float) -> List: |
82 | 85 | def func_gen(fit: af.Fit, minimum_weight: float) -> List[object]: |
83 | 86 | samples = fit.samples |
84 | 87 |
|
85 | | - weight_list = [] |
86 | | - |
87 | | - for sample in samples.sample_list: |
88 | | - if sample.weight > minimum_weight: |
89 | | - weight_list.append(sample.weight) |
90 | | - |
91 | | - return weight_list |
| 88 | + return [ |
| 89 | + sample.weight |
| 90 | + for sample, _ in self._valid_sample_instance_pairs( |
| 91 | + samples=samples, |
| 92 | + minimum_weight=minimum_weight, |
| 93 | + ) |
| 94 | + ] |
92 | 95 |
|
93 | 96 | func = partial(func_gen, minimum_weight=minimum_weight) |
94 | 97 |
|
@@ -119,22 +122,51 @@ def all_above_weight_gen_from(self, minimum_weight: float) -> Generator: |
119 | 122 | def func_gen(fit: af.Fit, minimum_weight: float) -> List[object]: |
120 | 123 | samples = fit.samples |
121 | 124 |
|
122 | | - all_above_weight_list = [] |
123 | | - |
124 | | - for sample in samples.sample_list: |
125 | | - if sample.weight > minimum_weight: |
126 | | - instance = sample.instance_for_model(model=samples.model) |
127 | | - |
128 | | - all_above_weight_list.append( |
129 | | - self.object_via_gen_from(fit=fit, instance=instance) |
130 | | - ) |
131 | | - |
132 | | - return all_above_weight_list |
| 125 | + return [ |
| 126 | + self.object_via_gen_from(fit=fit, instance=instance) |
| 127 | + for _, instance in self._valid_sample_instance_pairs( |
| 128 | + samples=samples, |
| 129 | + minimum_weight=minimum_weight, |
| 130 | + ) |
| 131 | + ] |
133 | 132 |
|
134 | 133 | func = partial(func_gen, minimum_weight=minimum_weight) |
135 | 134 |
|
136 | 135 | return self.aggregator.map(func=func) |
137 | 136 |
|
| 137 | + @staticmethod |
| 138 | + def _valid_sample_instance_pairs(samples, minimum_weight: float): |
| 139 | + """Return weighted samples whose model instances still reconstruct. |
| 140 | +
|
| 141 | + Constructor validation can become stricter after a result was written. |
| 142 | + Such historical points are not usable objects, but they must not make an |
| 143 | + entire aggregator query fail. ``FitException`` is the narrow model-point |
| 144 | + rejection contract; programming errors continue to propagate. |
| 145 | + """ |
| 146 | + pairs = [] |
| 147 | + rejected = 0 |
| 148 | + |
| 149 | + for sample in samples.sample_list: |
| 150 | + if sample.weight <= minimum_weight: |
| 151 | + continue |
| 152 | + try: |
| 153 | + instance = samples.model.instance_from_vector( |
| 154 | + sample.parameter_lists_for_model(model=samples.model) |
| 155 | + ) |
| 156 | + except af.exc.FitException: |
| 157 | + rejected += 1 |
| 158 | + continue |
| 159 | + pairs.append((sample, instance)) |
| 160 | + |
| 161 | + if rejected: |
| 162 | + logger.warning( |
| 163 | + "Skipped %d stored sample(s) rejected by current model " |
| 164 | + "validation while building aggregator objects.", |
| 165 | + rejected, |
| 166 | + ) |
| 167 | + |
| 168 | + return pairs |
| 169 | + |
138 | 170 | def randomly_drawn_via_pdf_gen_from(self, total_samples: int): |
139 | 171 | """ |
140 | 172 | Returns a generator which for every result generates a list of objects whose parameter values are drawn |
|
0 commit comments