Fixed merging conflict - log scale in ptycho_2d_dataset.py

This commit is contained in:
David Rower
2019-10-03 15:25:27 -04:00
19 changed files with 144 additions and 49 deletions
+9 -3
View File
@@ -32,8 +32,11 @@ import numpy as np
import torch as t
from copy import copy
import h5py
import pathlib
try:
import pathlib
except ImportError:
import pathlib2 as pathlib
from CDTools.tools import data as cdtdata
from CDTools.tools import plotting
from torch.utils import data as torchdata
@@ -98,7 +101,10 @@ class CDataset(torchdata.Dataset):
self.wavelength = wavelength
self.detector_geometry = copy(detector_geometry)
if mask is not None:
self.mask = t.tensor(mask)
if isinstance(mask, t.Tensor):
self.mask = mask.detach().to(dtype=t.bool)
else:
self.mask = t.BoolTensor(mask)
else:
self.mask = None
if background is not None:
+15 -5
View File
@@ -3,7 +3,10 @@ import numpy as np
import torch as t
from copy import copy
import h5py
import pathlib
try:
import pathlib
except ImportError:
import pathlib2 as pathlib
from CDTools.datasets import CDataset
from CDTools.tools import data as cdtdata
@@ -62,7 +65,7 @@ class Ptycho2DDataset(CDataset):
self.translations = t.tensor(translations)
self.patterns = t.tensor(patterns)
if self.mask is None:
self.mask = t.ones(self.patterns.shape[-2:]).to(dtype=t.uint8)
self.mask = t.ones(self.patterns.shape[-2:]).to(dtype=t.bool)
self.mask.masked_fill_(t.isnan(t.sum(self.patterns,dim=(0,))),0)
self.patterns.masked_fill_(t.isnan(self.patterns),0)
@@ -181,7 +184,7 @@ class Ptycho2DDataset(CDataset):
cdtdata.add_ptycho_translations(cxi_file, self.translations)
def inspect(self):
def inspect(self, logarithmic=True):
"""Launches an interactive plot for perusing the data
This launches an interactive plotting tool in matplotlib that
@@ -252,7 +255,10 @@ class Ptycho2DDataset(CDataset):
cb1.ax.set_title('Integrated Intensity', size="medium", pad=5)
cb1.ax.tick_params(labelrotation=20)
meas = axes[1].imshow(np.log(meas_data) / np.log(10) * mask)
if logarithmic:
meas = axes[1].imshow(np.log(meas_data) / np.log(10) * mask)
else:
meas = axes[1].imshow(meas_data * mask)
cb2 = plt.colorbar(meas, ax=axes[1], orientation='horizontal',format='%.2e',ticks=ticker.LinearLocator(numticks=5),pad=0.17,fraction=0.1)
cb2.ax.tick_params(labelrotation=20)
@@ -280,7 +286,11 @@ class Ptycho2DDataset(CDataset):
meas = axes[1].images[-1]
meas.set_data(np.log(meas_data) / np.log(10) * mask)
if logarithmic:
meas.set_data(np.log(meas_data) / np.log(10) * mask)
else:
meas.set_data(meas_data * mask)
update_colorbar(meas)