mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-11 22:12:38 +02:00
Make big changes to docs to update the tutorial section, and also to finally autoload the documentation of the models and datasets
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
import torch as t
|
||||
from matplotlib import pyplot as plt
|
||||
from cdtools.datasets import CDataset
|
||||
from cdtools.tools import data as cdtdata
|
||||
|
||||
|
||||
__all__ = ['BasicPtychoDataset']
|
||||
|
||||
|
||||
class BasicPtychoDataset(CDataset):
|
||||
"""The standard dataset for a 2D ptychography scan"""
|
||||
|
||||
def __init__(self, translations, patterns, *args, **kwargs):
|
||||
"""Initialize the dataset from python objects"""
|
||||
|
||||
super().__init__(*args, **kwargs)
|
||||
self.translations = t.Tensor(translations).clone()
|
||||
self.patterns = t.Tensor(patterns).clone()
|
||||
|
||||
def __len__(self):
|
||||
return self.patterns.shape[0]
|
||||
|
||||
def _load(self, index):
|
||||
return (index, self.translations[index]), self.patterns[index]
|
||||
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
"""Sends the relevant data to the given device and dtype"""
|
||||
super().to(*args,**kwargs)
|
||||
self.translations = self.translations.to(*args, **kwargs)
|
||||
self.patterns = self.patterns.to(*args, **kwargs)
|
||||
|
||||
|
||||
@classmethod
|
||||
def from_cxi(cls, cxi_file):
|
||||
"""Generates a new CDataset from a .cxi file directly"""
|
||||
|
||||
# Generate a base dataset
|
||||
dataset = CDataset.from_cxi(cxi_file)
|
||||
# Mutate the class to this subclass (BasicPtychoDataset)
|
||||
dataset.__class__ = cls
|
||||
|
||||
# Load the data that is only relevant for this class
|
||||
patterns, axes = cdtdata.get_data(cxi_file)
|
||||
translations = cdtdata.get_ptycho_translations(cxi_file)
|
||||
|
||||
# And now re-add it
|
||||
dataset.translations = t.Tensor(translations).clone()
|
||||
dataset.patterns = t.Tensor(patterns).clone()
|
||||
|
||||
return dataset
|
||||
|
||||
|
||||
def to_cxi(self, cxi_file):
|
||||
"""Saves out a BasicPtychoDataset as a .cxi file"""
|
||||
|
||||
super().to_cxi(cxi_file)
|
||||
cdtdata.add_data(cxi_file, self.patterns, axes=self.axes)
|
||||
cdtdata.add_ptycho_translations(cxi_file, self.translations)
|
||||
|
||||
|
||||
def inspect(self):
|
||||
"""Plots a random diffraction pattern"""
|
||||
|
||||
index = t.randint(len(self), (1,))[0]
|
||||
plt.figure()
|
||||
plt.imshow(self.patterns[index,:,:].cpu().numpy())
|
||||
Reference in New Issue
Block a user