Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 7 additions & 8 deletions autofit/database/aggregator/aggregator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
42 changes: 40 additions & 2 deletions autofit/non_linear/paths/database.py
Original file line number Diff line number Diff line change
@@ -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):
Expand Down Expand Up @@ -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()
Expand Down
21 changes: 21 additions & 0 deletions test_autofit/aggregator/test_aggregator.py
Original file line number Diff line number Diff line change
@@ -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
):
Expand Down
Loading