OSCR

Aberration-aware 3D localization microscopy via self-supervised neural-physics learning.

Code ↔ Paper

16 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 16 matches
  1. [1] § Methods › Training data generation ↔ ailoc/simulation/camera.py, lines 51–95 · score 0.74 · readout noise, electron multiplication, EMCCD camera, Poisson, simulation
  2. [2] § Methods › LUNAR architecture › Network encoder ↔ ailoc/lunar/transformer.py, lines 357–424 · score 0.69 · Transformer block, spatial dimension, aggregate, dropout, layer, masked
  3. [3] § Methods › Synchronized learning strategy ↔ ailoc/decode/loss.py, lines 35–76 · score 0.65 · Gaussian mixture model, log likelihood, GMM, loss
  4. [4] § Methods › Synchronized learning strategy ↔ ailoc/lunar/loss.py, lines 35–76 · score 0.65 · Gaussian mixture model, log likelihood, GMM, loss
  5. [5] § Results › Principle of LUNAR ↔ ailoc/lunar/network.py, lines 204–264 · score 0.62 · temporal attention module, feature extraction module, FEM, TAM, Transformer, LUNAR
  6. [6] § Methods › Microscope ↔ src/main/java/de/embl/rieslab/htsmlm/MainFramehtSMLM.java, lines 243–310 · score 0.58 · filter wheel, focus lock, SMART, beam, quadrant, TOPTICA
  7. [7] § Methods › Training data generation ↔ ailoc/common/notebook_gui.py, lines 727–804 · score 0.58 · readout noise, sCMOS, EMCCD, camera
  8. [8] § Methods › Sample preparation › Nuclear pore complex labeling ↔ usages/pyscripts/lunar_usage.py, lines 24–160 · score 0.57 · U2OS Nup96 SNAP, AF647, BG, blocked, cells
  9. [9] § Methods › LUNAR architecture › Network encoder ↔ ailoc/lunar/network.py, lines 204–264 · score 0.57 · Transformer block, dropout, FEM, TAM, layer, temporal
  10. [10] § Methods › Derivation of general objective ↔ ailoc/simulation/vectorpsf.py, lines 1514–1557 · score 0.56 · neural network, Zernike coefficients, PSF model, Derivation, optimize, objective
  11. [11] § Methods › Sample preparation › Cell culture and induced neuronal cells generation ↔ usages/pyscripts/lunar_usage.py, lines 24–160 · score 0.55 · Nup96 SNAP, U2OS, mm, cells
  12. [12] § Methods › Training data generation ↔ ailoc/simulation/mol_sampler.py, lines 230–282 · score 0.53 · predefined ranges, uniform distributions, batch, training, molecule, position
  13. [13] § Methods › Microscope ↔ ailoc/simulation/vectorpsf.py, lines 1514–1557 · score 0.52 · focal plane, microscope, NA, objective, camera, pixel
  14. [14] § Methods › PSF engineering ↔ ailoc/simulation/vectorpsf.py, lines 971–1076 · score 0.52 · optimized Zernike coefficients, pupil, wavelength, phase, nm, simulation
  15. [15] § Methods › Synchronized learning strategy ↔ ailoc/decode/loss.py, lines 35–76 · score 0.51 · negative log likelihood, GMM, loss, Gaussian, predicted, localization
  16. [16] § Methods › Synchronized learning strategy ↔ ailoc/lunar/loss.py, lines 35–76 · score 0.51 · negative log likelihood, GMM, loss, Gaussian, predicted, localization

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,560 lines · 79 KB · GPL-3.0 · 3 matches

  1. import ctypes
  2. import sys
  3. from abc import ABC, abstractmethod # abstract class
  4. import torch
  5. import torch.nn as nn
  6. import numpy as np
  7. import os
  8. from deprecated import deprecated
  9. import ailoc.common
  10. class VectorPSF(ABC):
  11. """
  12. Abstract base class for vector psf simulators
  13. """
  14. @abstractmethod
  15. def simulate(self, *args, **kwargs):
  16. """
  17. run the simulation
  18. """
  19. raise NotImplementedError
  20. @staticmethod
  21. def _gauss2D_kernel(shape=(3, 3), sigmax=0.5, sigmay=0.5, data_type=torch.float32):
  22. """
  23. 2D gaussian mask for VectorPSF otf rescale
  24. """
  25. m, n = [(ss - 1.) / 2. for ss in shape]
  26. y, x = torch.meshgrid(ailoc.common.gpu(torch.arange(-m, m + 1, 1), data_type=data_type),
  27. ailoc.common.gpu(torch.arange(-n, n + 1, 1), data_type=data_type), indexing='ij')
  28. h = torch.exp(-(x * x) / (2. * sigmax * sigmax + 1e-6) - (y * y) / (2. * sigmay * sigmay + 1e-6))
  29. torch.clamp(h, min=0)
  30. # h[h < torch.finfo(h.dtype).eps * h.max()] = 0
  31. with torch.no_grad():
  32. sumh = h.sum()
  33. output = h / sumh if sumh != 0 else h
  34. return output
  35. @staticmethod
  36. def otf_rescale(psfdata, sigma_xy):
  37. """
  38. Convolve the psf with a gaussian kernel, namely otf rescale
  39. Args:
  40. psfdata (torch.Tensor): psf data to be blured
  41. sigma_xy (torch.Tensor): sigma x y of the gaussian kernel, unit pixel
  42. Returns:
  43. torch.Tensor: blured psf data
  44. """
  45. kernel = VectorPSF._gauss2D_kernel(shape=(5, 5),
  46. sigmax=sigma_xy[0],
  47. sigmay=sigma_xy[1],
  48. data_type=psfdata.dtype).reshape((1, 1, 5, 5))
  49. psf_size = psfdata.shape[1]
  50. psfdata = psfdata.view(-1, 1, psf_size, psf_size)
  51. tmp = nn.functional.conv2d(psfdata, kernel, padding=2, stride=1)
  52. outdata = tmp.view(-1, psf_size, psf_size)
  53. return outdata
  54. class VectorPSFCUDA(VectorPSF):
  55. """
  56. Vector psf simulator using cuda, used for training due to its speed
  57. """
  58. def __init__(self, psf_params):
  59. """
  60. Args:
  61. psf_params (dict): PSF parameters for simulation except emitter positions.
  62. na: numerical aperture;
  63. wavelength: wavelength, unit nm;
  64. refmed: refractive index of the sample medium;
  65. refcov: refractive index of the coverslip;
  66. refimm: refractive index of the immersion oil;
  67. zernike_mode: zernike orders, 2D array, first column is radial order,
  68. second column is azimuthal order;
  69. zernike_coef: amplitude for each zernike mode, unit nm;
  70. zernike_coef_map: spatially variant amplitude map for each zernike mode
  71. with shape (num zernike, row, column), unit nm;
  72. objstage0: initial objStage position,relative to focus at coverslip, unit nm;
  73. zemit0: initial emitter z position, distance relative to coverslip, unit nm;
  74. pixel_size_xy: pixel sizes for x and y direction, unit nm;
  75. otf_rescale_xy: sigma x y of the gaussian kernel for otf rescale, unit pixel;
  76. npupil: number of pupil sampling points for simulation, unit pixel;
  77. psf_size: number of image pixels;
  78. Returns:
  79. VectorPSFCUDA: an instance of VectorPSFCUDA
  80. """
  81. self.na = psf_params['na']
  82. self.wavelength = psf_params['wavelength']
  83. self.refmed = psf_params['refmed']
  84. self.refcov = psf_params['refcov']
  85. self.refimm = psf_params['refimm']
  86. self.zernike_mode = psf_params['zernike_mode']
  87. try:
  88. self.zernike_coef = psf_params['zernike_coef']
  89. except KeyError:
  90. self.zernike_coef = None
  91. try:
  92. self.zernike_coef_map = psf_params['zernike_coef_map']
  93. except KeyError:
  94. self.zernike_coef_map = None
  95. if (self.zernike_coef is None) != (self.zernike_coef_map is None):
  96. pass
  97. else:
  98. raise ValueError('you must define either zernike_coef xor zernike_coef_map')
  99. self.pixel_size_xy = psf_params['pixel_size_xy']
  100. self.otf_rescale_xy = psf_params['otf_rescale_xy']
  101. self.npupil = psf_params['npupil']
  102. self.psf_size = psf_params['psf_size']
  103. self.objstage0 = psf_params['objstage0']
  104. try:
  105. self.zemit0 = psf_params['zemit0']
  106. except KeyError:
  107. self.zemit0 = -self.objstage0 / self.refimm * self.refmed
  108. thispath = os.path.dirname(os.path.abspath(__file__))
  109. self.dll_path = thispath + '/../extensions/psf_simu_gpu.dll'
  110. def simulate(self, x, y, z, photons, zernike_coefs=None):
  111. """
  112. Run the simulation to generate the vector psfs with the given positions
  113. Args:
  114. x (torch.Tensor): x positions of the psfs, unit nm
  115. y (torch.Tensor): y positions of the psfs, unit nm
  116. z (torch.Tensor): z positions of the psfs, unit nm
  117. photons (torch.Tensor): photon counts of the psfs, unit photons
  118. zernike_coefs (torch.Tensor or None): if not None, each psf can be assigned a different zernike
  119. coefficients from this array with shape (npsf, 21), otherwise use the common class
  120. property self.zernike_coef, unit nm
  121. Returns:
  122. torch.Tensor: psfs, unit photons
  123. """
  124. psf_dll = ctypes.CDLL(self.dll_path, winmode=0)
  125. class _PSFParams(ctypes.Structure):
  126. _fields_ = [
  127. ('aberrations_', ctypes.POINTER(ctypes.c_float)),
  128. ('NA_', ctypes.c_float),
  129. ('refmed_', ctypes.c_float),
  130. ('refcov_', ctypes.c_float),
  131. ('refimm_', ctypes.c_float),
  132. ('lambdaX_', ctypes.c_float),
  133. ('objStage0_', ctypes.c_float),
  134. ('zemit0_', ctypes.c_float),
  135. ('pixelSizeX_', ctypes.c_float),
  136. ('pixelSizeY_', ctypes.c_float),
  137. ('sizeX_', ctypes.c_float),
  138. ('sizeY_', ctypes.c_float),
  139. ('PupilSize_', ctypes.c_float),
  140. ('Npupil_', ctypes.c_float),
  141. ('zernikeModesN_', ctypes.c_int),
  142. ('xemit_', ctypes.POINTER(ctypes.c_float)),
  143. ('yemit_', ctypes.POINTER(ctypes.c_float)),
  144. ('zemit_', ctypes.POINTER(ctypes.c_float)),
  145. ('objstage_', ctypes.POINTER(ctypes.c_float)),
  146. ('aberrationsParas_', ctypes.POINTER(ctypes.c_float)),
  147. ('psfOut_', ctypes.POINTER(ctypes.c_float)),
  148. ('aberrationOut_', ctypes.POINTER(ctypes.c_float)),
  149. ('Nmol_', ctypes.c_int),
  150. ('showAberrationNumber_', ctypes.c_int)
  151. ]
  152. param_struct = _PSFParams()
  153. # the x y positions and pixelsize should be inversed to ensure the
  154. # input X_os corresponds to the column
  155. param_struct.Npupil_ = self.npupil
  156. param_struct.Nmol_ = x.shape[0]
  157. param_struct.NA_ = self.na
  158. param_struct.refmed_ = self.refmed
  159. param_struct.refcov_ = self.refcov
  160. param_struct.refimm_ = self.refimm
  161. param_struct.lambdaX_ = self.wavelength
  162. param_struct.objStage0_ = self.objstage0
  163. param_struct.zemit0_ = self.zemit0
  164. param_struct.pixelSizeX_ = self.pixel_size_xy[1]
  165. param_struct.pixelSizeY_ = self.pixel_size_xy[0]
  166. param_struct.zernikeModesN_ = self.zernike_mode.shape[0]
  167. param_struct.sizeX_ = self.psf_size
  168. param_struct.sizeY_ = self.psf_size
  169. param_struct.PupilSize_ = 1.0
  170. param_struct.showAberrationNumber_ = 1
  171. np_xemit = ailoc.common.cpu(y).astype('float32') # nm
  172. np_yemit = ailoc.common.cpu(x).astype('float32') # nm
  173. np_zemit = ailoc.common.cpu(z).astype('float32') # nm
  174. np_objstage = 0 * np_zemit # nm
  175. np_aberrations = np.array(np.pad(self.zernike_mode, pad_width=((0, 0), (0, 1)), mode='constant',
  176. constant_values=((0, 0), (0, 0))).flatten('F'), dtype=np.float32)
  177. param_struct.xemit_ = np_xemit.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  178. param_struct.yemit_ = np_yemit.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  179. param_struct.zemit_ = np_zemit.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  180. param_struct.objstage_ = np_objstage.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  181. param_struct.aberrations_ = np_aberrations.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  182. n_mol = param_struct.Nmol_
  183. outdata_buffer = np.empty((n_mol, self.psf_size, self.psf_size), dtype=np.float32, order='C')
  184. outpupil_buffer = np.empty((int(param_struct.Npupil_), int(param_struct.Npupil_)), dtype=np.float32, order='C')
  185. input_zernikecoefs_buffer = np.empty((n_mol, param_struct.zernikeModesN_), dtype=np.float32, order='C')
  186. if zernike_coefs is None:
  187. assert self.zernike_coef is not None, \
  188. 'both self.zernike_coef and input zernike_coefs are None, please define either of them'
  189. zernike_coefs = np.tile(self.zernike_coef, reps=(n_mol, 1)).astype('float32')
  190. else:
  191. zernike_coefs = ailoc.common.cpu(zernike_coefs).astype('float32')
  192. for n in range(0, n_mol):
  193. input_zernikecoefs_buffer[n] = zernike_coefs[n]
  194. # input_zernikecoefs_buffer = input_zernikecoefs_buffer * self.wavelength # multiply with lambda
  195. param_struct.aberrationsParas_ = input_zernikecoefs_buffer.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  196. param_struct.psfOut_ = outdata_buffer.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  197. param_struct.aberrationOut_ = outpupil_buffer.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
  198. # run simulation
  199. psf_dll.vectorPSFF1(param_struct)
  200. # check nan in the output PSFs
  201. assert not np.isnan(outdata_buffer).any(), "nan in the gpu psf, something wrong about the .dll call"
  202. psfs_out = ailoc.common.gpu(outdata_buffer)
  203. # otf rescale
  204. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1] != 0:
  205. psfs_out = self.otf_rescale(psfdata=psfs_out, sigma_xy=ailoc.common.gpu(self.otf_rescale_xy))
  206. # normalize the psf to 1, then multiply with the photon number
  207. # psfs_out /= psfs_out.sum(-1).sum(-1)[:, None, None]
  208. psfs_out *= photons[:, None, None]
  209. return psfs_out
  210. class PartialOptimizeFunction(torch.autograd.Function):
  211. @staticmethod
  212. def forward(ctx, coef, indices_to_optimize):
  213. # save needed tensors and indices
  214. ctx.save_for_backward(coef, indices_to_optimize)
  215. return coef # directly return coef in the forward pass
  216. @staticmethod
  217. def backward(ctx, grad_output):
  218. # get the saved tensors and indices
  219. coef, indices_to_optimize = ctx.saved_tensors
  220. # create a zero gradient tensor
  221. grad_input = torch.zeros_like(coef)
  222. # only assign gradients to specified indices
  223. grad_input[indices_to_optimize] = grad_output[indices_to_optimize]
  224. return grad_input, None
  225. # 包装函数
  226. def partial_optimize(coef, indices_to_optimize):
  227. return PartialOptimizeFunction.apply(coef, indices_to_optimize)
  228. class VectorPSFTorch(VectorPSF):
  229. """
  230. Vector psf simulator using pytorch, thus psf parameters can be optimized by pytorch
  231. """
  232. def __init__(self, psf_params, req_grad=False, data_type=torch.float64, zernike_idx_learn=None):
  233. """
  234. Args:
  235. psf_params (dict): PSF parameters for simulation except emitter positions
  236. req_grad (bool): whether the PSF parameters are required gradient
  237. data_type (torch.dtype): data type for the PSF parameters
  238. zernike_idx_learn (list of bool): specified zernike coefficients to learn for LUNAR SL.
  239. Returns:
  240. VectorPSFTorch: an instance of VectorPSFTorch
  241. """
  242. self.data_type = data_type
  243. if data_type == torch.float64:
  244. self.complex_type = torch.complex128
  245. elif data_type == torch.float32:
  246. self.complex_type = torch.complex64
  247. elif data_type == torch.float16:
  248. self.complex_type = torch.complex32
  249. else:
  250. raise ValueError(f'unsupported data type {data_type}')
  251. self.na = torch.tensor(psf_params['na'], device='cuda', dtype=self.data_type)
  252. self.wavelength = torch.tensor(psf_params['wavelength'], device='cuda', dtype=self.data_type)
  253. self.refmed = torch.tensor(psf_params['refmed'], device='cuda', dtype=self.data_type,)
  254. self.refcov = torch.tensor(psf_params['refcov'], device='cuda', dtype=self.data_type,)
  255. self.refimm = torch.tensor(psf_params['refimm'], device='cuda', dtype=self.data_type,)
  256. self.zernike_mode = torch.tensor(psf_params['zernike_mode'], device='cuda', dtype=self.data_type)
  257. self.zernike_coef = torch.tensor(psf_params['zernike_coef'], device='cuda', dtype=self.data_type,
  258. requires_grad=req_grad)
  259. if zernike_idx_learn is None:
  260. self.zernike_idx_learn = torch.arange(self.zernike_coef.shape[0])
  261. else:
  262. self.zernike_idx_learn = torch.tensor(zernike_idx_learn)
  263. self.zernike_coef_map = None
  264. self.objstage0 = torch.tensor(psf_params['objstage0'], device='cuda', dtype=self.data_type,)
  265. try:
  266. self.zemit0 = torch.tensor(psf_params['zemit0'], device='cuda', dtype=self.data_type, )
  267. except KeyError:
  268. self.zemit0 = torch.tensor(-psf_params['objstage0']/psf_params['refimm']*psf_params['refmed'],
  269. device='cuda',
  270. dtype=self.data_type,)
  271. self.pixel_size_xy = torch.tensor(psf_params['pixel_size_xy'], device='cuda', dtype=self.data_type)
  272. self.otf_rescale_xy = torch.tensor(psf_params['otf_rescale_xy'], device='cuda', dtype=self.data_type,)
  273. self.npupil = psf_params['npupil']
  274. self.psf_size = psf_params['psf_size']
  275. self.focus_norm = psf_params.get('focus_norm', False)
  276. self._pre_compute()
  277. @staticmethod
  278. def get_zernike(orders, xpupil, ypupil):
  279. """
  280. Calculate zernike polynomials on pupil plane
  281. Args:
  282. orders (torch.Tensor): zernike orders, 2D array, first column is radial order,
  283. second column is azimuthal order
  284. xpupil (torch.Tensor): x coordinate of pupil plane
  285. ypupil (torch.Tensor): y coordinate of pupil plane
  286. Returns:
  287. torch.Tensor: zernike polynomials
  288. """
  289. xpupil = torch.real(xpupil)
  290. ypupil = torch.real(ypupil)
  291. zersize = orders.shape
  292. Nzer = zersize[0]
  293. radormax = int(max(orders[:, 0]))
  294. azormax = int(max(abs(orders[:, 1])))
  295. [Nx, Ny] = xpupil.shape
  296. # zerpol = np.zeros( [radormax+1,azormax+1,Nx,Ny] )
  297. zerpol = torch.zeros([21, 6, Nx, Ny], device='cuda')
  298. rhosq = xpupil ** 2 + ypupil ** 2
  299. rho = torch.sqrt(rhosq)
  300. zerpol[0, 0, :, :] = torch.ones_like(xpupil)
  301. for jm in range(1, azormax + 2 + 1):
  302. m = jm - 1
  303. if m > 0:
  304. zerpol[jm - 1, jm - 1, :, :] = rho * torch.squeeze(zerpol[jm - 1 - 1, jm - 1 - 1, :, :])
  305. zerpol[jm + 2 - 1, jm - 1, :, :] = ((m + 2) * rhosq - m - 1) * torch.squeeze(zerpol[jm - 1, jm - 1, :, :])
  306. for p in range(2, radormax - m + 2 + 1):
  307. n = m + 2 * p
  308. jn = n + 1
  309. zerpol[jn - 1, jm - 1, :, :] = (2 * (n - 1) * (n * (n - 2) * (2 * rhosq - 1) - m ** 2) * torch.squeeze(
  310. zerpol[jn - 2 - 1, jm - 1, :, :]) -
  311. n * (n + m - 2) * (n - m - 2) * torch.squeeze(
  312. zerpol[jn - 4 - 1, jm - 1, :, :])) / ((n - 2) * (n + m) * (n - m))
  313. phi = torch.atan2(ypupil, xpupil)
  314. allzernikes = torch.zeros([Nzer, Nx, Ny], device='cuda')
  315. for j in range(1, Nzer + 1):
  316. n = int(orders[j - 1, 0])
  317. m = int(orders[j - 1, 1])
  318. if m >= 0:
  319. allzernikes[j - 1, :, :] = torch.squeeze(zerpol[n + 1 - 1, m + 1 - 1, :, :]) * torch.cos(m * phi)
  320. else:
  321. allzernikes[j - 1, :, :] = torch.squeeze(zerpol[n + 1 - 1, -m + 1 - 1, :, :]) * torch.sin(-m * phi)
  322. # plt.figure(constrained_layout=True)
  323. # for i in range(21):
  324. # plt.subplot(3, 7, i + 1)
  325. # plt.imshow(cpu(allzernikes[i]))
  326. # plt.show()
  327. return allzernikes
  328. def czt(self, datain, A, B, D):
  329. """
  330. Execute the chirp-z transform
  331. Args:
  332. datain (torch.Tensor): input data, 2D array
  333. A (torch.Tensor): chirp parameter, 1D array
  334. B (torch.Tensor): chirp parameter, 1D array
  335. D (torch.Tensor): chirp parameter, 1D array
  336. Returns:
  337. torch.Tensor: output data, 2D array
  338. """
  339. N = A.shape[1]
  340. M = B.shape[1]
  341. L = D.shape[1]
  342. K = datain.shape[0]
  343. # torch.repeat_interleave is too slow
  344. # t0 = time.time()
  345. # Amt = torch.repeat_interleave(A, K, 0)
  346. # Bmt = torch.repeat_interleave(B, K, 0)
  347. # Dmt = torch.repeat_interleave(D, K, 0)
  348. # print('torch czt: ', time.time() - t0)
  349. Amt = A.expand(K, N)
  350. Bmt = B.expand(K, M)
  351. Dmt = D.expand(K, L)
  352. cztin = torch.zeros([K, L], dtype=self.complex_type, device='cuda')
  353. cztin[:, 0:N] = Amt * datain
  354. tmp = Dmt * torch.fft.fft(cztin)
  355. cztout = torch.fft.ifft(tmp)
  356. dataout = Bmt * cztout[:, 0:M]
  357. return dataout
  358. def czt_parallel(self, datain, A, B, D):
  359. """
  360. Execute the chirp-z transform
  361. Args:
  362. datain (torch.Tensor): input data, 2D array
  363. A (torch.Tensor): chirp parameter, 1D array
  364. B (torch.Tensor): chirp parameter, 1D array
  365. D (torch.Tensor): chirp parameter, 1D array
  366. Returns:
  367. torch.Tensor: output data, 2D array
  368. """
  369. N = A.shape[1]
  370. M = B.shape[1]
  371. L = D.shape[1]
  372. K = datain.shape[-2]
  373. n_mol = datain.shape[-3]
  374. # torch.repeat_interleave is too slow
  375. # t0 = time.time()
  376. # Amt = torch.repeat_interleave(A, K, 0)
  377. # Bmt = torch.repeat_interleave(B, K, 0)
  378. # Dmt = torch.repeat_interleave(D, K, 0)
  379. # print('torch czt: ', time.time() - t0)
  380. Amt = A.expand(K, N)
  381. Bmt = B.expand(K, M)
  382. Dmt = D.expand(K, L)
  383. cztin = torch.zeros([2, 3, n_mol, K, L], dtype=self.complex_type, device='cuda')
  384. cztin[:, :, :, :, 0:N] = Amt[None, None, None] * datain
  385. try:
  386. tmp = Dmt * torch.fft.fft(cztin, dim=-1)
  387. except:
  388. print('fft error')
  389. cztout = torch.fft.ifft(tmp, dim=-1)
  390. dataout = Bmt[None, None, None] * cztout[:, :, :, :, 0:M]
  391. return dataout
  392. def prechirpz(self, xsize, qsize, N, M):
  393. """
  394. Calculate the auxiliary vectors for chirp-z.
  395. Args:
  396. xsize (float): normalized pupil radius 1.0
  397. qsize (float): the original pixel number that could cover the region of interest
  398. in the image plane if using FFT
  399. N (int): the sampling number on the pupil
  400. M (int): the sampling number on the region of interest on the image plane
  401. Returns:
  402. (torch.Tensor,torch.Tensor,torch.Tensor): auxiliary vectors
  403. """
  404. L = N + M - 1
  405. sigma = 2 * np.pi * xsize * qsize / N / M
  406. Afac = torch.exp(2 * 1j * sigma * (1 - M))
  407. Bfac = torch.exp(2 * 1j * sigma * (1 - N))
  408. sqW = torch.exp(2 * 1j * sigma)
  409. W = sqW ** 2
  410. # fixed phase factor and amplitude factor
  411. Gfac = (2 * xsize / N) * torch.exp(1j * sigma * (1 - N) * (1 - M))
  412. # integration about n
  413. Utmp = torch.zeros([1, N], dtype=self.complex_type, device='cuda')
  414. A = torch.zeros([1, N], dtype=self.complex_type, device='cuda')
  415. Utmp[0, 0] = sqW * Afac
  416. A[0, 0] = 1.0
  417. for i in range(1, N):
  418. A[0, i] = Utmp[0, i - 1] * A[0, i - 1]
  419. Utmp[0, i] = Utmp[0, i - 1] * W
  420. # the factor before the summation
  421. Utmp = torch.zeros([1, M], dtype=self.complex_type, device='cuda')
  422. B = torch.ones([1, M], dtype=self.complex_type, device='cuda')
  423. Utmp[0, 0] = sqW * Bfac
  424. B[0, 0] = Gfac
  425. for i in range(1, M):
  426. B[0, i] = Utmp[0, i - 1] * B[0, i - 1]
  427. Utmp[0, i] = Utmp[0, i - 1] * W
  428. # for circular convolution
  429. Utmp = torch.zeros([1, max(N, M) + 1], dtype=self.complex_type, device='cuda')
  430. Vtmp = torch.zeros([1, max(N, M) + 1], dtype=self.complex_type, device='cuda')
  431. Utmp[0, 0] = sqW
  432. Vtmp[0, 0] = 1.0
  433. # Utmp_cp = Utmp.clone()
  434. # Vtmp_cp = Vtmp.clone()
  435. for i in range(1, max(N, M) + 1):
  436. Vtmp[0, i] = Utmp[0, i - 1] * Vtmp[0, i - 1]
  437. Utmp[0, i] = Utmp[0, i - 1] * W
  438. # Vtmp[0, i] = Utmp_cp[0, i - 1] * Vtmp_cp[0, i - 1]
  439. # Utmp[0, i] = Utmp_cp[0, i - 1] * W
  440. # Vtmp_cp[0, i] = Vtmp[0, i].clone()
  441. # Utmp_cp[0, i] = Utmp[0, i].clone()
  442. D = torch.ones([1, L], dtype=self.complex_type, device='cuda')
  443. for i in range(0, M):
  444. D[0, i] = torch.conj(Vtmp[0, i])
  445. for i in range(0, N):
  446. D[0, L - 1 - i] = torch.conj(Vtmp[0, i + 1])
  447. D = torch.fft.fft(D, axis=1)
  448. return A, B, D
  449. # old version
  450. @deprecated(reason="the same as matlab code, using for loop is slow")
  451. def _pre_compute_v1(self):
  452. """
  453. Compute the common intermediate variables in advance, this can save time for PSFs simulation
  454. """
  455. # pupil radius (in diffraction units) and pupil coordinate sampling
  456. pupil_size = 1.0
  457. dxypupil = 2 * pupil_size / self.npupil
  458. xypupil = torch.arange(-pupil_size + dxypupil / 2, pupil_size, dxypupil, device='cuda', dtype=self.data_type)
  459. [xpupil, ypupil] = torch.meshgrid(xypupil, xypupil, indexing='ij')
  460. ypupil = torch.complex(ypupil, torch.zeros_like(ypupil))
  461. xpupil = torch.complex(xpupil, torch.zeros_like(xpupil))
  462. # calculation of relevant Fresnel-coefficients for the interfaces
  463. costhetamed = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refmed ** 2))
  464. costhetacov = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refcov ** 2))
  465. costhetaimm = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refimm ** 2))
  466. fresnelpmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetacov + self.refcov * costhetamed)
  467. fresnelsmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetamed + self.refcov * costhetacov)
  468. fresnelpcovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetaimm + self.refimm * costhetacov)
  469. fresnelscovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetacov + self.refimm * costhetaimm)
  470. fresnelp = fresnelpmedcov * fresnelpcovimm
  471. fresnels = fresnelsmedcov * fresnelscovimm
  472. # apodization
  473. apod = 1 / torch.sqrt(costhetaimm)
  474. # define aperture
  475. aperturemask = torch.where((xpupil ** 2 + ypupil ** 2).real < 1.0, 1.0, 0.0)
  476. self.amplitude = aperturemask * apod
  477. # setting of vectorial functions
  478. phi = torch.atan2(torch.real(ypupil), torch.real(xpupil))
  479. cosphi = torch.cos(phi)
  480. sinphi = torch.sin(phi)
  481. costheta = costhetamed
  482. sintheta = torch.sqrt(1 - costheta ** 2)
  483. pvec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  484. pvec[0] = fresnelp * costheta * cosphi
  485. pvec[1] = fresnelp * costheta * sinphi
  486. pvec[2] = -fresnelp * sintheta
  487. svec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  488. svec[0] = -fresnels * sinphi
  489. svec[1] = fresnels * cosphi
  490. svec[2] = 0 * cosphi
  491. polarizationvector = torch.empty([2, 3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  492. for ipol in range(3):
  493. polarizationvector[0, ipol] = cosphi * pvec[ipol] - sinphi * svec[ipol]
  494. polarizationvector[1, ipol] = sinphi * pvec[ipol] + cosphi * svec[ipol]
  495. self.wavevector = torch.empty([2, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  496. self.wavevector[0] = 2 * np.pi * self.na / self.wavelength * xpupil
  497. self.wavevector[1] = 2 * np.pi * self.na / self.wavelength * ypupil
  498. self.wavevectorzimm = 2 * np.pi * self.refimm / self.wavelength * costhetaimm
  499. self.wavevectorzmed = 2 * np.pi * self.refmed / self.wavelength * costhetamed
  500. # calculate aberration function
  501. waberration = torch.zeros_like(xpupil, dtype=self.complex_type, device='cuda')
  502. normfac = torch.sqrt(
  503. 2 * (self.zernike_mode[:, 0] + 1) / (1 + torch.where(self.zernike_mode[:, 1] == 0, 1.0, 0.0)))
  504. zernikecoefs_norm = self.zernike_coef * normfac
  505. allzernikes = self.get_zernike(self.zernike_mode, xpupil, ypupil)
  506. for izer in range(self.zernike_mode.shape[0]):
  507. waberration += zernikecoefs_norm[izer] * allzernikes[izer]
  508. waberration *= aperturemask
  509. self.zernike_phase = torch.exp(1j * 2 * np.pi * waberration / self.wavelength)
  510. self.pupilmatrix = torch.empty([2, 3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  511. for imat in range(2):
  512. for jmat in range(3):
  513. self.pupilmatrix[imat, jmat] = self.amplitude * self.zernike_phase * polarizationvector[imat, jmat]
  514. # czt transform(fft the pupil)
  515. xrange = self.pixel_size_xy[0] * self.psf_size / 2
  516. yrange = self.pixel_size_xy[1] * self.psf_size / 2
  517. imagesizex = xrange * self.na / self.wavelength
  518. imagesizey = yrange * self.na / self.wavelength
  519. # calculate the auxiliary vectors for chirp-z, pixelsize_xy should be inverse to match row and column
  520. self.ax, self.bx, self.dx = self.prechirpz(pupil_size, imagesizey, self.npupil, self.psf_size)
  521. self.ay, self.by, self.dy = self.prechirpz(pupil_size, imagesizex, self.npupil, self.psf_size)
  522. # calculate intensity normalization function using the PSF at focus
  523. fieldmatrix_norm = torch.empty([2, 3, self.psf_size, self.psf_size], dtype=self.complex_type, device='cuda')
  524. for itel in range(2):
  525. for jtel in range(3):
  526. Pupilfunction_norm = self.amplitude * polarizationvector[itel, jtel]
  527. inter_image_norm = torch.transpose(self.czt(Pupilfunction_norm, self.ax, self.bx, self.dx), 1, 0)
  528. fieldmatrix_norm[itel, jtel] = torch.transpose(self.czt(inter_image_norm, self.ay, self.by, self.dy), 1,
  529. 0)
  530. int_focus = torch.zeros([self.psf_size, self.psf_size], dtype=self.data_type, device='cuda')
  531. for jtel in range(3):
  532. for itel in range(2):
  533. int_focus += 1 / 3 * (torch.abs(fieldmatrix_norm[itel, jtel])) ** 2
  534. self.norm_intensity = torch.sum(int_focus)
  535. @deprecated(reason="parallel version of v1, but not compatible with simulate v3, where zernike phase "
  536. "is not computed in advance")
  537. def _pre_compute_v2(self):
  538. """
  539. Compute the common intermediate variables in advance, this can save time for PSFs simulation
  540. """
  541. # pupil radius (in diffraction units) and pupil coordinate sampling
  542. pupil_size = 1.0
  543. dxypupil = 2 * pupil_size / self.npupil
  544. xypupil = torch.arange(-pupil_size + dxypupil / 2, pupil_size, dxypupil, device='cuda', dtype=self.data_type)
  545. [xpupil, ypupil] = torch.meshgrid(xypupil, xypupil, indexing='ij')
  546. ypupil = torch.complex(ypupil, torch.zeros_like(ypupil))
  547. xpupil = torch.complex(xpupil, torch.zeros_like(xpupil))
  548. # calculation of relevant Fresnel-coefficients for the interfaces
  549. costhetamed = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refmed ** 2))
  550. costhetacov = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refcov ** 2))
  551. costhetaimm = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refimm ** 2))
  552. fresnelpmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetacov + self.refcov * costhetamed)
  553. fresnelsmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetamed + self.refcov * costhetacov)
  554. fresnelpcovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetaimm + self.refimm * costhetacov)
  555. fresnelscovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetacov + self.refimm * costhetaimm)
  556. fresnelp = fresnelpmedcov * fresnelpcovimm
  557. fresnels = fresnelsmedcov * fresnelscovimm
  558. # apodization
  559. # apod = 1 / torch.sqrt(costhetaimm) # previous version, for the simulated test dataset, should be deprecated
  560. # apod = 1 / torch.sqrt(costhetamed)
  561. apod = torch.sqrt(costhetaimm) / costhetamed # Sjoerd Stallinga version
  562. # define aperture
  563. aperturemask = torch.where((xpupil ** 2 + ypupil ** 2).real < 1.0, 1.0, 0.0)
  564. self.amplitude = aperturemask * apod
  565. # setting of vectorial functions
  566. phi = torch.atan2(torch.real(ypupil), torch.real(xpupil))
  567. cosphi = torch.cos(phi)
  568. sinphi = torch.sin(phi)
  569. costheta = costhetamed
  570. sintheta = torch.sqrt(1 - costheta ** 2)
  571. pvec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  572. pvec[0] = fresnelp * costheta * cosphi
  573. pvec[1] = fresnelp * costheta * sinphi
  574. pvec[2] = -fresnelp * sintheta
  575. svec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  576. svec[0] = -fresnels * sinphi
  577. svec[1] = fresnels * cosphi
  578. svec[2] = 0 * cosphi
  579. polarizationvector = torch.empty([2, 3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  580. polarizationvector[0,] = cosphi * pvec - sinphi * svec
  581. polarizationvector[1,] = sinphi * pvec + cosphi * svec
  582. self.wavevector = torch.empty([2, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  583. self.wavevector[0] = 2 * np.pi * self.na / self.wavelength * xpupil
  584. self.wavevector[1] = 2 * np.pi * self.na / self.wavelength * ypupil
  585. self.wavevectorzimm = 2 * np.pi * self.refimm / self.wavelength * costhetaimm
  586. self.wavevectorzmed = 2 * np.pi * self.refmed / self.wavelength * costhetamed
  587. # calculate aberration function
  588. waberration = torch.zeros_like(xpupil, dtype=self.complex_type, device='cuda')
  589. normfac = torch.sqrt(
  590. 2 * (self.zernike_mode[:, 0] + 1) / (1 + torch.where(self.zernike_mode[:, 1] == 0, 1.0, 0.0)))
  591. zernikecoefs_norm = self.zernike_coef * normfac
  592. allzernikes = self.get_zernike(self.zernike_mode, xpupil, ypupil)
  593. waberration += torch.sum(zernikecoefs_norm[:, None, None]*allzernikes, dim=0)
  594. waberration *= aperturemask
  595. self.zernike_phase = torch.exp(1j * 2 * np.pi * waberration / self.wavelength)
  596. self.pupilmatrix = torch.empty([2, 3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  597. self.pupilmatrix = self.amplitude[None, None] * self.zernike_phase[None, None] * polarizationvector
  598. # czt transform(fft the pupil)
  599. xrange = self.pixel_size_xy[0] * self.psf_size / 2
  600. yrange = self.pixel_size_xy[1] * self.psf_size / 2
  601. imagesizex = xrange * self.na / self.wavelength
  602. imagesizey = yrange * self.na / self.wavelength
  603. # calculate the auxiliary vectors for chirp-z, pixelsize_xy should be inverse to match row and column
  604. self.ax, self.bx, self.dx = self.prechirpz(pupil_size, imagesizey, self.npupil, self.psf_size)
  605. self.ay, self.by, self.dy = self.prechirpz(pupil_size, imagesizex, self.npupil, self.psf_size)
  606. # calculate intensity normalization function using the PSF at focus
  607. fieldmatrix_norm = torch.empty([2, 3, self.psf_size, self.psf_size], dtype=self.complex_type, device='cuda')
  608. Pupilfunction_norm = self.amplitude[None, None, None] * polarizationvector[:, :, None]
  609. inter_image_norm = torch.transpose(self.czt_parallel(Pupilfunction_norm, self.ax, self.bx, self.dx), -1, -2)
  610. fieldmatrix_norm = torch.transpose(self.czt_parallel(inter_image_norm, self.ay, self.by, self.dy), -1, -2)
  611. int_focus = torch.zeros([self.psf_size, self.psf_size], dtype=self.data_type, device='cuda')
  612. int_focus += 1 / 3 * torch.sum(torch.abs(fieldmatrix_norm) ** 2, dim=(0, 1, 2))
  613. self.norm_intensity = torch.sum(int_focus)
  614. def _pre_compute(self):
  615. """
  616. Compute the common intermediate variables in advance, this can save time for PSFs simulation
  617. """
  618. # pupil radius (in diffraction units) and pupil coordinate sampling
  619. pupil_size = 1.0
  620. dxypupil = 2 * pupil_size / self.npupil
  621. xypupil = torch.arange(-pupil_size + dxypupil / 2, pupil_size, dxypupil, device='cuda', dtype=self.data_type)
  622. [xpupil, ypupil] = torch.meshgrid(xypupil, xypupil, indexing='ij')
  623. ypupil = torch.complex(ypupil, torch.zeros_like(ypupil))
  624. xpupil = torch.complex(xpupil, torch.zeros_like(xpupil))
  625. # calculation of relevant Fresnel-coefficients for the interfaces
  626. costhetamed = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refmed ** 2))
  627. costhetacov = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refcov ** 2))
  628. costhetaimm = torch.sqrt(1.0 - (xpupil ** 2 + ypupil ** 2) * (self.na ** 2) / (self.refimm ** 2))
  629. fresnelpmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetacov + self.refcov * costhetamed)
  630. fresnelsmedcov = 2 * self.refmed * costhetamed / (self.refmed * costhetamed + self.refcov * costhetacov)
  631. fresnelpcovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetaimm + self.refimm * costhetacov)
  632. fresnelscovimm = 2 * self.refcov * costhetacov / (self.refcov * costhetacov + self.refimm * costhetaimm)
  633. fresnelp = fresnelpmedcov * fresnelpcovimm
  634. fresnels = fresnelsmedcov * fresnelscovimm
  635. # apodization
  636. # apod = 1 / torch.sqrt(costhetaimm) # previous version, for the simulated test dataset, should be deprecated
  637. # apod = 1 / torch.sqrt(costhetamed)
  638. apod = torch.sqrt(costhetaimm) / costhetamed # Sjoerd Stallinga version
  639. # define aperture
  640. aperturemask = torch.where((xpupil ** 2 + ypupil ** 2).real < 1.0, 1.0, 0.0)
  641. self.amplitude = aperturemask * apod
  642. # setting of vectorial functions
  643. phi = torch.atan2(torch.real(ypupil), torch.real(xpupil))
  644. cosphi = torch.cos(phi)
  645. sinphi = torch.sin(phi)
  646. costheta = costhetamed
  647. sintheta = torch.sqrt(1 - costheta ** 2)
  648. pvec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  649. pvec[0] = fresnelp * costheta * cosphi
  650. pvec[1] = fresnelp * costheta * sinphi
  651. pvec[2] = -fresnelp * sintheta
  652. svec = torch.empty([3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  653. svec[0] = -fresnels * sinphi
  654. svec[1] = fresnels * cosphi
  655. svec[2] = 0 * cosphi
  656. self.polarizationvector = torch.empty([2, 3, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  657. self.polarizationvector[0,] = cosphi * pvec - sinphi * svec
  658. self.polarizationvector[1,] = sinphi * pvec + cosphi * svec
  659. self.wavevector = torch.empty([2, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  660. self.wavevector[0] = 2 * np.pi * self.na / self.wavelength * xpupil
  661. self.wavevector[1] = 2 * np.pi * self.na / self.wavelength * ypupil
  662. self.wavevectorzimm = 2 * np.pi * self.refimm / self.wavelength * costhetaimm
  663. self.wavevectorzmed = 2 * np.pi * self.refmed / self.wavelength * costhetamed
  664. # calculate aberration function
  665. normfac = torch.sqrt(
  666. 2 * (self.zernike_mode[:, 0] + 1) / (1 + torch.where(self.zernike_mode[:, 1] == 0, 1.0, 0.0)))
  667. self.allzernikes = self.get_zernike(self.zernike_mode, xpupil, ypupil) * normfac[:, None, None] * aperturemask[None]
  668. # czt transform(fft the pupil)
  669. xrange = self.pixel_size_xy[0] * self.psf_size / 2
  670. yrange = self.pixel_size_xy[1] * self.psf_size / 2
  671. imagesizex = xrange * self.na / self.wavelength
  672. imagesizey = yrange * self.na / self.wavelength
  673. # calculate the auxiliary vectors for chirp-z, pixelsize_xy should be inverse to match row and column
  674. self.ax, self.bx, self.dx = self.prechirpz(pupil_size, imagesizey, self.npupil, self.psf_size)
  675. self.ay, self.by, self.dy = self.prechirpz(pupil_size, imagesizex, self.npupil, self.psf_size)
  676. # calculate intensity normalization function using the PSF at focus
  677. pupilfunction_norm = self.amplitude[None, None, None] * self.polarizationvector[:, :, None]
  678. inter_image_norm = torch.transpose(self.czt_parallel(pupilfunction_norm, self.ax, self.bx, self.dx), -1, -2)
  679. fieldmatrix_norm = torch.transpose(self.czt_parallel(inter_image_norm, self.ay, self.by, self.dy), -1, -2)
  680. int_focus = torch.zeros([self.psf_size, self.psf_size], dtype=self.data_type, device='cuda')
  681. int_focus += 1 / 3 * torch.sum(torch.abs(fieldmatrix_norm) ** 2, dim=(0, 1, 2))
  682. self.norm_intensity = torch.sum(int_focus)
  683. @deprecated(reason="the same as matlab code, using for loop is slow")
  684. def simulate_v1(self, x, y, z, photons, objstage=None):
  685. """
  686. Run the simulation to generate the vector PSFs with the given positions
  687. Args:
  688. x (torch.Tensor): x positions of the PSFs, unit nm
  689. y (torch.Tensor): y positions of the PSFs, unit nm
  690. z (torch.Tensor): z positions of the PSFs, unit nm
  691. photons (torch.Tensor): photon counts of the PSFs, unit photons
  692. objstage (torch.Tensor): objective stage positions relative to the cover-slip (0),
  693. the closer to the sample, the smaller this value is (-), unit nm
  694. Returns:
  695. torch.Tensor: PSFs, unit photons
  696. """
  697. n_mol = x.shape[0]
  698. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda') if objstage is None else objstage
  699. field_matrix = torch.empty([2, 3, n_mol, self.psf_size, self.psf_size],
  700. dtype=self.complex_type, device='cuda')
  701. for jz in range(n_mol):
  702. # xyz induced phase, x,y should be inverse to match the column, row
  703. if z[jz] + self.zemit0 >= 0:
  704. phase_xyz = -y[jz] * self.wavevector[0] - x[jz] * self.wavevector[1] + \
  705. (z[jz] + self.zemit0) * self.wavevectorzmed
  706. position_phase = torch.exp(1j * (phase_xyz + (objstage[jz] + self.objstage0) *
  707. self.wavevectorzimm))
  708. else:
  709. # print("warning! the emitter's position may not have physical meaning")
  710. phase_xyz = -y[jz] * self.wavevector[0] - x[jz] * self.wavevector[1]
  711. position_phase = torch.exp(1j * (phase_xyz + (objstage[jz] + self.objstage0 + z[jz]
  712. + self.zemit0) * self.wavevectorzimm))
  713. for itel in range(2):
  714. for jtel in range(3):
  715. pupil_tmp = position_phase * self.pupilmatrix[itel, jtel]
  716. inter_image = torch.transpose(self.czt(pupil_tmp, self.ay, self.by, self.dy), 1, 0)
  717. field_matrix[itel, jtel, jz] = torch.transpose(self.czt(inter_image, self.ax, self.bx, self.dx), 1,
  718. 0)
  719. psfs_out = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  720. for jz in range(n_mol):
  721. for jtel in range(3):
  722. for itel in range(2):
  723. psfs_out[jz, :, :] += 1 / 3 * (torch.abs(field_matrix[itel, jtel, jz])) ** 2
  724. psfs_out /= self.norm_intensity
  725. # otf rescale
  726. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1] != 0:
  727. psfs_out = self.otf_rescale(psfdata=psfs_out, sigma_xy=self.otf_rescale_xy)
  728. # normalize the psf to 1, then multiply with the photon number
  729. # psfs_out /= psfs_out.sum(-1).sum(-1)[:, None, None]
  730. psfs_out *= photons[:, None, None]
  731. return psfs_out
  732. @deprecated(reason="parallel version of v1, but not compatible with simulate v3, where zernike phase "
  733. "is not computed in advance")
  734. def simulate_v2(self, x, y, z, photons, objstage=None):
  735. """
  736. Run the simulation to generate the vector PSFs with the given positions
  737. Args:
  738. x (torch.Tensor): x positions of the PSFs, unit nm
  739. y (torch.Tensor): y positions of the PSFs, unit nm
  740. z (torch.Tensor): z positions of the PSFs, unit nm
  741. photons (torch.Tensor): photon counts of the PSFs, unit photons
  742. objstage (torch.Tensor): objective stage positions relative to the cover-slip (0),
  743. the closer to the sample, the smaller this value is (-), unit nm
  744. Returns:
  745. torch.Tensor: PSFs, unit photons
  746. """
  747. if not hasattr(self, "device"):
  748. self.device = x.device
  749. n_mol = x.shape[0]
  750. if n_mol == 0:
  751. return torch.zeros([0, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  752. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda') if objstage is None else objstage
  753. # batch_size in parallel to save GPU memory
  754. slice_list = []
  755. batch_size = 100
  756. for i in np.arange(0, n_mol, batch_size):
  757. slice_list.append(slice(i, min(i + batch_size, n_mol)))
  758. psfs_out = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  759. for slice_tmp in slice_list:
  760. length_tmp = slice_tmp.stop - slice_tmp.start
  761. position_phase = torch.empty([length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  762. idx = torch.where(z[slice_tmp] + self.zemit0 >= 0)[0]
  763. phase_xyz_tmp = -y[slice_tmp][idx][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx][:, None, None] * \
  764. self.wavevector[1][None] + \
  765. (z[slice_tmp][idx] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  766. position_phase[idx, :, :] = torch.exp(
  767. 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx][:, None, None] + self.objstage0) *
  768. self.wavevectorzimm[None]))
  769. idx = torch.where(z[slice_tmp] + self.zemit0 < 0)[0]
  770. phase_xyz_tmp = -y[slice_tmp][idx][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx][:, None, None] * \
  771. self.wavevector[1][None]
  772. position_phase[idx, :, :] = torch.exp(
  773. 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx][:, None, None] + self.objstage0 + z[slice_tmp][idx][:, None, None]
  774. + self.zemit0) * self.wavevectorzimm[None]))
  775. pupil_tmp = position_phase[None, None] * self.pupilmatrix[:, :, None]
  776. inter_image = torch.transpose(self.czt_parallel(pupil_tmp, self.ay, self.by, self.dy), -1, -2)
  777. field_matrix = torch.transpose(self.czt_parallel(inter_image, self.ax, self.bx, self.dx), -1, -2)
  778. psfs_out[slice_tmp] += 1 / 3 * torch.sum((torch.abs(field_matrix[:, :])) ** 2, dim=(0, 1))
  779. # # all in parallel, but may cause GPU memory overflow
  780. # position_phase = torch.empty([n_mol, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  781. #
  782. # idx = torch.where(z + self.zemit0 >= 0)[0]
  783. # phase_xyz_tmp = -y[idx][:, None, None] * self.wavevector[0][None] - x[idx][:, None, None] * self.wavevector[1][None] + \
  784. # (z[idx] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  785. # position_phase[idx, :, :] = torch.exp(1j * (phase_xyz_tmp + (objstage[idx][:, None, None] + self.objstage0) *
  786. # self.wavevectorzimm[None]))
  787. #
  788. # idx = torch.where(z + self.zemit0 < 0)[0]
  789. # phase_xyz_tmp = -y[idx][:, None, None] * self.wavevector[0][None] - x[idx][:, None, None] * self.wavevector[1][None]
  790. # position_phase[idx, :, :] = torch.exp(1j * (phase_xyz_tmp + (objstage[idx][:, None, None] + self.objstage0 + z[idx][:, None, None]
  791. # + self.zemit0) * self.wavevectorzimm[None]))
  792. #
  793. # pupil_tmp = position_phase[None, None] * self.pupilmatrix[:, :, None]
  794. # inter_image = torch.transpose(self.czt_parallel(pupil_tmp, self.ay, self.by, self.dy), -1, -2)
  795. # field_matrix = torch.transpose(self.czt_parallel(inter_image, self.ax, self.bx, self.dx), -1, -2)
  796. #
  797. # psfs_out = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  798. # psfs_out += 1 / 3 * torch.sum((torch.abs(field_matrix[:, :])) ** 2, dim=(0, 1))
  799. # intensity normalization
  800. psfs_out /= self.norm_intensity
  801. # otf rescale
  802. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1] != 0:
  803. psfs_out = self.otf_rescale(psfdata=psfs_out, sigma_xy=self.otf_rescale_xy)
  804. # normalize the psf to 1, then multiply with the photon number
  805. # psfs_out /= psfs_out.sum(-1).sum(-1)[:, None, None]
  806. psfs_out *= photons[:, None, None]
  807. return psfs_out
  808. def simulate(self, x, y, z, photons, objstage=None, zernike_coefs=None):
  809. """
  810. Run the simulation to generate the vector PSFs with the given positions
  811. Args:
  812. x (torch.Tensor): x positions of the PSFs, unit nm
  813. y (torch.Tensor): y positions of the PSFs, unit nm
  814. z (torch.Tensor): z positions of the PSFs, unit nm
  815. photons (torch.Tensor): photon counts of the PSFs, unit photons
  816. objstage (torch.Tensor): objective stage positions relative to the cover-slip (0),
  817. the closer to the sample, the smaller this value is (-), unit nm
  818. zernike_coefs (torch.Tensor or None): if not None, each psf can be assigned a different zernike
  819. coefficients from this array with shape (npsf, 21), otherwise use the common class
  820. property self.zernike_coef, unit nm
  821. Returns:
  822. torch.Tensor: PSFs, unit photons
  823. """
  824. n_mol = x.shape[0]
  825. if n_mol == 0:
  826. return torch.zeros([0, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  827. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda') if objstage is None else objstage
  828. # # temporally test, change the z to obj
  829. # objstage = z
  830. # z = torch.zeros_like(z, dtype=self.data_type, device='cuda')
  831. if zernike_coefs is None:
  832. try:
  833. # for partially optimized zernike coefficients in LUNAR SL physics learning
  834. zernike_coef_part_optm = partial_optimize(self.zernike_coef, self.zernike_idx_learn)
  835. zernike_phase = torch.exp(1j * 2 * np.pi *
  836. torch.sum(zernike_coef_part_optm[:, None, None] * self.allzernikes, dim=0)
  837. / self.wavelength)
  838. except:
  839. zernike_phase = torch.exp(1j * 2 * np.pi *
  840. torch.sum(self.zernike_coef[:, None, None] * self.allzernikes, dim=0)
  841. / self.wavelength)
  842. pupilmatrix = (self.amplitude[None, None] *
  843. zernike_phase[None, None] *
  844. self.polarizationvector)[:, :, None]
  845. # batch_size in parallel to save GPU memory
  846. slice_list = []
  847. batch_size = 100
  848. for i in np.arange(0, n_mol, batch_size):
  849. slice_list.append(slice(i, min(i + batch_size, n_mol)))
  850. psfs_out = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=torch.float32)
  851. for slice_tmp in slice_list:
  852. if zernike_coefs is not None:
  853. zernike_phase = torch.exp(1j * 2 * np.pi *
  854. torch.sum(zernike_coefs[slice_tmp, :, None, None] * self.allzernikes[None], dim=1)
  855. / self.wavelength)
  856. pupilmatrix = (zernike_phase[None, None] *
  857. self.polarizationvector[:, :, None] *
  858. self.amplitude[None, None, None])
  859. # length_tmp = slice_tmp.stop - slice_tmp.start
  860. # position_phase = torch.empty([length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  861. # idx = torch.where(z[slice_tmp] + self.zemit0 >= 0)[0]
  862. # phase_xyz_tmp = -y[slice_tmp][idx][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx][:, None, None] * \
  863. # self.wavevector[1][None] + \
  864. # (z[slice_tmp][idx] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  865. # position_phase[idx, :, :] = torch.exp(
  866. # 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx][:, None, None] + self.objstage0) *
  867. # self.wavevectorzimm[None]))
  868. #
  869. # idx = torch.where(z[slice_tmp] + self.zemit0 < 0)[0]
  870. # phase_xyz_tmp = -y[slice_tmp][idx][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx][:, None, None] * \
  871. # self.wavevector[1][None]
  872. # position_phase[idx, :, :] = torch.exp(
  873. # 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx][:, None, None] + self.objstage0 + z[slice_tmp][idx][:, None, None]
  874. # + self.zemit0) * self.wavevectorzimm[None]))
  875. phase_xyz_tmp = -y[slice_tmp][:, None, None] * self.wavevector[0][None] - x[slice_tmp][:, None, None] * \
  876. self.wavevector[1][None] + (z[slice_tmp] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  877. position_phase = torch.exp(1j * (phase_xyz_tmp + (objstage[slice_tmp][:, None, None] + self.objstage0) *
  878. self.wavevectorzimm[None]))
  879. pupil_tmp = position_phase[None, None] * pupilmatrix
  880. inter_image = torch.transpose(self.czt_parallel(pupil_tmp, self.ay, self.by, self.dy), -1, -2)
  881. field_matrix = torch.transpose(self.czt_parallel(inter_image, self.ax, self.bx, self.dx), -1, -2)
  882. psfs_out[slice_tmp] += 1 / 3 * torch.sum((torch.abs(field_matrix[:, :])) ** 2, dim=(0, 1))
  883. if self.focus_norm:
  884. # intensity normalization by focus
  885. psfs_out /= self.norm_intensity
  886. else:
  887. # normalize by themselves
  888. norm_factor = psfs_out.sum(dim=(-1, -2))
  889. psfs_out /= norm_factor[:, None, None]
  890. # otf rescale
  891. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1]:
  892. psfs_out = self.otf_rescale(psfdata=psfs_out, sigma_xy=self.otf_rescale_xy)
  893. # multiply with the photon number
  894. psfs_out *= photons[:, None, None]
  895. return psfs_out
  896. def compute_crlb(self, x, y, z, photons, bgs):
  897. """
  898. Calculate the CRLB of this PSF model at give positions, photons and backgrounds.
  899. Args:
  900. x (torch.Tensor): x positions of the PSFs, unit nm
  901. y (torch.Tensor): y positions of the PSFs, unit nm
  902. z (torch.Tensor): z positions of the PSFs, unit nm
  903. photons (torch.Tensor): photon counts of the PSFs, unit photons
  904. bgs (torch.Tensor): background counts of the PSFs, unit photons
  905. Returns:
  906. (torch.Tensor,torch.Tensor): CRLB xyz (nPSFs, 3) and model PSFs, unit nm and photons
  907. """
  908. n_mol = x.shape[0]
  909. # calculate the derivatives
  910. [dudt, model] = self._compute_derivative(x, y, z, photons, bgs)
  911. # dudt = self._torch_jacobian_derivative(x, y, z, photons, bgs)
  912. # calculate hessian matrix, here only consider not shared parameters: x, y, z, photons, background
  913. num_pars = n_mol * 5
  914. t2 = 1 / model
  915. hessian = torch.zeros([num_pars, num_pars], device='cuda', dtype=self.data_type)
  916. for p1 in range(num_pars):
  917. temp1_zind = int(np.floor(p1 / 5)) # the index of data
  918. temp1_pind = int(p1 % 5) # the index of parameter type
  919. temp1 = dudt[:, :, :, temp1_pind] # the derivative of data concerning this parameter type
  920. for p2 in range(p1, num_pars):
  921. temp2_zind = int(np.floor(p2 / 5))
  922. temp2_pind = int(p2 % 5)
  923. temp2 = dudt[:, :, :, temp2_pind]
  924. # since all parameters are not shared, only the same molecule data makes sense
  925. # when multiply gradients of two parameters
  926. if temp1_zind == temp2_zind:
  927. temp = t2[temp1_zind, :, :] * temp1[temp1_zind, :, :] * temp2[temp2_zind, :, :]
  928. hessian[p1, p2] = torch.sum(temp)
  929. hessian[p2, p1] = hessian[p1, p2]
  930. # calculate local fisher matrix and crlb
  931. xyz_crlb = torch.zeros([n_mol, 3], device='cuda')
  932. for j in range(n_mol):
  933. fisher_tmp = hessian[j * 5:j * 5 + 5, j * 5:j * 5 + 5]
  934. sqrt_crlb_tmp = torch.sqrt(torch.diag(torch.inverse(fisher_tmp)))
  935. xyz_crlb[j] = sqrt_crlb_tmp[0:3]
  936. return xyz_crlb, model
  937. def compute_crlb_mf(self, x, y, z, photons, bgs, attn_length):
  938. #todo: need test
  939. """
  940. Calculate the CRLB of this PSF model at give positions, photons and backgrounds.
  941. Args:
  942. x (torch.Tensor): x positions of the PSFs, unit nm
  943. y (torch.Tensor): y positions of the PSFs, unit nm
  944. z (torch.Tensor): z positions of the PSFs, unit nm
  945. photons (torch.Tensor): photon counts of the PSFs, unit photons
  946. bgs (torch.Tensor): background counts of the PSFs, unit photons
  947. attn_length (int): attention length of the network, used to multiply the Fisher matrix
  948. Returns:
  949. (torch.Tensor,torch.Tensor): CRLB xyz (nPSFs, 3) and model PSFs, unit nm and photons
  950. """
  951. n_mol = x.shape[0]
  952. # calculate the derivatives
  953. [dudt, model] = self._compute_derivative_parallel(x, y, z, photons, bgs)
  954. # calculate hessian matrix, here only consider not shared parameters: x, y, z, photons, background
  955. num_pars = n_mol * 5
  956. t2 = 1 / model
  957. hessian = torch.zeros([num_pars, num_pars], device='cuda', dtype=self.data_type)
  958. for p1 in range(num_pars):
  959. temp1_zind = int(np.floor(p1 / 5)) # the index of data
  960. temp1_pind = int(p1 % 5) # the index of parameter type
  961. temp1 = dudt[:, :, :, temp1_pind] # the derivative of data concerning this parameter type
  962. for p2 in range(p1, num_pars):
  963. temp2_zind = int(np.floor(p2 / 5))
  964. temp2_pind = int(p2 % 5)
  965. temp2 = dudt[:, :, :, temp2_pind]
  966. # since all parameters are not shared, only the same molecule data makes sense
  967. # when multiply gradients of two parameters
  968. if temp1_zind == temp2_zind:
  969. temp = t2[temp1_zind, :, :] * temp1[temp1_zind, :, :] * temp2[temp2_zind, :, :]
  970. hessian[p1, p2] = torch.sum(temp)
  971. hessian[p2, p1] = hessian[p1, p2]
  972. # calculate local fisher matrix and crlb
  973. xyz_crlb = torch.zeros([n_mol, 3], device='cuda')
  974. for j in range(n_mol):
  975. fisher_tmp = hessian[j * 5:j * 5 + 5, j * 5:j * 5 + 5]
  976. fisher_tmp[:3, :3] *= attn_length
  977. sqrt_crlb_tmp = torch.sqrt(torch.diag(torch.inverse(fisher_tmp)))
  978. xyz_crlb[j] = sqrt_crlb_tmp[0:3]
  979. return xyz_crlb, model
  980. @deprecated(reason='the same as matlab code, using for loop is slow')
  981. def _compute_derivative_v1(self, x, y, z, photons, bgs):
  982. """
  983. Calculate the analytical derivatives of the PSFs at given parameters with respect to x,y,z,photons,bg
  984. """
  985. n_mol = x.shape[0]
  986. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda')
  987. field_matrix = torch.empty([2, 3, n_mol, self.psf_size, self.psf_size],
  988. dtype=self.complex_type, device='cuda')
  989. field_matrix_ders = torch.empty([2, 3, n_mol, 3, self.psf_size, self.psf_size],
  990. dtype=self.complex_type, device='cuda')
  991. for jz in range(n_mol):
  992. # xyz induced phase
  993. if z[jz] + self.zemit0 >= 0:
  994. phase_xyz = -y[jz] * self.wavevector[0] - x[jz] * self.wavevector[1] + \
  995. (z[jz] + self.zemit0) * self.wavevectorzmed
  996. position_phase = torch.exp(1j * (phase_xyz + (objstage[jz] + self.objstage0) *
  997. self.wavevectorzimm))
  998. else:
  999. # print("warning! the emitter's position may not have physical meaning")
  1000. phase_xyz = -y[jz] * self.wavevector[0] - x[jz] * self.wavevector[1]
  1001. position_phase = torch.exp(1j * (phase_xyz + (objstage[jz] + self.objstage0 + z[jz]
  1002. + self.zemit0) * self.wavevectorzimm))
  1003. for itel in range(2):
  1004. for jtel in range(3):
  1005. pupil_tmp = position_phase * self.pupilmatrix[itel, jtel]
  1006. inter_image = torch.transpose(self.czt(pupil_tmp, self.ay, self.by, self.dy), 1, 0)
  1007. field_matrix[itel, jtel, jz] = torch.transpose(self.czt(inter_image, self.ax, self.bx, self.dx), 1,
  1008. 0)
  1009. # derivatives with respect to x,y,z
  1010. pupilfunction_x = -1j * self.wavevector[1] * position_phase * self.pupilmatrix[itel, jtel]
  1011. inter_image_x = torch.transpose(self.czt(pupilfunction_x, self.ay, self.by, self.dy), 1, 0)
  1012. field_matrix_ders[itel, jtel, jz, 0] = torch.transpose(self.czt(inter_image_x, self.ax,
  1013. self.bx, self.dx), 1, 0)
  1014. pupilfunction_y = -1j * self.wavevector[0] * position_phase * self.pupilmatrix[itel, jtel]
  1015. inter_image_y = torch.transpose(self.czt(pupilfunction_y, self.ay, self.by, self.dy), 1, 0)
  1016. field_matrix_ders[itel, jtel, jz, 1] = torch.transpose(self.czt(inter_image_y, self.ax,
  1017. self.bx, self.dx), 1, 0)
  1018. if z[jz] + self.zemit0 >= 0:
  1019. pupilfunction_z = 1j * self.wavevectorzmed * position_phase * self.pupilmatrix[itel, jtel]
  1020. else:
  1021. pupilfunction_z = 1j * self.wavevectorzimm * position_phase * self.pupilmatrix[itel, jtel]
  1022. inter_image_z = torch.transpose(self.czt(pupilfunction_z, self.ay, self.by, self.dy), 1, 0)
  1023. field_matrix_ders[itel, jtel, jz, 2] = torch.transpose(self.czt(inter_image_z, self.ax,
  1024. self.bx, self.dx), 1, 0)
  1025. psfs = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=self.data_type)
  1026. psfs_ders = torch.zeros([n_mol, self.psf_size, self.psf_size, 3], device='cuda', dtype=self.data_type)
  1027. for jz in range(n_mol):
  1028. for jtel in range(3):
  1029. for itel in range(2):
  1030. psfs[jz, :, :] += 1 / 3 * (torch.abs(field_matrix[itel, jtel, jz])) ** 2
  1031. for jder in range(3):
  1032. psfs_ders[jz, :, :, jder] = psfs_ders[jz, :, :, jder] + 2 / 3 * \
  1033. torch.real(torch.conj(field_matrix[itel, jtel, jz]) *
  1034. field_matrix_ders[itel, jtel, jz, jder])
  1035. psfs /= self.norm_intensity
  1036. psfs_ders /= self.norm_intensity
  1037. # otf rescale
  1038. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1] != 0:
  1039. psfs = self.otf_rescale(psfdata=psfs, sigma_xy=self.otf_rescale_xy)
  1040. psfs_ders[:, :, :, 0] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 0], sigma_xy=self.otf_rescale_xy)
  1041. psfs_ders[:, :, :, 1] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 1], sigma_xy=self.otf_rescale_xy)
  1042. psfs_ders[:, :, :, 2] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 2], sigma_xy=self.otf_rescale_xy)
  1043. psfs_out = ailoc.common.gpu(psfs * photons[:, None, None] + bgs[:, None, None])
  1044. ders_out = torch.zeros([n_mol, self.psf_size, self.psf_size, 5], device='cuda', dtype=self.data_type)
  1045. ders_out[:, :, :, 0:3] = psfs_ders * photons[:, None, None, None]
  1046. ders_out[:, :, :, 3] = psfs
  1047. ders_out[:, :, :, 4] = torch.ones_like(psfs)
  1048. return ders_out, psfs_out
  1049. @deprecated(reason='parallel version of v1, but not compatible with the _pre_compute_v3 and simulate_v3')
  1050. def _compute_derivative_v2(self, x, y, z, photons, bgs):
  1051. """
  1052. Calculate the analytical derivatives of the PSFs at given parameters with respect to x,y,z,photons,bg
  1053. """
  1054. n_mol = x.shape[0]
  1055. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda')
  1056. # batch_size in parallel to save GPU memory
  1057. slice_list = []
  1058. batch_size = 100
  1059. for i in np.arange(0, n_mol, batch_size):
  1060. slice_list.append(slice(i, min(i + batch_size, n_mol)))
  1061. psfs = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=self.data_type)
  1062. psfs_ders = torch.zeros([n_mol, self.psf_size, self.psf_size, 3], device='cuda', dtype=self.data_type)
  1063. for slice_tmp in slice_list:
  1064. length_tmp = slice_tmp.stop - slice_tmp.start
  1065. position_phase = torch.empty([length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  1066. idx_0 = torch.where(z[slice_tmp] + self.zemit0 >= 0)[0]
  1067. phase_xyz_tmp = -y[slice_tmp][idx_0][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx_0][:, None, None] * \
  1068. self.wavevector[1][None] + \
  1069. (z[slice_tmp][idx_0] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  1070. position_phase[idx_0, :, :] = torch.exp(
  1071. 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx_0][:, None, None] + self.objstage0) *
  1072. self.wavevectorzimm[None]))
  1073. idx_1 = torch.where(z[slice_tmp] + self.zemit0 < 0)[0]
  1074. phase_xyz_tmp = -y[slice_tmp][idx_1][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx_1][:, None, None] * \
  1075. self.wavevector[1][None]
  1076. position_phase[idx_1, :, :] = torch.exp(
  1077. 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx_1][:, None, None] + self.objstage0 + z[slice_tmp][idx_1][:, None, None]
  1078. + self.zemit0) * self.wavevectorzimm[None]))
  1079. pupil_tmp = position_phase[None, None] * self.pupilmatrix[:, :, None]
  1080. inter_image = torch.transpose(self.czt_parallel(pupil_tmp, self.ay, self.by, self.dy), -1, -2)
  1081. field_matrix = torch.transpose(self.czt_parallel(inter_image, self.ax, self.bx, self.dx), -1, -2)
  1082. psfs[slice_tmp] += 1 / 3 * torch.sum((torch.abs(field_matrix[:, :])) ** 2, dim=(0, 1))
  1083. # derivatives with respect to x,y,z
  1084. pupil_tmp_x = -1j * self.wavevector[1] * position_phase[None, None] * self.pupilmatrix[:, :, None]
  1085. inter_image_x = torch.transpose(self.czt_parallel(pupil_tmp_x, self.ay, self.by, self.dy), -1, -2)
  1086. field_matrix_x = torch.transpose(self.czt_parallel(inter_image_x, self.ax, self.bx, self.dx), -1, -2)
  1087. psfs_ders[slice_tmp, :, :, 0] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_x),
  1088. dim=(0, 1))
  1089. pupil_tmp_y = -1j * self.wavevector[0] * position_phase[None, None] * self.pupilmatrix[:, :, None]
  1090. inter_image_y = torch.transpose(self.czt_parallel(pupil_tmp_y, self.ay, self.by, self.dy), -1, -2)
  1091. field_matrix_y = torch.transpose(self.czt_parallel(inter_image_y, self.ax, self.bx, self.dx), -1, -2)
  1092. psfs_ders[slice_tmp, :, :, 1] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_y),
  1093. dim=(0, 1))
  1094. pupil_tmp_z = torch.empty([2, 3, length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  1095. pupil_tmp_z[:, :, idx_0] = 1j * self.wavevectorzmed * position_phase[idx_0] * self.pupilmatrix[:, :, None]
  1096. pupil_tmp_z[:, :, idx_1] = 1j * self.wavevectorzimm * position_phase[idx_1] * self.pupilmatrix[:, :, None]
  1097. inter_image_z = torch.transpose(self.czt_parallel(pupil_tmp_z, self.ay, self.by, self.dy), -1, -2)
  1098. field_matrix_z = torch.transpose(self.czt_parallel(inter_image_z, self.ax, self.bx, self.dx), -1, -2)
  1099. psfs_ders[slice_tmp, :, :, 2] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_z),
  1100. dim=(0, 1))
  1101. psfs /= self.norm_intensity
  1102. psfs_ders /= self.norm_intensity
  1103. # otf rescale
  1104. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1] != 0:
  1105. psfs = self.otf_rescale(psfdata=psfs, sigma_xy=self.otf_rescale_xy)
  1106. psfs_ders[:, :, :, 0] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 0], sigma_xy=self.otf_rescale_xy)
  1107. psfs_ders[:, :, :, 1] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 1], sigma_xy=self.otf_rescale_xy)
  1108. psfs_ders[:, :, :, 2] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 2], sigma_xy=self.otf_rescale_xy)
  1109. psfs_out = ailoc.common.gpu(psfs * photons[:, None, None] + bgs[:, None, None])
  1110. ders_out = torch.zeros([n_mol, self.psf_size, self.psf_size, 5], device='cuda', dtype=self.data_type)
  1111. ders_out[:, :, :, 0:3] = psfs_ders * photons[:, None, None, None]
  1112. ders_out[:, :, :, 3] = psfs
  1113. ders_out[:, :, :, 4] = torch.ones_like(psfs)
  1114. return ders_out, psfs_out
  1115. def _compute_derivative(self, x, y, z, photons, bgs):
  1116. """
  1117. Calculate the analytical derivatives of the PSFs at given parameters with respect to x,y,z,photons,bg
  1118. """
  1119. n_mol = x.shape[0]
  1120. objstage = torch.zeros(x.shape[0], dtype=self.data_type, device='cuda')
  1121. zernike_phase = torch.exp(1j * 2 * np.pi *
  1122. torch.sum(self.zernike_coef[:, None, None]*self.allzernikes, dim=0)
  1123. / self.wavelength)
  1124. pupilmatrix = (self.amplitude[None, None] *
  1125. zernike_phase[None, None] *
  1126. self.polarizationvector)[:, :, None]
  1127. # batch_size in parallel to save GPU memory
  1128. slice_list = []
  1129. batch_size = 100
  1130. for i in np.arange(0, n_mol, batch_size):
  1131. slice_list.append(slice(i, min(i + batch_size, n_mol)))
  1132. psfs = torch.zeros([n_mol, self.psf_size, self.psf_size], device='cuda', dtype=self.data_type)
  1133. psfs_ders = torch.zeros([n_mol, self.psf_size, self.psf_size, 3], device='cuda', dtype=self.data_type)
  1134. for slice_tmp in slice_list:
  1135. # length_tmp = slice_tmp.stop - slice_tmp.start
  1136. # position_phase = torch.empty([length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  1137. # idx_0 = torch.where(z[slice_tmp] + self.zemit0 >= 0)[0]
  1138. # phase_xyz_tmp = -y[slice_tmp][idx_0][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx_0][:, None, None] * \
  1139. # self.wavevector[1][None] + \
  1140. # (z[slice_tmp][idx_0] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  1141. # position_phase[idx_0, :, :] = torch.exp(
  1142. # 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx_0][:, None, None] + self.objstage0) *
  1143. # self.wavevectorzimm[None]))
  1144. #
  1145. # idx_1 = torch.where(z[slice_tmp] + self.zemit0 < 0)[0]
  1146. # phase_xyz_tmp = -y[slice_tmp][idx_1][:, None, None] * self.wavevector[0][None] - x[slice_tmp][idx_1][:, None, None] * \
  1147. # self.wavevector[1][None]
  1148. # position_phase[idx_1, :, :] = torch.exp(
  1149. # 1j * (phase_xyz_tmp + (objstage[slice_tmp][idx_1][:, None, None] + self.objstage0 + z[slice_tmp][idx_1][:, None, None]
  1150. # + self.zemit0) * self.wavevectorzimm[None]))
  1151. phase_xyz_tmp = -y[slice_tmp][:, None, None] * self.wavevector[0][None] - x[slice_tmp][:, None, None] * \
  1152. self.wavevector[1][None] + (z[slice_tmp] + self.zemit0)[:, None, None] * self.wavevectorzmed[None]
  1153. position_phase = torch.exp(1j * (phase_xyz_tmp + (objstage[slice_tmp][:, None, None] + self.objstage0) *
  1154. self.wavevectorzimm[None]))
  1155. pupil_tmp = position_phase[None, None] * pupilmatrix
  1156. inter_image = torch.transpose(self.czt_parallel(pupil_tmp, self.ay, self.by, self.dy), -1, -2)
  1157. field_matrix = torch.transpose(self.czt_parallel(inter_image, self.ax, self.bx, self.dx), -1, -2)
  1158. psfs[slice_tmp] += 1 / 3 * torch.sum((torch.abs(field_matrix[:, :])) ** 2, dim=(0, 1))
  1159. # derivatives with respect to x,y,z
  1160. pupil_tmp_x = -1j * self.wavevector[1] * position_phase[None, None] * pupilmatrix
  1161. inter_image_x = torch.transpose(self.czt_parallel(pupil_tmp_x, self.ay, self.by, self.dy), -1, -2)
  1162. field_matrix_x = torch.transpose(self.czt_parallel(inter_image_x, self.ax, self.bx, self.dx), -1, -2)
  1163. psfs_ders[slice_tmp, :, :, 0] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_x),
  1164. dim=(0, 1))
  1165. pupil_tmp_y = -1j * self.wavevector[0] * position_phase[None, None] * pupilmatrix
  1166. inter_image_y = torch.transpose(self.czt_parallel(pupil_tmp_y, self.ay, self.by, self.dy), -1, -2)
  1167. field_matrix_y = torch.transpose(self.czt_parallel(inter_image_y, self.ax, self.bx, self.dx), -1, -2)
  1168. psfs_ders[slice_tmp, :, :, 1] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_y),
  1169. dim=(0, 1))
  1170. # pupil_tmp_z = torch.empty([2, 3, length_tmp, self.npupil, self.npupil], dtype=self.complex_type, device='cuda')
  1171. # pupil_tmp_z[:, :, idx_0] = 1j * self.wavevectorzmed * position_phase[idx_0] * pupilmatrix
  1172. # pupil_tmp_z[:, :, idx_1] = 1j * self.wavevectorzimm * position_phase[idx_1] * pupilmatrix
  1173. pupil_tmp_z = 1j * self.wavevectorzmed * position_phase * pupilmatrix
  1174. inter_image_z = torch.transpose(self.czt_parallel(pupil_tmp_z, self.ay, self.by, self.dy), -1, -2)
  1175. field_matrix_z = torch.transpose(self.czt_parallel(inter_image_z, self.ax, self.bx, self.dx), -1, -2)
  1176. psfs_ders[slice_tmp, :, :, 2] += 2 / 3 * torch.sum(torch.real(torch.conj(field_matrix) * field_matrix_z),
  1177. dim=(0, 1))
  1178. if self.focus_norm:
  1179. # normalize by focus intensity
  1180. psfs /= self.norm_intensity
  1181. psfs_ders /= self.norm_intensity
  1182. else:
  1183. # normalize by themselves
  1184. norm_factor = psfs.sum(dim=(-1, -2))
  1185. psfs /= norm_factor[:, None, None]
  1186. psfs_ders /= norm_factor[:, None, None, None]
  1187. # otf rescale
  1188. if self.otf_rescale_xy[0] or self.otf_rescale_xy[1]:
  1189. psfs = self.otf_rescale(psfdata=psfs, sigma_xy=self.otf_rescale_xy)
  1190. psfs_ders[:, :, :, 0] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 0], sigma_xy=self.otf_rescale_xy)
  1191. psfs_ders[:, :, :, 1] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 1], sigma_xy=self.otf_rescale_xy)
  1192. psfs_ders[:, :, :, 2] = self.otf_rescale(psfdata=psfs_ders[:, :, :, 2], sigma_xy=self.otf_rescale_xy)
  1193. psfs_out = ailoc.common.gpu(psfs * photons[:, None, None] + bgs[:, None, None])
  1194. ders_out = torch.zeros([n_mol, self.psf_size, self.psf_size, 5], device='cuda', dtype=self.data_type)
  1195. ders_out[:, :, :, 0:3] = psfs_ders * photons[:, None, None, None]
  1196. ders_out[:, :, :, 3] = psfs
  1197. ders_out[:, :, :, 4] = torch.ones_like(psfs)
  1198. return ders_out, psfs_out
  1199. def _torch_jacobian_derivative(self, x, y, z, photons, bgs):
  1200. """Calculate the derivatives of the PSFs at given parameters with respect to x,y,z,photons,bg
  1201. using torch.autograd.functional.jacobian"""
  1202. x.requires_grad = True
  1203. y.requires_grad = True
  1204. z.requires_grad = True
  1205. photons.requires_grad = True
  1206. bgs.requires_grad = True
  1207. jacobian = torch.autograd.functional.jacobian(self.simulate_parallel, (x, y, z, photons))
  1208. ders_out = torch.zeros([x.shape[0], self.psf_size, self.psf_size, 5], device='cuda', dtype=self.data_type)
  1209. for i in range(len(jacobian)):
  1210. for j in range(x.shape[0]):
  1211. ders_out[j, :, :, i] = jacobian[i][j, :, :, j]
  1212. ders_out[:, :, :, 4] = torch.ones([x.shape[0], self.psf_size, self.psf_size])
  1213. return ders_out
  1214. def optimize_crlb(self, x, y, z, photons, bgs, tolerance):
  1215. """
  1216. Optimize the CRLB of this PSF model with respect to zernike coefficients
  1217. at give positions, photons and backgrounds. The instance should be initialized with req_grad=True
  1218. Args:
  1219. x (torch.Tensor): x positions of the PSFs, unit nm
  1220. y (torch.Tensor): y positions of the PSFs, unit nm
  1221. z (torch.Tensor): z positions of the PSFs, unit nm
  1222. photons (torch.Tensor): photon counts of the PSFs, unit photon
  1223. bgs (torch.Tensor): background counts of the PSFs, unit photon
  1224. tolerance (float): stop criteria, the difference of CRLB between two iterations
  1225. Returns:
  1226. """
  1227. # # ADAM optimizer
  1228. # n_mol = x.shape[0]
  1229. # crlb_optimizer = torch.optim.Adam([self.zernike_coef], lr=0.1)
  1230. # loss_1 = 1e6
  1231. # iter_num = 1
  1232. # # for iter_num in range(iterations):
  1233. # while True:
  1234. # xyz_crlb, model = self.compute_crlb(x, y, z, photons, bgs)
  1235. # crlb_3d_avg = torch.sum(xyz_crlb[:, 0]**2 + xyz_crlb[:, 1]**2 + xyz_crlb[:, 2]**2)/n_mol
  1236. # print(f"iter: {iter_num}, crlb_3d_avg: {crlb_3d_avg}")
  1237. # crlb_optimizer.zero_grad()
  1238. # crlb_3d_avg.backward()
  1239. # crlb_optimizer.step()
  1240. # self._pre_compute()
  1241. # if torch.abs((loss_1 - crlb_3d_avg) / loss_1) < tolerance:
  1242. # break
  1243. # loss_1 = crlb_3d_avg.detach()
  1244. # iter_num += 1
  1245. # print('CRLB optimization done')
  1246. # LBFGS optimizer
  1247. n_mol = x.shape[0]
  1248. crlb_optimizer = torch.optim.LBFGS([self.zernike_coef], lr=0.1, max_iter=20, tolerance_grad=1e-7,
  1249. tolerance_change=1e-9)
  1250. loss_1 = 1e6
  1251. iter_num = 1
  1252. def closure():
  1253. crlb_optimizer.zero_grad()
  1254. xyz_crlb, model = self.compute_crlb(x, y, z, photons, bgs)
  1255. crlb_3d_avg = torch.sum(xyz_crlb[:, 0] ** 2 + xyz_crlb[:, 1] ** 2 + xyz_crlb[:, 2] ** 2) / n_mol
  1256. crlb_3d_avg.backward()
  1257. return crlb_3d_avg
  1258. while True:
  1259. crlb_3d_avg = crlb_optimizer.step(closure)
  1260. print(f"iter: {iter_num}, crlb_3d_avg: {crlb_3d_avg}")
  1261. self._pre_compute()
  1262. if torch.abs((loss_1 - crlb_3d_avg) / loss_1) < tolerance:
  1263. break
  1264. loss_1 = crlb_3d_avg.detach()
  1265. iter_num += 1
  1266. print('CRLB optimization done')
  1267. class VectorPSFTorch_2channel(VectorPSFTorch):
  1268. """
  1269. Vectorial PSF model with two detection channels. The two channels have different focal planes.
  1270. Args:
  1271. wavelength (float): emission wavelength, unit nm
  1272. na (float): numerical aperture of the objective
  1273. nmed (float): refractive index of the medium
  1274. nimm (float): refractive index of the immersion medium
  1275. pixel_size (float): camera pixel size, unit nm
  1276. psf_size (int): size of the PSF image, should be an odd number
  1277. z_range (tuple of float): z range of the PSF, unit nm
  1278. z_step (float): z step of the PSF, unit nm
  1279. zemit0 (float, optional): initial axial position of the emitter, default 0, unit nm
  1280. objstage0 (float, optional): initial axial position of the objective stage, default 0, unit nm
  1281. n_zernike (int, optional): number of zernike modes used to model the aberrations,
  1282. default 15 (up to 5th order excluding piston, tip and tilt)
  1283. zernike_coef (torch.Tensor, optional): initial zernike coefficients in a torch tensor,
  1284. default None which means all zeros
  1285. focus_norm (bool, optional): whether to normalize the PSF intensity by the in-focus intensity,
  1286. default True. If False, the PSF is normalized by itself.
  1287. Note that focus_norm=True is more relevant to localization microscopy,
  1288. while focus_norm=False is more relevant to particle tracking microscopy.
  1289. otf_rescale_xy (tuple of float, optional): sigma_x and sigma_y for Gaussian rescaling of the OTF,
  1290. default (0,0) means no rescaling
  1291. req_grad (bool, optional): whether the zernike coefficients require gradients,
  1292. default False. Set True when optimizing the coefficients.
  1293. data_type (torch.dtype, optional): data type for real numbers, default torch.float32
  1294. device (str, optional): device to use, default 'cuda'
  1295. """
  1296. def __init__(self, psf_params, req_grad=False, data_type=torch.float64, zernike_idx_learn=None):
  1297. super(VectorPSFTorch_2channel, self).__init__(psf_params, req_grad, data_type, zernike_idx_learn)
  1298. self.z_offset = psf_params.get('z_offset', 300) # nm
  1299. self.reflection_ratio = psf_params.get('reflection_ratio', 0.5) # ratio of the main channel/sum of two channels
  1300. def simulate(self, x, y, z, photons, objstage=None, zernike_coefs=None):
  1301. main_psfs_out = super().simulate(x, y, z, photons*self.reflection_ratio, objstage, zernike_coefs)
  1302. auxi_psfs_out = super().simulate(x, y, z+self.z_offset, photons*(1-self.reflection_ratio), objstage, zernike_coefs)
  1303. psfs_out = torch.cat((main_psfs_out[None], auxi_psfs_out[None]), dim=0)
  1304. # todo: need to think how to calculate the derivatives for two channels and the CRLB optimization, further combined with neural network
  1305. return psfs_out

vectorpsf.py at commit a603000, under GPL-3.0 · at the source

Overview

Authors: Shuang Fu1, Wei Shi1, Eugene A Katrukha2, Xi Chen3, Yue Fei1, Ke Fang1,4, Ruixiong Wang1, Tianlun Zhang1, Donghan Ma3, Yiming Li1
  1. Department of Biomedical Engineering, Southern University of Science and Technology, Shenzhen, China
  2. Cell Biology, Neurobiology and Biophysics, Department of Biology, Faculty of Science, Utrecht University, Utrecht, the Netherlands
  3. School of Optoelectronic Engineering and Instrumentation Science, Dalian University of Technology, Dalian, China
  4. Center for Health Research, Guangzhou Institutes of Biomedicine and Health, Chinese Academy of Sciences, Guangzhou, China
Journal: Nature communications, volume 17, issue 1, article 6504
Dates: received 7 October 2025; accepted 30 April 2026; published online 16 May 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41467-026-73045-9 · PMID 42143028 · PMCID PMC13376372 · OpenAlex W7161408484
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: histology / microscopy (modality), cellular / molecular (subfield)
Methods: Spectral & time-frequency, Physiology & signal measures, Machine learning
Keywords: Super-resolution microscopy, Fluorescence imaging
MeSH: Imaging, Three-Dimensional*, Neural Networks, Computer*, Single Molecule Imaging*, Animals, Deep Learning, Microscopy, Fluorescence, Mitochondria, Neurons (* major topic)
Topic: Advanced Fluorescence Microscopy Techniques (Biophysics, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: National Natural Science Foundation of China (National Science Foundation of China) (62375116); China Postdoctoral Science Foundation (GZC20240651, 2025M772887, GZC20250546, 2025T180788)
Citations: cited by 2 papers (Europe PMC); 56 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repositories

Its files are read in the Code ↔ Paper reader above, with 16 matches between paragraphs and lines of code.

Li-Lab-SUSTech/LUNAR

License: GPL-3.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: a6030006892a2be8a862aa3e8d7bf03f7d3a0f12, 29 June 2026
Languages: Python (46), Jupyter (6)
Size: 65 files, 52 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (requirements.txt), documentation, 6 notebooks
Not found: CITATION.cff, tests, continuous integration
Tools: PyTorch (35 files), NumPy (34 files), Matplotlib (13 files), SciPy (12 files), tifffile (10 files), imageio (6 files), napari (2 files), OpenCV (2 files), Pillow (2 files), pandas (1 file), scikit-image (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
54 files

jdeschamps/htSMLM

License: LGPL-2.1
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 30b6bfd923a13076c08011ee001cc3a0e83c5434, 8 January 2024
Languages: Java (88), Shell (3)
Size: 135 files, 91 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, tests
Not found: CITATION.cff, environment file, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
93 files

Zenodo 19586991

License: GPL-3.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
At the source:

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-73045-9.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 3 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 143 scripts, each with its path and the digest of its content;
  • 16 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data availability statement

The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

  • it points to a dataset: Zenodo 14709467
  • it says that the data are available on request

Read it in the paper: doi.org/10.1038/s41467-026-73045-9.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 10 authors, 2 keywords, 8 MeSH terms, 2 funders, 47 references.

Cite

This paper

Fu, S., Shi, W., Katrukha, E. A., Chen, X., Fei, Y., Fang, K., Wang, R., Zhang, T., Ma, D., & Li, Y. (2026). Aberration-aware 3D localization microscopy via self-supervised neural-physics learning. Nature communications, 17(1), 6504. https://doi.org/10.1038/s41467-026-73045-9

BibTeX

@article{fu2026aberration,
author = {Fu, Shuang and Shi, Wei and Katrukha, Eugene A and Chen, Xi and Fei, Yue and Fang, Ke and Wang, Ruixiong and Zhang, Tianlun and Ma, Donghan and Li, Yiming},
title = {{Aberration-aware 3D localization microscopy via self-supervised neural-physics learning}},
journal = {Nature communications},
year = {2026},
month = may,
volume = {17},
number = {1},
pages = {6504},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-73045-9},
url = {https://doi.org/10.1038/s41467-026-73045-9},
pmid = {42143028},
pmcid = {PMC13376372}
}

RIS

TY - JOUR
AU - Fu, Shuang
AU - Shi, Wei
AU - Katrukha, Eugene A
AU - Chen, Xi
AU - Fei, Yue
AU - Fang, Ke
AU - Wang, Ruixiong
AU - Zhang, Tianlun
AU - Ma, Donghan
AU - Li, Yiming
TI - Aberration-aware 3D localization microscopy via self-supervised neural-physics learning
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/05/16
VL - 17
IS - 1
SP - 6504
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-73045-9
UR - https://doi.org/10.1038/s41467-026-73045-9
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-73045-9",
"type": "article-journal",
"title": "Aberration-aware 3D localization microscopy via self-supervised neural-physics learning",
"container-title": "Nature communications",
"author": [
{
"family": "Fu",
"given": "Shuang"
},
{
"family": "Shi",
"given": "Wei"
},
{
"family": "Katrukha",
"given": "Eugene A"
},
{
"family": "Chen",
"given": "Xi"
},
{
"family": "Fei",
"given": "Yue"
},
{
"family": "Fang",
"given": "Ke"
},
{
"family": "Wang",
"given": "Ruixiong"
},
{
"family": "Zhang",
"given": "Tianlun"
},
{
"family": "Ma",
"given": "Donghan"
},
{
"family": "Li",
"given": "Yiming"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "6504",
"DOI": "10.1038/s41467-026-73045-9",
"PMID": "42143028",
"PMCID": "PMC13376372",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-73045-9",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
16
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1371/journal.pcbi.1014571 [code]
SynAPSeg: A novel dataset and image analysis framework for deep learning-based synapse detection and quantification.
Journal: PLoS computational biology
In common: napari, imageio, tifffile, 8 other tools
[2] doi:10.1038/s41467-026-73389-2 [code]
Physics-informed multi-encoder adaptive optics enables rapid aberration correction for intravital microscopy of deep complex tissue.
Journal: Nature communications
In common: OpenCV, SciPy, Matplotlib, 1 other tool, histology / microscopy, 6 references
[3] doi:10.64898/2026.03.12.711239 [code]
Back-illumination Phase Imaging Enables Nanoscale Drift Stabilization in Non-transparent Biological Tissues
Journal: bioRxiv (preprint)
In common: napari, tifffile, OpenCV, 6 other tools, histology / microscopy, 1 reference
[4] doi: [code]
Real-time closed-loop feedback system for mouse mesoscale cortical signal and movement control
Journal: eLife
In common: napari, imageio, tifffile, 7 other tools
[5] doi:10.1038/s41467-026-72709-w [code]
An epifluorescence microscope design for naturalistic behavior and cellular activity in freely moving Caenorhabditis elegans.
Journal: Nature communications
In common: imageio, tifffile, OpenCV, 7 other tools, histology / microscopy, cellular / molecular
[6] doi:10.1016/j.isci.2026.116769 [code]
Graph-based modeling of optical system enables adaptive optics on dynamic samples with self-calibration.
Journal: iScience
In common: imageio, tifffile, scikit-image, 5 other tools, 2 references
[7] doi:10.1038/s41467-026-71614-6 [code]
Interferometric ultra-high resolution 3D imaging through brain sections.
Journal: Nature communications
In common: histology / microscopy, cellular / molecular, 7 references
[8] doi:10.1038/s41598-026-57519-w [code]
Automated segmentation of neurons and spinal cord structures in immunofluorescence images using SpineDL.
Journal: Scientific reports
In common: imageio, tifffile, OpenCV, 7 other tools, histology / microscopy
[9] doi:10.1038/s41467-026-75352-7 [code]
Mechanosensory encoding of surface mechanics optimizes locomotion.
Journal: Nature communications
In common: napari, imageio, tifffile, 6 other tools, cellular / molecular
[10] doi:10.1364/boe.605322 [code]
Generalized plaque digitization framework for multi-dimensional mesoscopic images.
Journal: Biomedical optics express
In common: imageio, tifffile, OpenCV, 7 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.