mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
Merge pull request #35 from gnzng/float64
handling of 64-bit floats in Ptycho2DDataset cxi import
This commit is contained in:
@@ -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
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user