Merge pull request #35 from gnzng/float64

handling of 64-bit floats in Ptycho2DDataset cxi import
This commit is contained in:
Dayne Yoshiki Sasaki
2025-06-17 16:21:56 -07:00
committed by GitHub
2 changed files with 62 additions and 19 deletions
+17 -12
View File
@@ -1,15 +1,16 @@
import warnings
from copy import copy, deepcopy
import pathlib
import h5py
import numpy as np
import torch as t
from copy import copy
import h5py
import pathlib
from cdtools.datasets import CDataset
from cdtools.datasets.random_selection import random_selection
from cdtools.tools import data as cdtdata
from cdtools.tools import plotting
from matplotlib import pyplot as plt
from cdtools.tools import analysis
from copy import deepcopy
__all__ = ['Ptycho2DDataset']
@@ -124,9 +125,9 @@ class Ptycho2DDataset(CDataset):
self.translations = self.translations.to(*args, **kwargs)
self.patterns = self.patterns.to(*args, **kwargs)
# It sucks that I can't reuse the base factory method here,
# perhaps there is a way but I couldn't figure it out.
@classmethod
def from_cxi(cls, cxi_file, cut_zeros=True, load_patterns=True):
"""Generates a new Ptycho2DDataset from a .cxi file directly
@@ -148,7 +149,7 @@ class Ptycho2DDataset(CDataset):
"""
# If a bare string is passed
if isinstance(cxi_file, str) or isinstance(cxi_file, pathlib.Path):
with h5py.File(cxi_file,'r') as f:
with h5py.File(cxi_file, 'r') as f:
return cls.from_cxi(f, cut_zeros=cut_zeros, load_patterns=load_patterns)
# Generate a base dataset
@@ -164,10 +165,15 @@ class Ptycho2DDataset(CDataset):
patterns, axes = cdtdata.get_data(cxi_file, cut_zeros=cut_zeros)
dataset.patterns = t.as_tensor(patterns)
if dataset.patterns.dtype == t.float64:
raise NotImplementedError('64-bit floats are not supported and precision will not be retained in reconstructions! Please explicitly convert your data to 32-bit or submit a pull request')
# If the data is 64-bit, we need to convert it to 32-bit
# because 64-bit floats are not supported in reconstructions
dataset.patterns = dataset.patterns.to(dtype=t.float32)
warnings.warn(
"64-bit floats are not supported and precision will not be retained in reconstructions and were converted to t.float32! "
"If you would like to have 64-bit support, please open an issue or submit a pull request."
)
dataset.axes = axes
if dataset.mask is None:
dataset.mask = t.ones(dataset.patterns.shape[-2:]).to(dtype=t.bool)
@@ -176,9 +182,8 @@ class Ptycho2DDataset(CDataset):
dataset.intensities = t.as_tensor(intensities, dtype=t.float32)
except KeyError:
dataset.intensities = None
return dataset
return dataset
def to_cxi(self, cxi_file):
"""Saves out a Ptycho2DDataset as a .cxi file
+45 -7
View File
@@ -1,12 +1,16 @@
import datetime
import itertools
import os
from copy import deepcopy
import h5py
import numpy as np
import pytest
import torch as t
from cdtools.datasets import CDataset, Ptycho2DDataset
from cdtools.tools import data as cdtdata
import numpy as np
import torch as t
import h5py
import datetime
from copy import deepcopy
import pytest
import itertools
#
# We start by testing the CDataset base class
@@ -220,6 +224,40 @@ def test_Ptycho2DDataset_from_cxi(test_ptycho_cxis):
assert t.allclose(t.tensor(expected['translations']),dataset.translations)
def test_Ptycho2DDataset_from_cxi_64bit(test_ptycho_cxis):
"""Test that we can load a 64-bit cxi file. Should issue
a warning, but still load the data."""
# create test patterns and translations
np.random.seed(42)
patterns = np.random.rand(20, 256, 256).astype(np.float64)
translations = np.random.rand(20, 3).astype(np.float64)
dataset = Ptycho2DDataset(translations, patterns)
dataset.detector_geometry = {
'distance': 0.1, # in meters
'basis': t.tensor([
[-0e-06, -13.5e-06 * 4],
[-13.5e-06 * 4, 0e-06],
[0e-06, 0e-06]
]),
'corner': None
}
dataset.wavelength = 1.6891579427792915e-09 # in meters
# and save to a temp file
dataset.to_cxi('test_Ptycho2DDataset_from_cxi_64bit.cxi')
with pytest.warns(UserWarning, match='64-bit floats'):
dataset_64bit = Ptycho2DDataset.from_cxi('test_Ptycho2DDataset_from_cxi_64bit.cxi')
# Check that the data is loaded correctly
assert dataset_64bit.patterns.dtype == t.float32
assert dataset_64bit.translations.dtype == t.float32
# delete the created test file
os.remove('test_Ptycho2DDataset_from_cxi_64bit.cxi')
def test_Ptycho2DDataset_to_cxi(test_ptycho_cxis, tmp_path):
for cxi, expected in test_ptycho_cxis:
print('loading dataset')