From ddcc5bbab469025ea07c95e80ef8431b458c1902 Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Tue, 8 Mar 2022 10:01:26 -0500 Subject: [PATCH] Fix ptycho_2d_dataset --- CDTools/datasets/ptycho_2d_dataset.py | 2 +- CDTools/models/fancy_ptycho.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/CDTools/datasets/ptycho_2d_dataset.py b/CDTools/datasets/ptycho_2d_dataset.py index e81e990..40f81e4 100644 --- a/CDTools/datasets/ptycho_2d_dataset.py +++ b/CDTools/datasets/ptycho_2d_dataset.py @@ -165,7 +165,7 @@ class Ptycho2DDataset(CDataset): dataset.mask = t.ones(dataset.patterns.shape[-2:]).to(dtype=t.bool) try: - intensities = cdtdata.get_shot_to_shot_info('intensities') + intensities = cdtdata.get_shot_to_shot_info(cxi_file, 'intensities') dataset.intensities = t.as_tensor(intensities, dtype=t.float32) except KeyError: dataset.intensities = None diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index 02b817e..42aa733 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -35,7 +35,7 @@ class FancyPtycho(CDIModel): det_geo['distance'] = t.tensor(det_geo['distance'], dtype=t.float32) if 'basis' in det_geo: det_geo['basis'] = t.tensor(det_geo['basis'], dtype=t.float32) - if 'corner' in det_geo: + if 'corner' in det_geo and det_geo['corner'] is not None: det_geo['corner'] = t.tensor(det_geo['corner'], dtype=t.float32) self.min_translation = t.tensor(min_translation) @@ -347,7 +347,7 @@ class FancyPtycho(CDIModel): det_geo['distance'] = det_geo['distance'].to(*args, **kwargs) if 'basis' in det_geo: det_geo['basis'] = det_geo['basis'].to(*args, **kwargs) - if 'corner' in det_geo: + if 'corner' in det_geo and det_geo['corner'] is not None: det_geo['corner'] = det_geo['corner'].to(*args, **kwargs) if self.mask is not None: