mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
280 lines
11 KiB
Python
280 lines
11 KiB
Python
from __future__ import division, print_function, absolute_import
|
|
|
|
import pytest
|
|
import numpy as np
|
|
from scipy import fftpack as ffts
|
|
import torch as t
|
|
from itertools import combinations
|
|
|
|
from CDTools.tools import analysis, cmath, initializers
|
|
|
|
|
|
def test_orthogonalize_probes():
|
|
|
|
# The test strategy should be to define a few non-orthogonal probes
|
|
# and orthogonalize them. Then we can test two features of the results:
|
|
|
|
# 1) Are they orthogonal?
|
|
# 2) Is the total intensity at each point the same as it was originally?
|
|
|
|
probe_xs = np.arange(128) - 64
|
|
probe_ys = np.arange(150) - 75
|
|
probe_Ys, probe_Xs = np.meshgrid(probe_ys, probe_xs)
|
|
probe_Rs = np.sqrt(probe_Xs**2 + probe_Ys**2)
|
|
|
|
probes = np.array([10*np.exp(-probe_Rs**2 / (2 * 10**2)),
|
|
3*np.exp(-probe_Rs**2 / (2 * 12**2)),
|
|
1*np.exp(-probe_Rs**2 / (2 * 15**2))]).astype(np.complex64)
|
|
|
|
# test that it works on numpy arrays
|
|
ortho_probes = analysis.orthogonalize_probes(probes)
|
|
|
|
# test that it also works on torch tensors
|
|
ortho_probes_t = cmath.torch_to_complex(analysis.orthogonalize_probes(cmath.complex_to_torch(probes)))
|
|
|
|
for p1,p2 in combinations(ortho_probes,2):
|
|
assert np.sum(np.conj(p1)*p2) / np.sum(np.abs(p1)**2) < 1e-6
|
|
|
|
for p1,p2 in combinations(ortho_probes_t,2):
|
|
assert np.sum(np.conj(p1)*p2) / np.sum(np.abs(p1)**2) < 1e-6
|
|
|
|
probe_intensity = np.sum(np.abs(probes)**2,axis=0)
|
|
ortho_probe_intensity = np.sum(np.abs(ortho_probes)**2,axis=0)
|
|
ortho_probe_t_intensity = np.sum(np.abs(ortho_probes_t)**2,axis=0)
|
|
|
|
assert np.allclose(probe_intensity,ortho_probe_intensity)
|
|
assert np.allclose(probe_intensity,ortho_probe_t_intensity)
|
|
|
|
|
|
|
|
def test_standardize():
|
|
|
|
# Start by making a probe and object that should meet the standardization
|
|
# conditions
|
|
probe = initializers.gaussian((230,240),(20,20),curvature=(0.01,0.01))
|
|
probe = cmath.torch_to_complex(probe)
|
|
probe = probe * np.sqrt(len(probe.ravel()) / np.sum(np.abs(probe)**2))
|
|
probe = probe * np.exp(-1j * np.angle(np.sum(probe)))
|
|
|
|
assert np.isclose(1, np.sum(np.abs(probe)**2)/ len(probe.ravel()))
|
|
assert np.isclose(0,np.angle(np.sum(probe)))
|
|
|
|
obj = 30 * np.random.rand(230,240) * np.exp(1j * (np.random.rand(230,240) - 0.5))
|
|
obj_slice = np.s_[(obj.shape[0]//8)*3:(obj.shape[0]//8)*5,
|
|
(obj.shape[1]//8)*3:(obj.shape[1]//8)*5]
|
|
|
|
obj = obj * np.exp(-1j * np.angle(np.sum(obj[obj_slice])))
|
|
assert np.isclose(0,np.angle(np.sum(obj[obj_slice])))
|
|
|
|
|
|
# Then make a nonstandard version of them and standardize it
|
|
# First, don't add a phase ramp and test
|
|
test_probe = probe * 37.6 * np.exp(1j*0.35)
|
|
test_obj = obj / 37.6 * np.exp(1j*1.43)
|
|
s_probe, s_obj = analysis.standardize(test_probe, test_obj)
|
|
assert np.allclose(probe, s_probe)
|
|
assert np.allclose(obj, s_obj)
|
|
|
|
# Test that it works on torch tensors
|
|
s_probe, s_obj = analysis.standardize(cmath.complex_to_torch(test_probe).to(t.float32), cmath.complex_to_torch(test_obj).to(t.float32))
|
|
s_probe = cmath.torch_to_complex(s_probe)
|
|
s_obj = cmath.torch_to_complex(s_obj)
|
|
assert np.allclose(probe, s_probe)
|
|
assert np.allclose(obj, s_obj)
|
|
|
|
|
|
# Then do one with a phase ramp
|
|
phase_ramp_dir = (np.random.rand(2) - 0.5)
|
|
|
|
probe_Xs, probe_Ys = np.mgrid[:probe.shape[0],:probe.shape[1]]
|
|
phase_ramp = np.exp(1j*probe_Ys * phase_ramp_dir[1]+
|
|
1j*probe_Xs * phase_ramp_dir[0])
|
|
test_probe = test_probe * phase_ramp
|
|
|
|
obj_Xs, obj_Ys = np.mgrid[:obj.shape[0],:obj.shape[1]]
|
|
obj_phase_ramp = np.exp(-1j*obj_Ys * phase_ramp_dir[1]+
|
|
-1j*obj_Xs * phase_ramp_dir[0])
|
|
test_obj = test_obj * obj_phase_ramp
|
|
|
|
s_probe, s_obj = analysis.standardize(test_probe, test_obj, correct_ramp=True)
|
|
|
|
assert np.max(s_probe - probe) / np.max(np.abs(probe)) < 1e-4
|
|
assert np.max(s_obj - obj) / np.max(np.abs(obj)) < 1e-4
|
|
|
|
# Finally a test with the phase ramp and multiple probes
|
|
subdominant_probe = 0.1*np.random.rand(230,240) * np.exp(1j * (np.random.rand(230,240) - 0.5))
|
|
subdominant_probe = subdominant_probe * np.exp(-1j * np.angle(np.sum(subdominant_probe)))
|
|
test_subdominant_probe = subdominant_probe * 37.6
|
|
test_subdominant_probe = test_subdominant_probe * phase_ramp
|
|
|
|
incoh_probe = np.array([test_probe,test_subdominant_probe])
|
|
|
|
s_probe, s_obj = analysis.standardize(incoh_probe, test_obj, correct_ramp=True)
|
|
|
|
assert np.max(s_probe[0] - probe) / np.max(np.abs(probe)) < 1e-4
|
|
assert np.max(s_obj - obj) / np.max(np.abs(obj)) < 1e-4
|
|
assert np.max(s_probe[1] - subdominant_probe) / np.max(np.abs(subdominant_probe)) < 1e-4
|
|
|
|
|
|
|
|
from matplotlib import pyplot as plt
|
|
def test_synthesize_reconstructions():
|
|
# I can only really test for a lack of failures, so I think my plan
|
|
# will be to create a dataset that just needs to be added and see that
|
|
# it successfully doesn't mess it up.
|
|
|
|
# Start by making a probe and object that should meet the standardization
|
|
# conditions
|
|
probe = initializers.gaussian((230,240),(20,20),curvature=(0.01,0.01))
|
|
probe = cmath.torch_to_complex(probe)
|
|
probe = probe * np.sqrt(len(probe.ravel()) / np.sum(np.abs(probe)**2))
|
|
probe = probe * np.exp(-1j * np.angle(np.sum(probe)))
|
|
|
|
assert np.isclose(1, np.sum(np.abs(probe)**2)/ len(probe.ravel()))
|
|
assert np.isclose(0,np.angle(np.sum(probe)))
|
|
|
|
obj = 30 * np.random.rand(230,240) * np.exp(1j * (np.random.rand(230,240) - 0.5))
|
|
obj_slice = np.s_[(obj.shape[0]//8)*3:(obj.shape[0]//8)*5,
|
|
(obj.shape[1]//8)*3:(obj.shape[1]//8)*5]
|
|
|
|
obj = obj * np.exp(-1j * np.angle(np.sum(obj[obj_slice])))
|
|
assert np.isclose(0,np.angle(np.sum(obj[obj_slice])))
|
|
|
|
# Now I make stacks of identical probes and objects
|
|
probes = [probe,probe,probe,probe]
|
|
probe = np.copy(probe)
|
|
objects = [obj,obj,obj,obj]
|
|
obj = np.copy(obj)
|
|
|
|
s_probe, s_obj, obj_stack = analysis.synthesize_reconstructions(probes,objects)
|
|
assert np.max(s_probe - probe) < 2e-5
|
|
assert np.max(s_obj - obj) < 2e-5
|
|
for t_obj in obj_stack:
|
|
assert np.max(t_obj - obj) < 5e-5
|
|
|
|
|
|
|
|
def test_calc_consistency_prtf():
|
|
|
|
# Create an object with a specific structure
|
|
obj = 30 * np.random.rand(1030,1040) * np.exp(1j * (np.random.rand(1030,1040) - 0.5))
|
|
|
|
#
|
|
synth_obj = np.sqrt(0.7) * obj
|
|
|
|
obj_stack = [obj,obj,obj,obj]
|
|
|
|
basis = np.array([[0,2,0],
|
|
[3,0,0]])
|
|
|
|
freqs, prtf = analysis.calc_consistency_prtf(synth_obj, obj_stack, basis)
|
|
assert np.allclose(prtf, 0.7)
|
|
|
|
freqs, prtf = analysis.calc_consistency_prtf(synth_obj, obj_stack, basis, nbins=30)
|
|
assert np.allclose(prtf, 0.7)
|
|
|
|
# Check that it also works with torch input
|
|
t_synth_obj = cmath.complex_to_torch(synth_obj)
|
|
t_obj_stack = [cmath.complex_to_torch(obj) for obj in obj_stack]
|
|
freqs, prtf = analysis.calc_consistency_prtf(t_synth_obj, t_obj_stack, basis, nbins=30)
|
|
assert np.allclose(prtf.numpy(), 0.7)
|
|
|
|
# And also when the basis is in torch
|
|
t_synth_obj = cmath.complex_to_torch(synth_obj)
|
|
t_obj_stack = [cmath.complex_to_torch(obj) for obj in obj_stack]
|
|
freqs, prtf = analysis.calc_consistency_prtf(t_synth_obj, t_obj_stack, t.Tensor(basis), nbins=30)
|
|
assert np.allclose(prtf.numpy(), 0.7)
|
|
|
|
|
|
# Check that is uses the right number of bins
|
|
assert len(prtf) == 30
|
|
assert len(freqs) == 30
|
|
|
|
# Check that the maximum frequency is correct for the basis
|
|
assert np.isclose(freqs[-1]-freqs[-2] + freqs[-1], np.sqrt(1/4**2 + 1/6**2))
|
|
|
|
|
|
def test_calc_deconvolved_cross_correlation():
|
|
|
|
obj1 = np.random.rand(200,300) + 1j * np.random.rand(200,300)
|
|
obj2 = np.random.rand(200,300) + 1j * np.random.rand(200,300)
|
|
|
|
cor_fft = np.fft.fft2(obj1) * np.conj(np.fft.fft2(obj2))
|
|
|
|
# Not sure if this is more or less stable than just the correlation
|
|
# maximum - requires some testing
|
|
np_cor = np.fft.ifft2(cor_fft / np.abs(cor_fft))
|
|
# test with numpy inputs
|
|
test_cor = analysis.calc_deconvolved_cross_correlation(obj1,obj2, im_slice=np.s_[:,:])
|
|
assert np.allclose(test_cor, np_cor)
|
|
|
|
# test with pytorch inputs
|
|
obj1_t = cmath.complex_to_torch(obj1)
|
|
obj2_t = cmath.complex_to_torch(obj2)
|
|
test_cor_t = analysis.calc_deconvolved_cross_correlation(obj1_t,obj2_t, im_slice=np.s_[:,:])
|
|
|
|
assert np.allclose(cmath.torch_to_complex(test_cor_t), np_cor)
|
|
|
|
|
|
|
|
def test_calc_frc():
|
|
|
|
obj1 = np.random.rand(270,230) + 1j * np.random.rand(270,230)
|
|
obj2 = np.random.rand(270,230) + 1j * np.random.rand(270,230)
|
|
|
|
basis = np.array([[0,2,0],
|
|
[3,0,0]])
|
|
|
|
nbins = 100
|
|
snr = 2
|
|
|
|
cor_fft = ffts.fftshift(ffts.fft2(obj1[10:-10,20:-20])) * \
|
|
ffts.fftshift(np.conj(ffts.fft2(obj2[10:-10,20:-20])))
|
|
|
|
F1 = np.abs(ffts.fftshift(ffts.fft2(obj1[10:-10,20:-20])))**2
|
|
F2 = np.abs(ffts.fftshift(ffts.fft2(obj2[10:-10,20:-20])))**2
|
|
|
|
|
|
di = np.linalg.norm(basis[:,0])
|
|
dj = np.linalg.norm(basis[:,1])
|
|
|
|
i_freqs = ffts.fftshift(ffts.fftfreq(cor_fft.shape[0],d=di))
|
|
j_freqs = ffts.fftshift(ffts.fftfreq(cor_fft.shape[1],d=dj))
|
|
|
|
Js,Is = np.meshgrid(j_freqs,i_freqs)
|
|
Rs = np.sqrt(Is**2+Js**2)
|
|
|
|
numerator, bins = np.histogram(Rs,bins=nbins,weights=cor_fft)
|
|
denominator_F1, bins = np.histogram(Rs,bins=nbins,weights=F1)
|
|
denominator_F2, bins = np.histogram(Rs,bins=nbins,weights=F2)
|
|
n_pix, bins = np.histogram(Rs,bins=nbins)
|
|
bins = bins[:-1]
|
|
|
|
frc = np.abs(numerator / np.sqrt(denominator_F1*denominator_F2))
|
|
# This moves from combined-image SNR to single-image SNR
|
|
snr /= 2
|
|
|
|
threshold = (snr + (2 * snr + 1) / np.sqrt(n_pix)) / \
|
|
(1 + snr + (2 * np.sqrt(snr)) / np.sqrt(n_pix))
|
|
|
|
test_bins, test_frc, test_threshold = analysis.calc_frc(obj1, obj2,
|
|
basis, im_slice=np.s_[10:-10,20:-20], nbins=100, snr=2)
|
|
|
|
assert np.allclose(bins, test_bins)
|
|
assert np.allclose(frc, test_frc)
|
|
assert np.allclose(threshold, test_threshold)
|
|
|
|
# try again with complex
|
|
obj1_torch = cmath.complex_to_torch(obj1)
|
|
obj2_torch = cmath.complex_to_torch(obj2)
|
|
basis_torch = t.tensor(basis)
|
|
|
|
test_bins_t, test_frc_t, test_threshold_t = analysis.calc_frc(obj1_torch,
|
|
obj2_torch,
|
|
basis_torch,
|
|
im_slice=np.s_[10:-10,20:-20], nbins=100, snr=2)
|
|
|
|
assert np.allclose(bins, test_bins_t.numpy())
|
|
assert np.allclose(frc, test_frc_t.numpy())
|
|
assert np.allclose(threshold, test_threshold_t.numpy())
|