diff --git a/autofit/database/aggregator/aggregator.py b/autofit/database/aggregator/aggregator.py index 4c27f0c3d..e331b9aa3 100644 --- a/autofit/database/aggregator/aggregator.py +++ b/autofit/database/aggregator/aggregator.py @@ -343,16 +343,15 @@ def __getitem__(self, item): if isinstance(item, int): return self.fits[item] elif isinstance(item, slice): + start_index = 0 if item.start is not None: - if item.start >= 0: - offset += item.start - else: - offset = len(self) + item.start + start_index = ( + item.start if item.start >= 0 else len(self) + item.start + ) + offset += start_index if item.stop is not None: - if item.stop >= 0: - limit = len(self) - item.stop - offset - else: - limit = len(self) + item.stop + stop_index = item.stop if item.stop >= 0 else len(self) + item.stop + limit = max(stop_index - start_index, 0) return self._new_with(offset=offset, limit=limit) def _fits_for_query(self, query: str) -> List[m.Fit]: diff --git a/autofit/non_linear/paths/database.py b/autofit/non_linear/paths/database.py index b146d1a2a..929e25d07 100644 --- a/autofit/non_linear/paths/database.py +++ b/autofit/non_linear/paths/database.py @@ -1,14 +1,15 @@ import shutil from typing import Optional, Union -from autoconf.output import conditional_output +from autoconf.output import conditional_output, should_output from autofit.database.sqlalchemy_ import sa from .abstract import AbstractPaths import numpy as np from autofit.database.model import Fit -from autoconf.dictable import to_dict +from autoconf.dictable import to_dict, from_dict from autofit.database.aggregator.info import Info +from autofit.non_linear.samples.summary import SamplesSummary class DatabasePaths(AbstractPaths): @@ -276,6 +277,43 @@ def save_samples(self, samples): self.fit.samples = samples self.fit.set_json("samples_info", samples.samples_info) + def save_samples_summary( + self, samples_summary: SamplesSummary, name="samples_summary" + ): + model = samples_summary.model + + filter_args = tuple( + arg_name + for arg_name in ( + "errors_at_sigma_1", + "errors_at_sigma_3", + "values_at_sigma_1", + "values_at_sigma_3", + "max_log_likelihood_sample", + "median_pdf_sample", + ) + if not should_output(arg_name) + ) + + samples_summary.model = None + self.fit.set_json( + name or "samples_summary", + to_dict(samples_summary, filter_args=filter_args), + ) + samples_summary.model = model + + def load_samples_summary(self) -> SamplesSummary: + try: + summary_dict = self.fit.get_json("samples_summary") + except KeyError: + raise FileNotFoundError( + f"No samples_summary saved for fit {self.identifier}" + ) + samples_summary = from_dict(summary_dict) + samples_summary.model = self.model + + return samples_summary + def save_latent_samples(self, latent_samples): if not self.save_all_samples: latent_samples = latent_samples.minimise() diff --git a/test_autofit/aggregator/test_aggregator.py b/test_autofit/aggregator/test_aggregator.py index f931006d1..6f7d656cf 100644 --- a/test_autofit/aggregator/test_aggregator.py +++ b/test_autofit/aggregator/test_aggregator.py @@ -1,3 +1,24 @@ +import pytest + +import autofit as af + + +@pytest.fixture(name="aggregator_x3") +def make_aggregator_x3(session): + fits = [af.db.Fit(id=f"fit_{i}", is_complete=True) for i in range(3)] + session.add_all(fits) + session.flush() + return af.Aggregator(session) + + +def test_slicing(aggregator_x3): + assert len(aggregator_x3[:2]) == 2 + assert len(aggregator_x3[1:3]) == 2 + assert len(aggregator_x3[:-1]) == 2 + assert len(aggregator_x3[-2:]) == 2 + assert len(aggregator_x3[2:]) == 1 + + def test_completed_aggregator( aggregator ):