% PTYCHO_SOLVER the main loop of ptychography. Calls the selected engine % apply additional constraints, and tries to remove ambiguities % % [outputs, fourier_error, fsc_score] = ptycho_solver(self, par, cache) % % ** self structure containing inputs: e.g. current reconstruction results, data, mask, positions, pixel size, .. % ** par structure containing parameters for the engines % ** cache structure with precalculated values to avoid unnecessary overhead % % returns: % ++ outputs self-like structure with final reconstruction % ++ fourier_error array [Npos,1] containing evolution of reconstruction error % ++ fsc_score [] or a structure with outputs from online estimation of FSC curve % % function [outputs, fourier_error, fsc_score] = ptycho_solver(self, par, cache) import engines.GPU.analysis.* import engines.GPU.shared.* import engines.GPU.initialize.* import engines.GPU.GPU_wrapper.* import math.* import utils.* import io.* % precalculate the parallel block sizes and sets [cache, par] = get_parallel_blocks(self, par, cache); global gpu use_gpu verbose( par.verbose_level ) if (nargout == 1 && verbose() == 0 && isinf(par.plot_results_every)) par.get_error = false; elseif ~isfield(par, 'get_error') par.get_error = true; end if use_gpu %modified by YJ to print out more info if isfield(par,'extraPrintInfo') verbose(struct('prefix',['GPU-',num2str(gpu.Index),'_', par.method,'_',par.extraPrintInfo])) else verbose(struct('prefix',['GPU-',num2str(gpu.Index),'_', par.method])) end %verbose(struct('prefix',['GPU-', par.method])) %PSI verbose(0,'Started solver using %s method on GPU %i', par.method, gpu.Index ) else verbose(struct('prefix',['CPU-', par.method])) verbose(0,'Started solver using %s method on CPU', par.method) end lastwarn('') fsc_score = cell(1,0); par.Nscans = length(self.reconstruct_ind); %% move everything on GPU if needed; if par.use_gpu % if not sparse solvers as ePIE, hPIE, MLs are used presplit data into % bunches (allow larger data to be processed ) %%%%split_data = is_method(par, {'MLc', 'DM'}); split_data = false; % modified by YJ, seems to avoid some errors from GPU [self, cache] = move_to_gpu(self,cache, par.keep_on_gpu, split_data); end if par.share_object && par.object_modes == 1 % enforce only a single object self.object = self.object(1,:); cache.illum_sum_0 = cache.illum_sum_0(1); end %% allocate memory fourier_error = Garray( nan(par.number_iterations, self.Npos)); if is_method(par, {'DM'}) psi_dash = cell(max(par.probe_modes,par.object_modes), length(cache.preloaded_indices_simple{1}.indices)); end if is_method(par, {'PIE', 'ML'}) if par.beta_LSQ cache.beta_object = ones(self.Npos,par.Nlayers,'single')*par.beta_object; cache.beta_probe = ones(self.Npos,par.Nlayers,'single')*par.beta_probe; else cache.beta_object = single(par.beta_object); cache.beta_probe = single(par.beta_probe); end switch lower(par.likelihood) case 'l1', cache.beta_xi = 1; % optimal step for gauss case 'poisson', cache.beta_xi = 0.5*ones(1,1,self.Npos,'single'); % for poisson it will be further refined end end %% in case of multilayer extension assume that the provided probe is positioned in middle of the sample -> shift it at the beginning if par.Nlayers > 1 && par.preshift_ML_probe probe_offset = -sum(self.z_distance(1:end-1))/2; for ii = 1:par.probe_modes self.probe{ii} = utils.prop_free_nf(self.probe{ii}, self.lambda , probe_offset ,self.pixel_size) ; end end %% in case of tilted plane ptychography, tilt the provided probe if any(par.p.sample_rotation_angles(1:2)) && check_option(par.p, 'apply_tilted_plane_correction', 'propagation') % apply propagators to the tilted plane for ii = 1:par.probe_modes self.probe{ii} = self.modes{ii}.tilted_plane_propagate_fwd(self.probe{ii}); end end global pprev; pprev = -1; mode_id = 1; % main mode (assume single most important mode for approchimations) t0 = tic; t_recon0 = tic; t_start = tic; check_recon_time = true; iter_start = 1; % object averaging for DM code for ll = 1:length(self.object) object_avg{ll} = 0; end N_object_avg = 0; %par.initial_probe_rescaling = true or false %% added by YJ update_position_weight = true; %calculate position weight the first time find_geom_correction() is called fourier_error_mean_previous = 1e10; %for convergence check if par.fourier_error_threshold 0 % PSI's old progress message %if verbose() == 0 % progressbar(iter, par.number_iterations, max(20,round(sqrt(par.number_iterations)))) %else % verbose(1,'Iteration %s: %i / %i (time %3.3g avg:%3.3g)', par.method, iter, par.number_iterations, toc(t_start), toc(t0)/(iter-1)) %end %Modified by YJ to show more details avgTimePerIter = toc(t0)/(iter-iter_start); if par.time_limit - toc(t_recon0) < (par.number_iterations-iter+1)*avgTimePerIter %if par.time_limit < inf % if reconstruction is time limited timeLeft = par.time_limit - toc(t_recon0); % recon is time-limited iterTotal = ceil(timeLeft/avgTimePerIter) + iter; %estimate total # of iteration else timeLeft = (par.number_iterations-iter+1)*avgTimePerIter; %recon is iteration-limited iterTotal = par.number_iterations; end if timeLeft>3600 verbose(0, 'Iteration: %i / %i (Time left:%3.3g hour. avg:%3.3g sec)', iter, iterTotal, timeLeft/3600, avgTimePerIter) elseif timeLeft>60 verbose(0,'Iteration: %i / %i (Time left:%3.3g min. avg:%3.3g sec)', iter, iterTotal, timeLeft/60, avgTimePerIter) else verbose(0,'Iteration: %i / %i (Time left:%3.3g sec. avg:%3.3g sec)', iter, iterTotal, timeLeft, avgTimePerIter) end end t_start = tic; if iter > 0.9*par.number_iterations && is_method(par, 'DM') for ll = 1:length(self.object) object_avg{ll} = object_avg{ll} + self.object{ll}; end N_object_avg = N_object_avg +1 ; verbose(1,'==== Averaging DM result ======') end %% GEOMETRICAL CORRECTIONS if (iter > par.probe_position_search || iter > par.detector_rotation_search) && is_method(par, {'PIE', 'ML'}) self = find_geom_correction(self,cache,par,iter,mode_id, update_position_weight); % Added by YJ if mod(iter-par.probe_position_search+1, par.update_pos_weight_every) == 0 update_position_weight = true; else update_position_weight = false; end end %% remove extra degree of freedom for OPRP and other optimizations if iter > par.probe_fourier_shift_search for kk = 1:par.Nscans ind = self.reconstruct_ind{kk}; self.modes{1}.probe_fourier_shift(ind,:) = self.modes{1}.probe_fourier_shift(ind,:) - mean(self.modes{1}.probe_fourier_shift(ind,:)); end end %% remove ambiguity related to the variable probe if par.variable_probe && iter > par.probe_change_start && is_method(par, 'ML') self = remove_variable_probe_ambiguities(self,par); end %% remove the ambiguity in the probe / object reconstruction => keep average object transmission around 1 if mod(iter,10)==1 && par.remove_object_ambiguity && ~is_used(par, {'fly_scan'}) && ~is_method(par, {'DM', 'PIE'}) % too slow for variable probe self = remove_object_ambiguity(self, cache, par) ; end if (mod(iter, 10) == 1 || iter < 5) && check_option(par, 'get_fsc_score') && ... (((par.Nscans > 1 ) && size(self.object,1) == par.Nscans) || ... ( check_option(self, 'object_orig') )) %% Fourier ring correlation between two scans with independend objects aux = online_FSC_estimate(self, par, cache, fsc_score(end,:), iter); fsc_score(end+1,1:size(aux,1), 1:size(aux,2)) = aux; end %% ADVANCED FLY SCAN. % disabled by YJ. seems redundant since it's already calculated init_solver.m % probably useful with position correction? %{ if is_used(par, 'fly_scan') if iter == 1 disp(['== AVG fly scan step ', num2str( median(sqrt(sum(diff(self.probe_positions_0).^2,1))) )]) end self = prepare_flyscan_positions(self, par); end %} %% update current probe positions (views) if iter <= 1 || iter >= par.probe_position_search %%%%%%%% crop only ROI of the full image for subsequent calculations %%%%%%%%%% for ll = 1:par.Nmodes if ~is_used(par, 'fly_scan') % && isempty(self.modes{1}.ASM_factor) %% conventional farfield/nearfield ptycho -> update view coordinates and keep subpixels shift < 1 px [cache.oROI_s{ll},cache.oROI{ll},sub_px_shift] = find_reconstruction_ROI( self.modes{1}.probe_positions,self.Np_o, self.Np_p); self.modes{1}.sub_px_shift = sub_px_shift; else % if is_used(par, 'fly_scan') %% flyscan farfield ptycho -> update view coordinates and keep subpixels shift < 1 px [cache.oROI_s{ll},cache.oROI{ll},sub_px_shift] = find_reconstruction_ROI( self.modes{1}.probe_positions,self.Np_o, self.Np_p); %% use the fftshift for much larger corrections in the case of the fly scan self.modes{ll}.sub_px_shift = self.modes{ll}.probe_positions - self.modes{1}.probe_positions_0 + sub_px_shift; % else % %% nearfield ptycho -> keep view coordinates and update subpixels shift only -> assume that positon correction was only minor % if iter <= 1 % [cache.oROI_s{ll},cache.oROI{ll},sub_px_shift] = find_reconstruction_ROI( self.modes{1}.probe_positions_0,self.Np_o, self.Np_p); % end % self.modes{ll}.sub_px_shift = self.modes{ll}.probe_positions - self.modes{1}.probe_positions_0; end end end %% update probe fft support window if mode.probe_scale_upd(end) ~= 0, important to avoid issues during subpixel probe rescaling (is pixel scale search) if self.modes{1}.probe_scale_upd(end) > 0 self.modes{1}.probe_scale_window = get_window(self.Np_p, 1+self.modes{1}.probe_scale_upd(end), 1) .* get_window(self.Np_p, 1+self.modes{1}.probe_scale_upd(end), 2); elseif self.modes{1}.probe_scale_upd(end) < 0 self.modes{1}.probe_scale_window = fftshift(get_window(self.Np_p, 1-self.modes{1}.probe_scale_upd(end), 1) .* get_window(self.Np_p, 1-self.modes{1}.probe_scale_upd(end), 2)) ; else self.modes{1}.probe_scale_window = []; end %% updated illumination if iter <= 1 || ( iter > par.probe_change_start && (mod(iter, 10) == 1 || iter < par.probe_change_start+10 )) aprobe2 = abs(self.probe{1}(:,:,1)).^2; for ll = 1:size(self.object,1) if par.share_object ind = [self.reconstruct_ind{:}]; else ind = self.reconstruct_ind{ll}; end % avoid oscilations by adding momentum term cache.illum_sum_0{ll} = set_views(cache.illum_sum_0{ll}, Garray(aprobe2), 1,1,ind, cache)/2; cache.illum_norm(ll) = norm2(cache.illum_sum_0{ll}); cache.MAX_ILLUM(ll) = max2(cache.illum_sum_0{ll}); end end %% improve convergence speed by gradient acceleration if is_method(par, 'MLc') && iter >= par.accelerated_gradients_start [self, cache] = accelerate_gradients(self, par, cache, iter); end %% %%%%%%%%%%%%%%%%%%%%%%%% PERFORM ONE ITERATION OF THE SELECTED METHOD %%%%%%%%%%%%%%%%%%%%% switch lower(par.method) case {'epie', 'hpie'} [self, cache, fourier_error] = engines.GPU.PIE(self,par,cache,fourier_error,iter); case { 'mls','mlc'} [self, cache, fourier_error] = engines.GPU.LSQML(self,par,cache,fourier_error,iter); case 'dm' [self, cache,psi_dash,fourier_error] = engines.GPU.DM(self,par,cache,psi_dash,fourier_error,iter); otherwise error('Not implemented method') end %%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%% if iter == 0 % interation 0 is used only to calibrate iinitial probe intensity iter = iter + 1; %added by YJ for while loop continue end if verbose() > 0 && any(~isnan(fourier_error(iter,:))) switch lower(par.likelihood) case 'l1', verbose(1,'===== Fourier error = %3.4g ', nanmedian(fourier_error(iter,:)) ); case 'poisson' err = fourier_error(iter,:) - fourier_error(1,:); verbose(1,'===== Log likelihood = %3.5g ', nanmedian(err)); end end %% %%%%%%%%%%%%%%%%%%%%%%%%%% convergence check added by YJ %%%%%%%%%%%%%%%%%%%%%%%%%% if any(~isnan(fourier_error(iter,:))) switch lower(par.likelihood) case 'l1' fourier_error_mean = mean(fourier_error(iter,:)); %maybe median is better? fourier_error_diff = (fourier_error_mean - fourier_error_mean_previous)/fourier_error_mean_previous; if fourier_error_diff > par.fourier_error_threshold %verbose(0, 'Convergence check is enabled.Threshold = %3.3g.', par.fourier_error_threshold) error('Current Fourier error exceeds the previous one. Terminate reconstruction.') %error('Current Fourier error exceeds the previous one.') end fourier_error_mean_previous = fourier_error_mean; case 'poisson' % err = fourier_error(iter,:) - fourier_error(1,:); % verbose(1,'===== Log likelihood = %3.5g ', nanmedian(err)); end end %% %%%%%%%%%%%%%%%%%%%%%%%%%% CORRECTIONS %%%%%%%%%%%%%%%%%%%%%%%%%%%%%% % apply probe constraints if iter >= par.probe_change_start self.probe{mode_id} = apply_probe_contraints(self.probe{mode_id}, self.modes{mode_id}); end % push low illum regions of object to zero if par.delta > 0 if iter > par.object_change_start for ll = 1:max(par.object_modes, par.Nscans) % push everywhere to zero, even out of the object region self.object{ll} = Gfun(@regular_object_out_regions, self.object{ll}, cache.illum_sum_0{ll}, cache.MAX_ILLUM(ll),par.delta); end end if iter > par.probe_change_start && ~par.variable_probe for ll = 1:par.probe_modes % push everywhere to zero, even out of the object region self.probe{ll} = self.probe{ll} .* (1-par.delta); end end end %% suppress uncontrained values !! if iter > par.object_change_start && par.object_regular(1) > 0 for kk = 1:par.Nscans for ll = 1:max(par.object_modes) self.object{kk,ll} = apply_smoothness_constraint(self.object{kk,ll},par.object_regular(1)); % blur only intensity, not phase end end end %% Added by YJ. Remove grid artifacts using Fourier windows if iter > par.object_change_start && all(par.rm_grid_artifact_step_size) for kk = 1:par.Nscans for ll = 1:par.object_modes object_temp = self.object{kk,ll}; object_temp_ph = angle(object_temp); object_temp_ph_roi = object_temp_ph(cache.object_ROI{:}); % only use the object region step_size = par.rm_grid_artifact_step_size; window_size = par.rm_grid_artifact_window_size; direction = par.rm_grid_artifact_direction; object_temp_ph_roi = remove_grid_artifact(object_temp_ph_roi, self.pixel_size(1), step_size, window_size, direction, false); object_temp_ph(cache.object_ROI{:}) = object_temp_ph_roi; self.object{kk,ll} = abs(object_temp).*exp(1i.*object_temp_ph); end end end %% Added by YJ. TV regularization on object if iter > par.object_change_start && isfield(par,'TV_lambda') && par.TV_lambda > 0 N_tv_iter = 10; for kk = 1:par.Nscans for ll = 1:par.object_modes self.object{kk,ll} = local_TV2D_chambolle(self.object{kk,ll}, par.TV_lambda, N_tv_iter); %self.object{ll}(cache.object_ROI{:}) = ... % Gfun(@local_TV2D_chambolle,self.object{ll}(cache.object_ROI{:}), par.TV_lambda, N_tv_iter); end end end %% weak positivity object if iter > par.object_change_start && any(par.positivity_constraint_object) for kk = 1:par.Nscans for ll = 1:par.object_modes self.object{kk,ll}(cache.object_ROI{:}) = ... Gfun(@positivity_constraint_object,self.object{kk,ll}(cache.object_ROI{:}), par.positivity_constraint_object); end end end %% probe orthogonalization if par.probe_modes > par.Nrec && (~is_method(par, 'DM') || iter == par.number_iterations) % orthogonalization of incoherent probe modes if is_used(par, 'fly_scan') probes = self.probe; % orthogonalize the modes with all the other shifted modes for i = 1:par.Nrec dx = mean(self.modes{i}.sub_px_shift - self.modes{1}.sub_px_shift); % apply average shift probes{i} = imshift_fft(probes{i}, dx); end probes = ortho_modes(probes); % perform othogonalization % update only the incoherent orthogonal modes self.probe(1+par.Nrec:par.probe_modes) = probes(1+par.Nrec:par.probe_modes); else %% orthogonalize the incoherent probe modes for ii = 1:par.probe_modes P(:,:,ii) = self.probe{ii}(:,:,1); end P = core.probe_modes_ortho(P); for ii = 1:par.probe_modes self.probe{ii}(:,:,1) = P(:,:,ii); end end end %% regularize multilayer reconstruction if par.regularize_layers > 0 && par.Nlayers > 1 % && mod(iter, 2) == 1 self = regulation_multilayers(self, par, cache); end %% PLOTTING %%%% plot results %%%%%%%%%%%% if par.p.use_display && mod(iter, par.plot_results_every)==0 && par.plot_results_every~=0 try if verbose() <= 0 % use cSAXS plorring rutines ptycho_plot_wrapper(self, par, fourier_error) else % use more detailed plotting rutines if (par.probe_modes > 1) %% probe incoherent modes plot_probe_modes(self,par); end if (par.Nlayers > 1) || (par.Nscans > 1 && ~par.share_object) %% object incoherent modes plot_object_modes(self, cache) end plot_results(self,cache, par, Ggather(fourier_error), ... self.modes{mode_id}.probe_positions) end % show variable modes if (par.variable_probe && par.variable_probe_modes > 0) plot_variable_probe(self, par) end % show position correction in the fourier plane if iter > par.probe_fourier_shift_search plotting.smart_figure(24654) clf hold all for ll = 1:par.Nmodes plot(self.modes{ll}.probe_fourier_shift) end hold off grid on axis tight xlabel('Position #') ylabel('Corrected probe shift in Fourier plane [px]') title('Fourier space probe shift') end % show position correction if iter > min([par.probe_position_search, par.estimate_NF_distance, par.detector_rotation_search, par.detector_scale_search]) ... && is_method(par, {'PIE', 'ML'}) plot_geom_corrections(self, self.modes{1}, Ggather(self.object{1}),iter, par, cache) if iter > min(par.detector_rotation_search, par.detector_scale_search) for i = 1:max(1,par.Nrec) verbose(1,sprintf(' Reconstruction id: %i ============= Detector pixel scale: %0.5g Detector rotation: %0.5g deg', ... i, 1-self.modes{i}.probe_scale_upd(end) , self.modes{i}.probe_rotation(end,1))) end end end if par.get_fsc_score && ~isempty(fsc_score) % plot score estimated by the fourier ring correlation plot_frc_analysis(fsc_score, par) end drawnow catch err warning(err.message) if verbose() > 1 keyboard end end end %% save intermediate images, added by YJ if isfield(par,'save_results_every') && (mod(iter, par.save_results_every ) == 0 && par.save_results_every ~=0) || iter == par.number_iterations if ~exist(par.fout, 'dir') mkdir(par.fout); end % save temporary outputs outputs = struct(); %% old format %for ll = 1:par.probe_modes % outputs.probe{ll} = Ggather(self.probe{ll}); %end %for ll = 1:par.object_modes % outputs.object{ll} = Ggather(self.object{ll}); %end %% new format for PSI's IO code probe_temp = Ggather(self.probe{1}); probe = zeros(size(probe_temp,1),size(probe_temp,2),par.probe_modes,par.variable_probe_modes+1); for ll = 1:par.probe_modes probe_temp = Ggather(self.probe{ll}); probe(:,:,ll,:) = probe_temp(:,:,1,:); end object_temp = Ggather(self.object{1}); object = zeros(size(object_temp,1),size(object_temp,2),par.Nlayers); for ll = 1:par.Nlayers %for multislice recon object_temp = Ggather(self.object{ll}); object(:,:,ll) = object_temp(:,:,1,1); end % data error if strcmp(par.likelihood, 'poisson') fourier_error_out = fourier_error- fourier_error(1,:); fourier_error_out = Ggather(fourier_error_out); else fourier_error_out = Ggather(fourier_error); end outputs.fourier_error_out = mean(fourier_error_out,2,'omitnan'); % omit nan by ZC %save the lastest error of all dp try iter_error = find(~isnan(mean(fourier_error_out,2))); iter_error = iter_error(end); outputs.fourier_error_dp = fourier_error_out(iter_error,:); clear fourier_error_out probe_temp object_temp iter_error %free up memory catch end if par.variable_probe && isfield(self,'probe_evolution') %self.probe_evolution=(S*V')'. dimension:[Nposition,eng.variable_probe_modes] %to obtain probes at each scan position: compute probe*self.probe_evolution' outputs.probe_evolution = self.probe_evolution; % end if par.get_fsc_score && ~isempty(fsc_score) outputs.fsc_score = fsc_score; end % scan positions if is_used(par, 'fly_scan') for ll=1:length(self.modes) outputs.probe_positions(:,:,ll) = self.modes{ll}.probe_positions; end else outputs.probe_positions = self.modes{1}.probe_positions; end try %probe_positions_model may not exist if geom refinement is not enabled outputs.probe_positions_model = self.modes{1}.probe_positions_model; catch end outputs.affine_matrix = self.affine_matrix{1}; if isfield(self,'affine_matrix_fit') %save input affine matrix outputs.affine_matrix_fit = self.affine_matrix_fit{1}; end if isfield(par,'affine_matrix_init') %save input affine matrix outputs.affine_matrix_init = par.affine_matrix_init; end outputs.probe_positions_0 = self.probe_positions_0; outputs.probe_positions_weight = double(self.modes{1}.probe_positions_weight); %global errors refinement results outputs.detector_rotation = self.modes{1}.probe_rotation(end,:); outputs.detector_scale = 1+self.modes{1}.probe_scale_upd(end); outputs.relative_pixel_scale = self.modes{1}.scales; outputs.rotation = self.modes{1}.rotation; outputs.shear = self.modes{1}.shear; outputs.asymmetry = self.modes{1}.asymmetry; % added by ZC outputs.z_distance = self.modes{1}.distances; % outputs.avgTimePerIter = avgTimePerIter; outputs.fourier_error_threshold = par.fourier_error_threshold; % additional parameter for PSI's IO code p = {}; p.object_ROI{1} = cache.object_ROI{1}; p.object_ROI{2} = cache.object_ROI{2}; p.binning = false; p.detector.binning = false; p.dx_spec = self.pixel_size; p.lambda = self.lambda; % store initial object file if given by an external file if isfield(par.p,'initial_iterate_object_file') && ~isempty(par.p.initial_iterate_object_file{1}) p.init_object_file = par.p.initial_iterate_object_file; end % store initial position file if isfield(par.p.scan,'custom_positions_source') && ~isempty(par.p.scan.custom_positions_source) p.init_position_file = par.p.scan.custom_positions_source; end % store initial probe if isfield(par.p,'initial_probe_file') && ~isempty(par.p.initial_probe_file) % if initial probe is given by a file p.init_probe_file = par.p.initial_probe_file; end if isfield(par.p,'normalize_init_probe') p.normalize_init_probe = par.p.normalize_init_probe; else p.normalize_init_probe = true; end if par.save_init_probe p.init_probe = probe_init; %store initial probe (after init_solver.m's pre-processing) end save(strcat(par.fout,'Niter',num2str(iter),'.mat'),'outputs','p','probe','object'); %% save object phase if any(ismember({'obj_ph','obj_ph_sum','obj_ph_stack'}, par.save_images)) object_temp = Ggather(self.object{1}); object_roi_temp = object_temp(cache.object_ROI{:},:); N_obj_roi = size(object_roi_temp); O_phase_roi = zeros(N_obj_roi(1), N_obj_roi(2), par.Nlayers); for ll=1:par.Nlayers object_temp = Ggather(self.object{ll}); O_phase_roi(:,:,ll) = phase_unwrap(angle(object_temp(cache.object_ROI{:}))); end if ismember('obj_ph_sum', par.save_images) O_phase_roi_sum = sum(O_phase_roi,3); saveName = strcat('obj_phase_roi_sum_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_phase_roi_sum/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end save_tiff_image(O_phase_roi_sum,strcat(saveDir,saveName)); end if ismember('obj_ph', par.save_images) O_phase_roi2 = zeros(N_obj_roi(1), N_obj_roi(2)*par.Nlayers); for ll=1:par.Nlayers x_lb = (ll-1)*N_obj_roi(2)+1; x_ub = ll*N_obj_roi(2); O_phase_roi2(:,x_lb:x_ub) = O_phase_roi(:,:,ll); end saveName = strcat('obj_phase_roi_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_phase_roi/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end save_tiff_image(O_phase_roi2,strcat(saveDir,saveName)); end if ismember('obj_ph_stack', par.save_images) saveName = strcat('obj_phase_roi_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_phase_roi_stack/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end imExportTiff(O_phase_roi,strcat(saveDir,saveName),'XY'); end end %% save object magnitude if any(ismember({'obj_mag','obj_ph_sum','obj_mag_stack'}, par.save_images)) object_temp = Ggather(self.object{1}); object_roi_temp = object_temp(cache.object_ROI{:},:); N_obj_roi = size(object_roi_temp); O_mag_roi = zeros(N_obj_roi(1), N_obj_roi(2), par.Nlayers); for ll=1:par.Nlayers object_temp = Ggather(self.object{ll}); O_mag_roi(:,:,ll) = abs(object_temp(cache.object_ROI{:})); end if ismember('obj_mag_sum', par.save_images) O_mag_roi_sum = prod(O_mag_roi,3); saveName = strcat('obj_mag_roi_sum_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_mag_roi_sum/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end save_tiff_image(O_mag_roi_sum,strcat(saveDir,saveName)); end if ismember('obj_mag', par.save_images) O_mag_roi2 = zeros(N_obj_roi(1), N_obj_roi(2)*par.Nlayers); for ll=1:par.Nlayers x_lb = (ll-1)*N_obj_roi(2)+1; x_ub = ll*N_obj_roi(2); O_mag_roi2(:,x_lb:x_ub) = O_mag_roi(:,:,ll); end saveName = strcat('obj_mag_roi_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_mag_roi/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end save_tiff_image(O_mag_roi2,strcat(saveDir,saveName)); end if ismember('obj_mag_stack', par.save_images) saveName = strcat('obj_mag_roi_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/obj_mag_roi_stack/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end imExportTiff(O_mag_roi,strcat(saveDir,saveName),'XY'); end end %% save probe if any(ismember({'probe_mag','probe'}, par.save_images)) probe_temp = zeros(size(probe,1),size(probe,2)*par.probe_modes); for jj=1:par.probe_modes x_lb = (jj-1)*size(probe,2)+1; x_ub = jj*size(probe,2); probe_temp(:,x_lb:x_ub) = probe(:,:,jj); end if ismember('probe_mag', par.save_images) saveName = strcat('probe_mag_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/probe_mag/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end imwrite(mat2gray(abs(probe_temp))*64,parula,strcat(saveDir,saveName),'tiff') end if ismember('probe', par.save_images) saveName = strcat('probe_Niter',num2str(iter),'.tiff'); saveDir = strcat(par.fout,'/probe/'); if ~exist(saveDir, 'dir'); mkdir(saveDir); end imwrite(convert_to_rgb(probe_temp),strcat(saveDir,saveName),'tiff') end end end iter = iter + 1; if check_recon_time && toc(t_recon0) > par.time_limit par.number_iterations = iter; % run one more iteration check_recon_time = false; end end %% return results if is_method(par, 'DM') % return average from last 10% iterations for ll = 1:length(self.object) self.object{ll} = object_avg{ll} /N_object_avg; end end iter_time = toc(t0) / par.number_iterations; verbose(1,' ==== Time per one iteration %3.3fs', iter_time) verbose(1,' ==== Total time %3.2fs', toc(t0)) % clip outliers from the low illum regions MAX_OBJ = 0; for ii = 1:length(self.object) MAX_OBJ = max(MAX_OBJ, max2(abs(self.object{ii}(cache.object_ROI{:})))); end for ii = 1:length(self.object) aobj = abs(self.object{ii}); self.object{ii} = min(MAX_OBJ, aobj) .* self.object{ii} ./ max(aobj, 1e-3); end %% in case of tilted plane ptychography, back-tilt to the detector probe if any(par.p.sample_rotation_angles(1:2)) && check_option(par.p, 'apply_tilted_plane_correction', 'propagation') % apply propagators to the tilted plane for ii = 1:par.probe_modes self.probe{ii} = self.modes{ii}.tilted_plane_propagate_back(self.probe{ii}); end end %% in case of multilayer extension assume return the reconstructed probe in middle of the sample !! if par.Nlayers > 1 && par.preshift_ML_probe probe_offset = sum(self.z_distance(1:end-1))/2; for ii = 1:par.probe_modes self.probe{ii} = utils.prop_free_nf(self.probe{ii}, self.lambda , probe_offset ,self.pixel_size) ; end end outputs = self; % avoid duplication in memory self.noise = []; self.diffraction= []; self.mask = []; % store useful parameters back to the main structure only for the 1st mode outputs.relative_pixel_scale = self.modes{1}.scales(end,:); outputs.rotation = self.modes{1}.rotation(end,:); outputs.shear = self.modes{1}.shear(end,:); outputs.z_distance = self.modes{1}.distances(end,:); outputs.shift_scans = self.modes{1}.shift_scans; outputs.probe_fourier_shift = self.modes{1}.probe_fourier_shift; outputs.probe_positions = self.modes{1}.probe_positions; outputs.detector_rotation = self.modes{1}.probe_rotation(end,:); outputs.detector_scale = 1+self.modes{1}.probe_scale_upd(end); if strcmp(par.likelihood, 'poisson') fourier_error = fourier_error - fourier_error(1,:); end fourier_error = Ggather(fourier_error); outputs.illum_sum = cache.illum_sum_0; outputs.diffraction = []; outputs.noise = []; outputs.mask = []; %% move everything back to RAM from GPU if par.use_gpu outputs = move_from_gpu(outputs); end %% report results try verbose(0,'==== REPORT ==== \n SNR %4.3g RES %3.2g (%3.3gnm) AuC: %3.3g\n\n', fsc_score{end,1}.SNR_avg,... fsc_score{end,1}.resolution,mean(self.pixel_size)*1e9/fsc_score{end,1}.resolution, fsc_score{end,1}.AUC) end end function x = positivity_constraint_object(x, relax) x = relax.*abs(x) + (1-relax).*x; end function object = regular_object_out_regions(object, illum, max_illum, delta) % push everywhere to zero, even out of the object region W = illum / max_illum; W = W ./ (0.1+W); object = object .* (W + (1-W).*(1-delta)); end function win = get_window(Np_p, scale, ax) % aux function % apodize window for img to prevent periodic boundary errors import engines.GPU.GPU_wrapper.Garray win = ones(floor(Np_p(ax)/scale/2-2)*2); win = utils.crop_pad(win, [Np_p(ax),1]); win = shiftdim(win, 1-ax); win = Garray(win); end