mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-19 09:02:09 +02:00
misc small changes
This commit is contained in:
@@ -148,15 +148,13 @@ def standardize(probe, obj, obj_slice=None, correct_ramp=False):
|
||||
|
||||
|
||||
if correct_ramp:
|
||||
# Need to check if this is actually working and, if noy, why not
|
||||
center_freq = ip.centroid_sq(cmath.fftshift(t.fft(probe[0],2)),comp=True)
|
||||
# Need to check if this is actually working and, if not, why not
|
||||
center_freq = ip.centroid(cmath.cabssq(cmath.fftshift(t.fft(probe[0],2))))
|
||||
center_freq -= (t.tensor(probe[0].shape[:-1]) // 2).to(t.float32)
|
||||
center_freq /= t.tensor(probe[0].shape[:-1]).to(t.float32)
|
||||
|
||||
|
||||
|
||||
Is, Js = np.mgrid[:probe[0].shape[0],:probe[0].shape[1]]
|
||||
probe_phase_ramp = cmath.expi(2*np.pi *
|
||||
probe_phase_ramp = cmath.expi(2 * np.pi *
|
||||
(center_freq[0] * t.tensor(Is).to(t.float32) +
|
||||
center_freq[1] * t.tensor(Js).to(t.float32)))
|
||||
probe = cmath.cmult(probe, cmath.cconj(probe_phase_ramp))
|
||||
@@ -165,7 +163,6 @@ def standardize(probe, obj, obj_slice=None, correct_ramp=False):
|
||||
(center_freq[0] * t.tensor(Is).to(t.float32) +
|
||||
center_freq[1] * t.tensor(Js).to(t.float32)))
|
||||
obj = cmath.cmult(obj, obj_phase_ramp)
|
||||
|
||||
|
||||
# Then, we set them to consistent absolute phases
|
||||
|
||||
@@ -475,7 +472,7 @@ def calc_frc(im1, im2, basis, im_slice=None, nbins=None, snr=1.):
|
||||
(im1.shape[1]//8)*3:(im1.shape[1]//8)*5]
|
||||
|
||||
if nbins is None:
|
||||
nbins = np.max(synth_obj[im_slice].shape) // 4
|
||||
nbins = np.max(im1[im_slice].shape) // 4
|
||||
|
||||
|
||||
cor_fft = cmath.cmult(cmath.fftshift(t.fft(im1[im_slice],2)),
|
||||
@@ -490,7 +487,7 @@ def calc_frc(im1, im2, basis, im_slice=None, nbins=None, snr=1.):
|
||||
|
||||
i_freqs = fftpack.fftshift(fftpack.fftfreq(cor_fft.shape[0],d=di))
|
||||
j_freqs = fftpack.fftshift(fftpack.fftfreq(cor_fft.shape[1],d=dj))
|
||||
|
||||
|
||||
Js,Is = np.meshgrid(j_freqs,i_freqs)
|
||||
Rs = np.sqrt(Is**2+Js**2)
|
||||
|
||||
|
||||
@@ -278,7 +278,7 @@ def get_mask(cxi_file):
|
||||
mask = np.array(i1['detector_1/mask']).astype(np.uint32)
|
||||
mask_on = np.equal(mask,np.uint32(0))
|
||||
mask_has_signal = np.equal(mask,np.uint32(0x00001000))
|
||||
return np.logical_or(mask_on,mask_has_signal).astype(np.uint8)
|
||||
return np.logical_or(mask_on,mask_has_signal).astype(np.bool)
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
@@ -73,6 +73,7 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None,
|
||||
|
||||
# In some edge cases this shape can be smaller than the detector shape
|
||||
full_shape = t.max(full_shape, det_shape)
|
||||
|
||||
|
||||
if opt_for_fft:
|
||||
full_shape = t.Tensor([next_fast_len(dim) for dim in full_shape]).to(t.int32)
|
||||
|
||||
@@ -50,11 +50,11 @@ def amplitude_mse(intensities, sim_intensities, mask=None):
|
||||
|
||||
if mask is None:
|
||||
return t.sum((t.sqrt(sim_intensities) -
|
||||
t.sqrt(intensities))**2) / intensities.view(-1).shape[0]
|
||||
t.sqrt(intensities))**2)
|
||||
else:
|
||||
masked_intensities = intensities.masked_select(mask)
|
||||
return t.sum((t.sqrt(sim_intensities.masked_select(mask)) -
|
||||
t.sqrt(masked_intensities))**2) / masked_intensities.shape[0]
|
||||
t.sqrt(masked_intensities))**2)
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user