diff --git a/setup.py b/setup.py index 4e62c0f..ae710d7 100644 --- a/setup.py +++ b/setup.py @@ -1,11 +1,21 @@ import setuptools +import os +import re with open("README.md", "r") as fh: long_description = fh.read() +# read version from src/cdtools/_version.py +version_file = os.path.join("src/cdtools", "_version.py") +with open(version_file) as f: + version_match = re.search(r"^__version__ = ['\"]([^'\"]*)['\"]", f.read(), re.M) +if not version_match: + raise RuntimeError("Unable to find version string.") +version = version_match.group(1) + setuptools.setup( name="cdtools-py", - version="0.3.0", + version=version, python_requires='>3.8', # recommended minimum version for pytorch 2.3.0 author="Abe Levitan", author_email="abraham.levitan@psi.ch", diff --git a/src/cdtools/__init__.py b/src/cdtools/__init__.py index 9132209..19dc4c6 100644 --- a/src/cdtools/__init__.py +++ b/src/cdtools/__init__.py @@ -6,6 +6,8 @@ warnings.filterwarnings("ignore", __all__ = ['tools', 'datasets', 'models', 'reconstructors'] +from ._version import __version__ + from cdtools import tools from cdtools import datasets from cdtools import models diff --git a/src/cdtools/_version.py b/src/cdtools/_version.py new file mode 100644 index 0000000..9163035 --- /dev/null +++ b/src/cdtools/_version.py @@ -0,0 +1 @@ +__version__ = "0.3.1.dev" diff --git a/src/cdtools/reconstructors/base.py b/src/cdtools/reconstructors/base.py index c869678..84a482a 100644 --- a/src/cdtools/reconstructors/base.py +++ b/src/cdtools/reconstructors/base.py @@ -243,7 +243,7 @@ class Reconstructor: sim_patterns = self.model.forward(*inp) # Calculate the loss - if hasattr(self, 'mask'): + if hasattr(self.model, 'mask'): loss = self.model.loss(pats, sim_patterns, mask=self.model.mask) diff --git a/tests/models/test_fancy_ptycho.py b/tests/models/test_fancy_ptycho.py index 8bc4d87..02893f9 100644 --- a/tests/models/test_fancy_ptycho.py +++ b/tests/models/test_fancy_ptycho.py @@ -52,6 +52,10 @@ def test_lab_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot): print('\nTesting performance on the standard transmission ptycho dataset') dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(lab_ptycho_cxi) + # Test the masking system + dataset.mask[110:115,65:70] = 0 + dataset.patterns[...,~dataset.mask] = t.max(dataset.patterns) + model = cdtools.models.FancyPtycho.from_dataset( dataset, n_modes=3, diff --git a/tests/test_version.py b/tests/test_version.py new file mode 100644 index 0000000..5b6a778 --- /dev/null +++ b/tests/test_version.py @@ -0,0 +1,18 @@ +from cdtools import __version__ +import re + + +def test_version_exists(): + """Test that version is defined and not empty.""" + assert __version__ + assert isinstance(__version__, str) + assert len(__version__) > 0 + + +def test_version_format(): + """Test that version follows semantic versioning format.""" + # Basic semantic versioning pattern (X.Y.Z with optional pre-release) + pattern = r"^\d+\.\d+\.\d+(?:[-.]?(?:alpha|beta|rc|dev)\d*)?$" + assert re.match( + pattern, __version__ + ), f"Version '{__version__}' doesn't follow semantic versioning"