Merge branch 'master' into multigpu

This commit is contained in:
yoshikisd
2025-11-09 04:19:51 +00:00
6 changed files with 37 additions and 2 deletions
+11 -1
View File
@@ -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",
+2
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
__version__ = "0.3.1.dev"
+1 -1
View File
@@ -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)
+4
View File
@@ -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,
+18
View File
@@ -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"