OSCR

Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy.

Code ↔ Paper

5 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 5 matches
  1. [1] § STAR★Methods › Methods details › Hyperparameters for training ↔ model.py, lines 496–553 · score 0.64 · Gibson Lanni model, tg, wavelength, NA
  2. [2] § STAR★Methods › Methods details › Blur generation model architecture ↔ model.py, lines 403–471 · score 0.60 · Gibson Lanni model, neural implicit, PSF, Blur
  3. [3] § STAR★Methods › Methods details › Network architecture ↔ model.py, lines 403–471 · score 0.60 · Gibson Lanni model, neural implicit, PSF, trained
  4. [4] § STAR★Methods › Quantification and statistical analysis › Hyperparameter sensitivity analysis ↔ figures/paper_figures_suppli.ipynb, lines 76–199 · score 0.56 · VQ loss, EWC loss, PSF loss, MRF
  5. [5] § STAR★Methods › Methods details › Optimization ↔ model.py, lines 328–400 · score 0.51 · neural implicit PSF, RProp, batch, trained

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 · 775 lines · 31 KB · no license · 4 matches

  1. import numpy as np
  2. import scipy
  3. import torch
  4. import torch.nn as nn
  5. import torch.nn.functional as F
  6. import torch.distributions as dist
  7. from torch.utils.checkpoint import checkpoint
  8. import torch.special as S
  9. from scipy.stats import lognorm
  10. from fft_conv_pytorch import fft_conv
  11. import matplotlib.pyplot as plt
  12. import time
  13. class JNetBlock0(nn.Module):
  14. def __init__(self, in_channels, out_channels):
  15. super().__init__()
  16. self.conv = nn.Conv3d(in_channels = in_channels ,
  17. out_channels = out_channels,
  18. kernel_size = 7 ,
  19. padding = 'same' ,
  20. padding_mode = 'replicate' ,)
  21. def forward(self, x):
  22. x = self.conv(x)
  23. return x
  24. class JNetBlock(nn.Module):
  25. def __init__(self, in_channels, hidden_channels, dropout):
  26. super().__init__()
  27. self.bn1 = nn.BatchNorm3d(num_features = in_channels)
  28. self.relu1 = nn.ReLU(inplace=True)
  29. self.conv1 = nn.Conv3d(in_channels = in_channels ,
  30. out_channels = hidden_channels,
  31. kernel_size = 3 ,
  32. padding = 'same' ,
  33. padding_mode = 'replicate' ,)
  34. self.bn2 = nn.BatchNorm3d(num_features = hidden_channels)
  35. self.relu2 = nn.ReLU(inplace=True)
  36. self.dropout1 = nn.Dropout(p = dropout)
  37. self.conv2 = nn.Conv3d(in_channels = hidden_channels,
  38. out_channels = in_channels ,
  39. kernel_size = 3 ,
  40. padding = 'same' ,
  41. padding_mode = 'replicate' ,)
  42. def forward(self, x):
  43. d = self.bn1(x)
  44. d = self.relu1(d)
  45. d = self.conv1(d)
  46. d = self.bn2(d)
  47. d = self.relu2(d)
  48. d = self.dropout1(d)
  49. d = self.conv2(d)
  50. x = x + d
  51. return x
  52. class JNetBlockN(nn.Module):
  53. def __init__(self, in_channels, out_channels):
  54. super().__init__()
  55. self.conv = nn.Conv3d(in_channels = in_channels ,
  56. out_channels = out_channels,
  57. kernel_size = 3 ,
  58. padding = 'same' ,
  59. padding_mode = 'replicate' ,)
  60. self.sigm = nn.Sigmoid()
  61. def forward(self, x):
  62. x = self.conv(x)
  63. x = self.sigm(x)
  64. return x
  65. class JNetPooling(nn.Module):
  66. def __init__(self, in_channels, out_channels):
  67. super().__init__()
  68. self.maxpool = nn.MaxPool3d(kernel_size = 2)
  69. self.conv = nn.Conv3d(in_channels = in_channels ,
  70. out_channels = out_channels ,
  71. kernel_size = 1 ,
  72. padding = 'same' ,
  73. padding_mode = 'replicate' ,)
  74. self.relu = nn.ReLU(inplace=True)
  75. def forward(self, x):
  76. x = self.maxpool(x)
  77. x = self.conv(x)
  78. x = self.relu(x)
  79. return x
  80. class JNetUnpooling(nn.Module):
  81. def __init__(self, in_channels, out_channels):
  82. super().__init__()
  83. self.upsample = nn.Upsample(scale_factor = 2 ,
  84. mode = 'trilinear' ,)
  85. self.conv = nn.Conv3d(in_channels = in_channels ,
  86. out_channels = out_channels ,
  87. kernel_size = 1 ,
  88. padding = 'same' ,
  89. padding_mode = 'replicate' ,)
  90. self.relu = nn.ReLU(inplace=True)
  91. def forward(self, x):
  92. x = self.upsample(x)
  93. x = self.conv(x)
  94. x = self.relu(x)
  95. return x
  96. class JNetUpsample(nn.Module):
  97. def __init__(self, scale_factor):
  98. super().__init__()
  99. self.upsample = nn.Upsample(scale_factor = scale_factor ,
  100. mode = 'trilinear' ,)
  101. def forward(self, x):
  102. return self.upsample(x)
  103. class SuperResolutionBlock(nn.Module):
  104. def __init__(self, scale_factor, in_channels, nblocks, dropout):
  105. super().__init__()
  106. self.upsample = nn.Upsample(scale_factor = scale_factor ,
  107. mode = 'trilinear' ,)
  108. self.post = nn.ModuleList([JNetBlock(in_channels = in_channels ,
  109. hidden_channels = in_channels ,
  110. dropout = dropout ,
  111. ) for _ in range(nblocks)])
  112. def forward(self, x):
  113. x = self.upsample(x)
  114. for f in self.post:
  115. x = f(x)
  116. return x
  117. class CrossAttentionBlock(nn.Module):
  118. """
  119. ### Transformer Layer
  120. """
  121. def __init__(self, channels: int, n_heads: int, d_cond: int):
  122. """
  123. :param d_model: is the input embedding size
  124. :param n_heads: is the number of attention heads
  125. :param d_head: is the size of a attention head
  126. :param d_cond: is the size of the conditional embeddings
  127. """
  128. super().__init__()
  129. self.attn = CrossAttention(d_model = channels,
  130. d_cond = d_cond,
  131. n_heads = n_heads,
  132. d_head = channels // n_heads,)
  133. self.norm = nn.LayerNorm(normalized_shape = channels,)
  134. def forward(self, x: torch.Tensor):
  135. """
  136. :param x: are the input embeddings of shape `[batch_size, height * width, d_model]`
  137. :param cond: is the conditional embeddings of shape `[batch_size, n_cond, d_cond]`
  138. """
  139. b, c, d, h, w = x.shape
  140. x = x.permute(0, 2, 3, 4, 1).view(b, d * h * w, c)
  141. x = self.attn(self.norm(x)) + x
  142. x = x.view(b, d, h, w, c).permute(0, 4, 1, 2, 3)
  143. return x
  144. class CrossAttention(nn.Module):
  145. """
  146. ### Cross Attention Layer
  147. This falls-back to self-attention when conditional embeddings are not specified.
  148. """
  149. def __init__(self, d_model: int, d_cond: int, n_heads: int, d_head: int, is_inplace: bool = True):
  150. """
  151. :param d_model: is the input embedding size
  152. :param n_heads: is the number of attention heads
  153. :param d_head: is the size of a attention head
  154. :param d_cond: is the size of the conditional embeddings
  155. :param is_inplace: specifies whether to perform the attention softmax computation inplace to
  156. save memory
  157. """
  158. super().__init__()
  159. self.is_inplace = is_inplace
  160. self.n_heads = n_heads
  161. self.d_head = d_head
  162. # Attention scaling factor
  163. self.scale = d_head ** -0.5
  164. # Query, key and value mappings
  165. d_attn = d_head * n_heads
  166. self.to_q = nn.Linear(in_features = d_model ,
  167. out_features = d_attn ,
  168. bias = False ,)
  169. self.to_k = nn.Linear(in_features = d_cond ,
  170. out_features = d_attn ,
  171. bias = False ,)
  172. self.to_v = nn.Linear(in_features = d_cond ,
  173. out_features = d_attn ,
  174. bias = False ,)
  175. # Final linear layer
  176. self.to_out = nn.Sequential(
  177. nn.Linear(in_features = d_attn ,
  178. out_features = d_model,),
  179. )
  180. def forward(self, x: torch.Tensor, cond=None):
  181. """
  182. :param x: are the input embeddings of shape `[batch_size, height * width, d_model]`
  183. :param cond: is the conditional embeddings of shape `[batch_size, n_cond, d_cond]`
  184. """
  185. has_cond = cond is not None
  186. if not has_cond:
  187. cond = x
  188. q = self.to_q(x)
  189. k = self.to_k(cond)
  190. v = self.to_v(cond)
  191. return self.normal_attention(q, k, v)
  192. def normal_attention(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor):
  193. """
  194. #### Normal Attention
  195. :param q: `[batch_size, seq, d_attn]`
  196. :param k: `[batch_size, seq, d_attn]`
  197. :param v: `[batch_size, seq, d_attn]`
  198. """
  199. q = q.view(*q.shape[:2], self.n_heads, -1)
  200. k = k.view(*k.shape[:2], self.n_heads, -1)
  201. v = v.view(*v.shape[:2], self.n_heads, -1)
  202. attn = torch.einsum('bihd,bjhd->bhij', q, k) * self.scale
  203. if self.is_inplace:
  204. half = attn.shape[0] // 2
  205. attn[half:] = attn[half:].softmax(dim=-1)
  206. attn[:half] = attn[:half].softmax(dim=-1)
  207. else:
  208. attn = attn.softmax(dim=-1)
  209. out = torch.einsum('bhij,bjhd->bihd', attn, v)
  210. # Reshape to `[batch_size, height * width * depth, n_heads * d_head]`
  211. out = out.reshape(*out.shape[:2], -1)
  212. # Map to `[batch_size, height * width, d_model]` with a linear layer
  213. return self.to_out(out)
  214. class VectorQuantizer(nn.Module):
  215. def __init__(self, threshold, device):
  216. super().__init__()
  217. self.t = threshold
  218. self.device = device
  219. def forward(self, x):
  220. x_quantized = (x >= self.t).to(self.device).float()
  221. x_quantized = x + (x_quantized - x).detach()
  222. quantize_loss = F.mse_loss(x_quantized.detach(), x)
  223. return x_quantized, quantize_loss
  224. class JNetLayer(nn.Module):
  225. def __init__(self, in_channels, hidden_channels_list,
  226. attn_list, nblocks, dropout):
  227. super().__init__()
  228. is_attn = attn_list.pop(0)
  229. hidden_channels = hidden_channels_list.pop(0)
  230. self.hidden_channels = hidden_channels
  231. self.pool = JNetPooling(in_channels = in_channels ,
  232. out_channels = hidden_channels,)
  233. self.conv = nn.Conv3d(in_channels = hidden_channels,
  234. out_channels = hidden_channels,
  235. kernel_size = 1 ,
  236. padding = 'same' ,
  237. padding_mode = 'replicate' ,)
  238. self.prev = nn.ModuleList([JNetBlock(in_channels = hidden_channels,
  239. hidden_channels = hidden_channels,
  240. dropout = dropout ,
  241. ) for _ in range(nblocks)])
  242. self.mid = JNetLayer(in_channels = hidden_channels ,
  243. hidden_channels_list = hidden_channels_list ,
  244. attn_list = attn_list ,
  245. nblocks = nblocks ,
  246. dropout = dropout ,
  247. ) if hidden_channels_list else nn.Identity()
  248. self.attn = CrossAttentionBlock(channels = hidden_channels ,
  249. n_heads = 8 ,
  250. d_cond = hidden_channels ,)\
  251. if is_attn else nn.Identity()
  252. self.post = nn.ModuleList([JNetBlock(in_channels = hidden_channels,
  253. hidden_channels = hidden_channels,
  254. dropout = dropout ,
  255. ) for _ in range(nblocks)])
  256. self.unpool = JNetUnpooling(in_channels = hidden_channels,
  257. out_channels = in_channels ,)
  258. def forward(self, x):
  259. d = self.pool(x)
  260. d = self.conv(d) # checkpoint
  261. for f in self.prev:
  262. d = f(d)
  263. d = self.mid(d)
  264. d = self.attn(d)
  265. for f in self.post:
  266. d = f(d)
  267. d = self.unpool(d) # checkpoint
  268. x = x + d
  269. return x
  270. class Emission(nn.Module):
  271. def __init__(self):
  272. super().__init__()
  273. def sample(self, x, params):
  274. b = x.shape[0]
  275. pz0 = dist.LogNormal(loc = torch.tensor(float(params["mu_z"])).view( b,1,1,1,1).expand(*x.shape),
  276. scale = torch.tensor(float(params["sig_z"])).view(b,1,1,1,1).expand(*x.shape),)
  277. x = x * pz0.sample().to(x.device)
  278. x = torch.clip(x, min=0., max=1.)
  279. return x
  280. class NeuralImplicitPSF(nn.Module):
  281. def __init__(self, config, psf):
  282. super().__init__()
  283. mid = config["mid"]
  284. self.layers = nn.Sequential(
  285. nn.BatchNorm1d(2) ,
  286. nn.Linear(2, mid) ,
  287. nn.Sigmoid() ,
  288. nn.BatchNorm1d(mid),
  289. nn.Linear(mid, 1) ,
  290. nn.Sigmoid()
  291. )
  292. self.config = config
  293. self._gen_coord(psf)
  294. self._gen_label(psf)
  295. def _init_weights(self):
  296. for m in self.parameters():
  297. torch.nn.init.normal_(m)
  298. def forward(self, coords):
  299. return self.layers(coords)
  300. def trainer(self):
  301. self.to(self.config["device"])
  302. self._train()
  303. def _gen_coord(self, psf):
  304. self.psf_shape = psf.shape
  305. z, r = self.psf_shape
  306. zs = torch.linspace(-1, 1, steps=z).abs()
  307. rs = torch.linspace(-1, 1, steps=r)
  308. grid_z, grid_r = torch.meshgrid(zs, rs, indexing='ij')
  309. self.coord = torch.stack((
  310. grid_z.flatten(), grid_r.flatten()), -1).to(self.config["device"])
  311. def _gen_label(self, psf):
  312. self.label = psf.flatten()[:, None]
  313. def _train(self):
  314. label = self.label.to(self.config["device"])
  315. coord = self.coord.to(self.config["device"])
  316. loss_fn = eval(self.config["loss_fn"])
  317. loss = torch.tensor(1.)
  318. while loss.item() >= self.config["nipsf_loss_target"]:
  319. self._init_weights()
  320. count = 0
  321. label = label.to(self.config["device"])
  322. coord = coord.to(self.config["device"])
  323. optim = torch.optim.Rprop(self.parameters(), lr=self.config["lr"])
  324. averaged_model = torch.optim.swa_utils.AveragedModel(
  325. self,
  326. multi_avg_fn=torch.optim.swa_utils.get_ema_multi_avg_fn(0.999))
  327. sched = torch.optim.lr_scheduler.ExponentialLR(
  328. optim,
  329. gamma=0.999,)
  330. while loss.item() >= self.config["nipsf_loss_target"]:
  331. optim.zero_grad()
  332. loss = loss_fn(self(coord), label)
  333. loss.backward()
  334. optim.step()
  335. averaged_model.update_parameters(self)
  336. sched.step()
  337. count += 1
  338. if count == self.config["num_iter_psf_pretrain"]:
  339. print("Neural Implicit PSF is not good enough. " \
  340. "Initializing NeuriPSF and trying again...")
  341. break
  342. print(f"NeuriPSF train done with loss target"\
  343. f" {self.config['nipsf_loss_target']}.")
  344. def array(self):
  345. return self(self.coord).view(self.psf_shape)
  346. class Blur(nn.Module):
  347. def __init__(self, params):
  348. super().__init__()
  349. self.device = params["device"]
  350. self.use_fftconv = params["use_fftconv"]
  351. if params["blur_mode"] == "gaussian":
  352. psf_model = GaussianModel(params)
  353. elif params["blur_mode"] == "gibsonlanni":
  354. psf_model = GibsonLanniModel(params)
  355. else:
  356. raise(NotImplementedError(
  357. f'blur_mode {params["blur_mode"]} is not implemented. Try "gaussian" or "gibsonlanni".'))
  358. self.init_psf_rz = torch.tensor(psf_model.PSF_rz, requires_grad=False).float().to(self.device)
  359. self.neuripsf = NeuralImplicitPSF(params, self.init_psf_rz)
  360. # self.neuripsf.trainer(self.init_psf_rz) <- moved to train_runner.py
  361. self.psf_rz_s0 = self.init_psf_rz.shape[0]
  362. self.size_z = params["size_z"]
  363. xy = torch.meshgrid(torch.arange(params["size_y"]),
  364. torch.arange(params["size_x"]),
  365. indexing='ij')
  366. r = torch.tensor(psf_model.r)
  367. x0 = (params["size_x"] - 1) / 2
  368. y0 = (params["size_y"] - 1) / 2
  369. r_pixel = torch.sqrt((xy[1] - x0) ** 2 + (xy[0] - y0) ** 2) * params["res_lateral"]
  370. rs0, = r.shape
  371. self.rps0, self.rps1 = r_pixel.shape
  372. r_e = r[:, None, None].expand(rs0, self.rps0, self.rps1)
  373. r_pixel_e = r_pixel[None].expand(rs0, self.rps0, self.rps1)
  374. r_index = torch.argmin(torch.abs(r_e- r_pixel_e), dim=0)
  375. r_index_fe = r_index.flatten().expand(self.psf_rz_s0, -1)
  376. self.r_index_fe = r_index_fe.to(self.device)
  377. self.z_pad = int((params["size_z"] - params["res_axial"] // params["res_lateral"] + 1) // 2)
  378. self.x_pad = (params["size_x"]) // 2
  379. self.y_pad = (params["size_y"]) // 2
  380. self.stride = (params["scale"], 1, 1)
  381. def forward(self, x):
  382. psf_rz = self.neuripsf.array()
  383. l2_psf_rz = torch.mean((psf_rz - self.init_psf_rz) ** 2)
  384. psf = torch.gather(psf_rz, 1, self.r_index_fe)
  385. psf = psf.reshape(self.psf_rz_s0, self.rps0, self.rps1)
  386. psf = F.interpolate(
  387. input = psf[None, None, :] ,
  388. size = (self.size_z, self.rps0, self.rps1) ,
  389. mode = "nearest" ,
  390. )[0, 0]
  391. psf = psf / torch.sum(psf)#;print("sum: ", torch.sum(psf));print("max: ", torch.max(psf))#/ torch.sum(psf)
  392. if self.use_fftconv:
  393. _x = fft_conv(signal = x ,
  394. kernel = psf ,
  395. stride = self.stride ,
  396. padding = (self.z_pad, self.x_pad, self.y_pad,),
  397. )
  398. else:
  399. _x = F.conv3d(input = x ,
  400. weight = psf ,
  401. stride = self.stride ,
  402. padding = (self.z_pad, self.x_pad, self.y_pad,),
  403. )
  404. return {"out" : _x ,
  405. "psf_loss": l2_psf_rz}
  406. def show_psf_3d(self):
  407. # with torch.no_grad():
  408. psf_rz = self.neuripsf.array()
  409. psf = torch.gather(psf_rz, 1, self.r_index_fe)
  410. psf = psf / torch.sum(psf)
  411. psf = psf.reshape(self.psf_rz_s0, self.rps0, self.rps1)
  412. return psf
  413. class GaussianModel():
  414. def __init__(self, params):
  415. oversampling = 1 # Defines the upsampling ratio on the image space grid for computations
  416. size_x = params["size_x"]
  417. size_y = params["size_y"]
  418. size_z = params["size_z"] // params["scale"]
  419. bet_xy = params["bet_xy"]
  420. bet_z = params["bet_z" ]
  421. x0 = (size_x - 1) / 2
  422. y0 = (size_y - 1) / 2
  423. z0 = (size_z - 1) / 2
  424. res_lateral = params["res_lateral"]#0.05 # microns # # # # param # # # #
  425. max_radius = round(np.sqrt((size_x - x0) * (size_x - x0) + (size_y - y0) * (size_y - y0)))
  426. self.r = res_lateral * np.arange(0, oversampling * max_radius) / oversampling
  427. xy = np.meshgrid(np.arange(size_z), np.arange(max_radius), indexing="ij")
  428. distance = np.sqrt((xy[1] / bet_xy) ** 2 + ((xy[0] - z0) / bet_z) ** 2)
  429. self.PSF_rz = np.exp(- distance ** 2)
  430. def __call__(self):
  431. return self.PSF_rz
  432. class GibsonLanniModel():
  433. def __init__(self, params):
  434. size_x = params["size_x"]#256 # # # # param # # # #
  435. size_y = params["size_y"]#256 # # # # param # # # #
  436. size_z = params["size_z"] // params["scale"]#128 # # # # param # # # #
  437. # Precision control
  438. num_basis = 100 # Number of rescaled Bessels that approximate the phase function
  439. num_samples = 1000 # Number of pupil samples along radial direction
  440. oversampling = 1 # Defines the upsampling ratio on the image space grid for computations
  441. # Microscope parameters
  442. NA = params["NA"] #1.1 # # # # param # # # #
  443. wavelength = params["wavelength"]#0.910 # microns # # # # param # # # #
  444. M = params["M"] #25 # magnification # # # # param # # # #
  445. ns = params["ns"] #1.33 # specimen refractive index (RI)
  446. ng0 = params["ng0"]#1.5 # coverslip RI design value
  447. ng = params["ng"] #1.5 # coverslip RI experimental value
  448. ni0 = params["ni0"]#1.5 # immersion medium RI design value
  449. ni = params["ni"] #1.5 # immersion medium RI experimental value
  450. ti0 = params["ti0"]#150 # microns, working distance (immersion medium thickness) design value
  451. tg0 = params["tg0"]#170 # microns, coverslip thickness design value
  452. tg = params["tg"] #170 # microns, coverslip thickness experimental value
  453. res_lateral = params["res_lateral"]#0.05 # microns # # # # param # # # #
  454. res_axial = params["res_axial"]#0.5 # microns # # # # param # # # #
  455. pZ = params["pZ"] # 2 microns, particle distance from coverslip
  456. # Scaling factors for the Fourier-Bessel series expansion
  457. min_wavelength = 0.436 # microns
  458. scaling_factor = NA * (3 * np.arange(1, num_basis + 1) - 2) * min_wavelength / wavelength
  459. x0 = (size_x - 1) / 2
  460. y0 = (size_y - 1) / 2
  461. max_radius = round(np.sqrt((size_x - x0) * (size_x - x0) + (size_y - y0) * (size_y - y0)))
  462. r = res_lateral * np.arange(0, oversampling * max_radius) / oversampling
  463. self.r = r
  464. a = min([NA, ns, ni, ni0, ng, ng0]) / NA
  465. rho = np.linspace(0, a, num_samples)
  466. z = res_axial * np.arange(-size_z / 2, size_z /2) + res_axial / 2
  467. OPDs = pZ * np.sqrt(ns * ns - NA * NA * rho * rho) # OPD in the sample
  468. OPDi = (z.reshape(-1,1) + ti0) * np.sqrt(ni * ni - NA * NA * rho * rho) - ti0 * np.sqrt(ni0 * ni0 - NA * NA * rho * rho) # OPD in the immersion medium
  469. OPDg = tg * np.sqrt(ng * ng - NA * NA * rho * rho) - tg0 * np.sqrt(ng0 * ng0 - NA * NA * rho * rho) # OPD in the coverslip
  470. W = 2 * np.pi / wavelength * (OPDs + OPDi + OPDg)
  471. phase = np.cos(W) + 1j * np.sin(W)
  472. J = scipy.special.jv(0, scaling_factor.reshape(-1, 1) * rho)
  473. C, residuals, _, _ = np.linalg.lstsq(J.T, phase.T, rcond=-1)
  474. b = 2 * np.pi * r.reshape(-1, 1) * NA / wavelength
  475. J0 = lambda x: scipy.special.j0(x)
  476. J1 = lambda x: scipy.special.j1(x)
  477. denom = scaling_factor * scaling_factor - b * b
  478. R = scaling_factor * J1(scaling_factor * a) * J0(b * a) * a - b * J0(scaling_factor * a) * J1(b * a) * a
  479. R /= denom
  480. PSF_rz = (np.abs(R.dot(C))**2).T
  481. self.PSF_rz = PSF_rz / np.max(PSF_rz)
  482. def __call__(self):
  483. return self.PSF_rz
  484. class Noise(nn.Module):
  485. def __init__(self, params):
  486. super().__init__()
  487. self.sig_eps = params["sig_eps"]
  488. self.a = params["poisson_weight"]
  489. def forward(self, x):
  490. x = x + torch.randn_like(x) * (x * self.a + self.sig_eps)
  491. return x
  492. class PreProcess(nn.Module):
  493. def __init__(self, min, max, params):
  494. super().__init__()
  495. self.min = min
  496. self.max = max
  497. self.background = torch.tensor(params["background"])
  498. self.gamma = dist.Gamma(torch.tensor([1.0]),
  499. torch.tensor(params["background"]))
  500. def forward(self, x):
  501. x = torch.clip(x, min=self.min, max=self.max)
  502. #x = (x - self.min) / (self.max - self.min)
  503. return x
  504. def sample(self, x):
  505. x = (x - self.min) / (self.max - self.min)
  506. # x = x + self.background
  507. # max_value = torch.quantile(x.flatten(), self.max)
  508. x = torch.clip(x, min=self.min, max=self.max) #max_value.item())
  509. return x
  510. class Hill(nn.Module):
  511. def __init__(self, n, ka, params):
  512. super().__init__()
  513. self.n_init = n
  514. self.ka_init = ka
  515. self.n = n
  516. self.k = ka
  517. self.x1 = torch.ones(1).to(device=params["device"]) * 0.5
  518. self.x2 = torch.ones(1).to(device=params["device"]) * 0.8
  519. def forward(self, x):
  520. n = F.sigmoid(self.n)
  521. ka = torch.clip(self.ka, min=0.)
  522. x = self.hill(x, n, ka)
  523. return x
  524. def hill_with_best_value(self, x):
  525. v = self.solve_hill(
  526. self.find_y(x, self.x1),
  527. self.find_y(x, self.x2),
  528. self.hill_ideal(self.x1),
  529. self.hill_ideal(self.x2),)
  530. x = self.hill(x, v['n'], v["ka"])
  531. return x
  532. def sample(self, x):
  533. x = self.hill(x, self.n_init, self.ka_init)
  534. return x
  535. def hill(self, x, n, ka):
  536. return (ka ** n + 1) * x ** n / (ka ** n + x ** n)
  537. def hill_ideal(self, x):
  538. return x ** 0.5 / (1 + x ** 0.5)
  539. def find_y(self, x, x1):
  540. return torch.quantile(x.flatten(), x1)
  541. def solve_hill(self, x1, x2, y1, y2):
  542. a = torch.log(x2) * (torch.log(1 - y1) - torch.log(y1))
  543. b = torch.log(x1) * (torch.log(1 - y2) - torch.log(y2))
  544. c = torch.log(y2) + torch.log(1 - y1)
  545. d = torch.log(y1) + torch.log(1 - y2)
  546. self.k = torch.exp((a - b) / (c - d))
  547. self.n = (torch.log(1 - y1) - torch.log(y1)) \
  548. / (torch.log(self.k) - torch.log(x1))
  549. return {"n" : self.n,
  550. "ka" : self.k }
  551. def inverse_hill(self, x):
  552. x = x.clip(min = 0., max = 1.)
  553. return ((self.k - x + 1) / (self.k * x)) ** (- 1 / self.n)
  554. class ImagingProcess(nn.Module):
  555. def __init__(self, params):
  556. super().__init__()
  557. self.device = params["device"]
  558. self.mu_z = params["mu_z"]
  559. self.sig_z = params["sig_z"]
  560. self.log_ez0 = nn.Parameter(
  561. (torch.tensor(params["mu_z"] + 0.5 \
  562. * params["sig_z"] ** 2)).to(self.device),
  563. requires_grad=True)
  564. self.emission = Emission()
  565. self.blur = Blur(params = params)
  566. self.noise = Noise(params)
  567. self.preprocess = PreProcess(min=0., max=1., params=params)
  568. self.hill = Hill(n=0.5, ka=1., params=params)
  569. def forward(self, x):
  570. out = self.blur(x)
  571. x = out["out"]
  572. x = self.preprocess(x)
  573. # x = self.hill.sample(x)
  574. out = {"out" : x ,
  575. "psf_loss" : out["psf_loss"] }
  576. return out
  577. class JNet(nn.Module):
  578. def __init__(self, params):
  579. super().__init__()
  580. t1 = time.time()
  581. print('initializing JNet model...')
  582. scale_factor = (params["scale"], 1, 1)
  583. hidden_channels_list = params["hidden_channels_list"].copy()
  584. attn_list = params["attn_list"].copy()
  585. hidden_channels = hidden_channels_list.pop(0)
  586. attn_list.pop(0)
  587. self.prev0 = JNetBlock0(in_channels = 1 ,
  588. out_channels = hidden_channels,)
  589. self.prev = nn.ModuleList(
  590. [JNetBlock(
  591. in_channels = hidden_channels,
  592. hidden_channels = hidden_channels,
  593. dropout = params["dropout"],
  594. ) for _ in range(params["nblocks"])
  595. ])
  596. self.mid = JNetLayer(
  597. in_channels = hidden_channels ,
  598. hidden_channels_list = hidden_channels_list ,
  599. attn_list = attn_list ,
  600. nblocks = params["nblocks"] ,
  601. dropout = params["dropout"] ,
  602. ) if hidden_channels_list else nn.Identity()
  603. self.postx = nn.ModuleList(
  604. [JNetBlock(in_channels = hidden_channels,
  605. hidden_channels = hidden_channels,
  606. dropout = params["dropout"],
  607. ) for _ in range(params["nblocks"])
  608. ])
  609. self.postx.append(JNetBlockN(
  610. in_channels = hidden_channels ,
  611. out_channels = 1 ,
  612. ))
  613. self.postz = nn.ModuleList(
  614. [JNetBlock(in_channels = hidden_channels,
  615. hidden_channels = hidden_channels,
  616. dropout = params["dropout"],
  617. ) for _ in range(params["nblocks"])
  618. ])
  619. self.postz.append(JNetBlockN(
  620. in_channels = hidden_channels ,
  621. out_channels = 1 ,
  622. ))
  623. self.image = ImagingProcess(params = params)
  624. self.upsample = JNetUpsample(scale_factor = scale_factor)
  625. self.activation = params["activation"]
  626. self.superres = params["superres"]
  627. self.reconstruct = params["reconstruct"]
  628. self.apply_vq = params["apply_vq"]
  629. self.vq = VectorQuantizer(threshold=params["threshold"],
  630. device=params["device"])
  631. t2 = time.time()
  632. print(f'JNet init done ({t2-t1:.2f} s)')
  633. self.use_x_quantized = params["use_x_quantized"]
  634. self.tau = 1.
  635. a = params["poisson_weight"]
  636. a_inv_sig = np.log(a / (1 - a))
  637. self.a = nn.Parameter(torch.tensor(a_inv_sig)
  638. .to(device=params["device"]))
  639. s = params["sig_eps"]
  640. s_inv_sig = np.log(s / (1 - s))
  641. self.sigma = nn.Parameter(torch.tensor(s_inv_sig)
  642. .to(device=params["device"]))
  643. def forward(self, x):
  644. if self.superres:
  645. x = self.upsample(x)
  646. x = self.prev0(x)
  647. for f in self.prev:
  648. x = f(x)
  649. _x = self.mid(x)
  650. x = _x
  651. z = _x
  652. for f in self.postx:
  653. x = f(x)
  654. for f in self.postz:
  655. z = f(z)
  656. if self.apply_vq:
  657. if self.use_x_quantized:
  658. x, qloss = self.vq(x)
  659. else:
  660. _, qloss = self.vq(x)
  661. if self.reconstruct:
  662. lu = x * z
  663. out = self.image(lu)
  664. r = out["out"]
  665. psf_loss = out["psf_loss"]
  666. else:
  667. r = x
  668. out = {"enhanced_image" : x ,
  669. "reconstruction" : r ,
  670. "mid" : _x ,
  671. "estim_luminance" : z ,
  672. "poisson_weight" : F.sigmoid(self.a) ,
  673. "gaussian_sigma" : F.sigmoid(self.sigma),
  674. }
  675. vqd = {"quantized_loss" : qloss} if self.apply_vq\
  676. else {"quantized_loss" : None}
  677. out = dict(**out, **vqd)
  678. psl = {"psf_loss": psf_loss} if self.reconstruct\
  679. else {"psf_loss": None}
  680. out = dict(**out, **psl)
  681. return out

model.py at commit 4d6bc0e, no license · at the source

Overview

Authors: Haruhiko Morita1, Shuto Hayashi1,2, Takahiro Tsuji3, Daisuke Kato4, Hiroaki Wake5,6, Teppei Shimamura1
ORCID iDs: Haruhiko Morita
  1. Department of Computational and Systems Biology, Division of Biological Data Science, Medical Research Laboratory, Institute of Integrated Research, Institute of Science Tokyo, Yushima, Bunkyoku, Tokyo 113-8510, Japan
  2. Innovation Center of NanoMedicine, Kawasaki Institute of Industrial Promotion, 3-25-14 Tonomachi, Kawasaki-ku, Kawasaki 210-0821, Japan
  3. Nagoya University Graduate School of Medicine, Department of Anatomy and Molecular Cell Biology, Tsurumaicho, Showa-ku, Nagoya, Aichi 466-8560, Japan
  4. Department of Physiology, Nippon Medical School Graduate School of Medicine, 1-25-16 Nezu, Bunkyo-ku, Tokyo 113-0031, Japan
  5. Department of Anatomy and Molecular Cell Biology, Nagoya University Graduate School of Medicine, 65 Tsurumai-cho, Showa-ku, Nagoya 466-8550, Japan
  6. Division of Multicellular Circuit Dynamics, National Institute for Physiological Sciences, National Institute of Natural Sciences, 38 Nishigonaka, Myodaiji-cho, Okazaki, Japan
Journal: Cell reports methods, volume 6, issue 8, article 101476
Dates: received 30 September 2025; accepted 6 May 2026; published online 5 June 2026; in print August 2026
Type: Research article · Language: English
License: CC BY-NC
Identifiers: DOI 10.1016/j.crmeth.2026.101476 · PMID 42248144 · PMCID PMC13494536 · OpenAlex W4415252322
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: histology / microscopy (modality), optical imaging (calcium, voltage, 2-photon) (modality), human (organism), mouse (organism), other condition (population), cellular / molecular (subfield)
Methods: Connectivity, Spectral & time-frequency, Evoked potentials, Machine learning, fMRI & imaging
Keywords: multiphoton excitation microscopy, unsupervised machine learning, autoencoder, microglia
MeSH: Deep Learning*, Microscopy, Fluorescence, Multiphoton*, Unsupervised Machine Learning*, Algorithms, Animals, Brain Neoplasms, Humans, Image Processing, Computer-Assisted, Imaging, Three-Dimensional, Mice, Microglia (* major topic)
Topic: Cell Image Analysis Techniques (Biophysics, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: Japan Agency for Medical Research and Development (JP26zf0127012, JP26tm0424226); Japan Society for the Promotion of Science (22H04925, 26K03026, 23H04938); Core Research for Evolutional Science and Technology (JP26gm2010002, JP26wm0625519, JP25wm0325068, JP26nk0101112); Japan Science and Technology Agency Moonshot Research and Development Program (JPMJMS2026); Japan Science and Technology Agency; University of Tokyo
Citations: not cited yet (Europe PMC); 36 references in the paper
Research resources: RRID:IMSR_JAX:005582

Abstract

Two-photon microscopy enables the non-invasive imaging of deep living tissue. Quantitative three-dimensional analysis is hampered by axial blur and anisotropic resolution in two-photon microscopy. We introduce the two-photon microscopy image enhancement network (TENET), a fully unsupervised framework that simultaneously performs deblurring, up to 12× resolution enhancement, and semantic segmentation on two-photon microscopy volumes in a single pass. TENET embeds a physics-informed blur-generation module with a trainable neural implicit point spread function (PSF), requiring only approximate PSF initialization rather than rigorous experimental measurement, paired “unblurred” images, or isotropy assumptions. On synthetic images, fluorescent beads, and in vivo microglia datasets, TENET surpasses the total-variation-regularized Richardson-Lucy (RLTV) algorithm, CARE, and Neuroclear in image fidelity and segmentation accuracy. Using time-lapse images of microglia surrounding metastatic brain tumors, TENET enables automated 3D morphometry that reveals dynamic microglia-tumor interactions. By converting blur-limited two-photon microscopy data into high-fidelity volumetric reconstructions with ready-to-use masks, TENET streamlines downstream analysis and expands the reach of deep-tissue live imaging.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Repositories

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

HaLU-9000/TENET

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4d6bc0e988969e3ca01c3f4fece7f8d96d6d9867, 27 May 2026
Languages: Python (15), Jupyter (6), Shell (2)
Size: 41 files, 23 scripts
Software Heritage: not archived
Found in: the resources table
Holds: README, environment (requirements.txt), 6 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (19 files), NumPy (18 files), Matplotlib (11 files), tifffile (10 files), OpenCV (5 files), pandas (4 files), scikit-image (3 files), SciPy (3 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
24 files

Zenodo 19663648

License: CC-BY-4.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the resources table
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: PyTorch (18 files), NumPy (17 files), Matplotlib (11 files), tifffile (9 files), OpenCV (5 files), pandas (4 files), scikit-image (3 files), SciPy (3 files)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
23 files

The paper's code and data availability statement is in the Data section.

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:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 45 scripts, each with its path and the digest of its content;
  • 5 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 and code availability

The code for reproducing the tables and figures is available at https://github.com/HaLU-9000/TENET/. The data used for training and raw data of the figures are available at https://zenodo.org/records/18869333 and https://zenodo.org/records/15545449. These repositories include all scripts, model weights, and representative datasets necessary to reproduce the results presented in this manuscript.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

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 2, 28 September 2026

  • Authors: added Haruhiko Morita (0009-0004-4639-6175); removed Haruhiko Morita

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 6 authors, 4 keywords, 11 MeSH terms, 6 funders, 36 references, 1 RRID.

Cite

This paper

Morita, H., Hayashi, S., Tsuji, T., Kato, D., Wake, H., & Shimamura, T. (2026). Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy. Cell reports methods, 6(8), 101476. https://doi.org/10.1016/j.crmeth.2026.101476

BibTeX

@article{morita2026unsupervised,
author = {Morita, Haruhiko and Hayashi, Shuto and Tsuji, Takahiro and Kato, Daisuke and Wake, Hiroaki and Shimamura, Teppei},
title = {{Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy}},
journal = {Cell reports methods},
year = {2026},
month = jun,
volume = {6},
number = {8},
pages = {101476},
publisher = {Elsevier},
issn = {2667-2375},
doi = {10.1016/j.crmeth.2026.101476},
url = {https://doi.org/10.1016/j.crmeth.2026.101476},
pmid = {42248144},
pmcid = {PMC13494536}
}

RIS

TY - JOUR
AU - Morita, Haruhiko
AU - Hayashi, Shuto
AU - Tsuji, Takahiro
AU - Kato, Daisuke
AU - Wake, Hiroaki
AU - Shimamura, Teppei
TI - Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy
T2 - Cell reports methods
J2 - Cell Rep Methods
PY - 2026
DA - 2026/06/05
VL - 6
IS - 8
SP - 101476
SN - 2667-2375
PB - Elsevier
DO - 10.1016/j.crmeth.2026.101476
UR - https://doi.org/10.1016/j.crmeth.2026.101476
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.crmeth.2026.101476",
"type": "article-journal",
"title": "Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy",
"container-title": "Cell reports methods",
"author": [
{
"family": "Morita",
"given": "Haruhiko"
},
{
"family": "Hayashi",
"given": "Shuto"
},
{
"family": "Tsuji",
"given": "Takahiro"
},
{
"family": "Kato",
"given": "Daisuke"
},
{
"family": "Wake",
"given": "Hiroaki"
},
{
"family": "Shimamura",
"given": "Teppei"
}
],
"container-title-short": "Cell Rep Methods",
"volume": "6",
"issue": "8",
"page": "101476",
"DOI": "10.1016/j.crmeth.2026.101476",
"PMID": "42248144",
"PMCID": "PMC13494536",
"ISSN": "2667-2375",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.crmeth.2026.101476",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
5
]
]
}
}

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.1016/j.isci.2026.117010 [code]
Deep learning-assisted mapping of dendritic spines using sequential 2D two-photon calcium imaging.
Journal: iScience
In common: tifffile, OpenCV, scikit-image, 5 other tools, optical imaging (calcium, voltage, 2-photon), histology / microscopy, cellular / molecular, 1 reference
[2] doi:10.7554/elife.105081 [code]
Movie reconstruction from mouse visual cortex activity.
Journal: eLife
In common: tifffile, OpenCV, PyTorch, 4 other tools, optical imaging (calcium, voltage, 2-photon), mouse, 2 references
[3] 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: tifffile, OpenCV, scikit-image, 5 other tools, histology / microscopy, other condition, mouse
[4] 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, optical imaging (calcium, voltage, 2-photon), histology / microscopy, mouse, 3 references
[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: tifffile, OpenCV, scikit-image, 5 other tools, histology / microscopy, cellular / molecular
[6] doi:10.1038/s41467-026-73045-9 [code]
Aberration-aware 3D localization microscopy via self-supervised neural-physics learning.
Journal: Nature communications
In common: tifffile, OpenCV, scikit-image, 5 other tools, histology / microscopy, cellular / molecular
[7] doi:10.1364/boe.600665 [code]
NeuroSeg-MF: robust neuron segmentation in two-photon Ca&lt;sup&gt;2+&lt;/sup&gt; imaging using multi-feature fusion and detection-guided SAM.
Journal: Biomedical optics express
In common: tifffile, OpenCV, PyTorch, 4 other tools, optical imaging (calcium, voltage, 2-photon), 1 reference
[8] doi:10.1371/journal.pcbi.1013441 [code]
Large vision model framework for automated C. elegans analysis: From static morphometry to dynamic neural activity.
Journal: PLoS computational biology
In common: tifffile, OpenCV, scikit-image, 5 other tools, optical imaging (calcium, voltage, 2-photon)
[9] doi:10.1016/j.isci.2026.116206 [code]
Gut distension evokes rapid neural dynamics in vagal and hindbrain populations of larval zebrafish.
Journal: iScience
In common: tifffile, OpenCV, scikit-image, 5 other tools, optical imaging (calcium, voltage, 2-photon)
[10] doi:10.1002/jbio.70260 [code]
Cell-MICS: Detecting Immune Cells With Label-Free Two-Photon Autofluorescence and Deep Learning.
Journal: Journal of biophotonics
In common: tifffile, PyTorch, pandas, 3 other tools, optical imaging (calcium, voltage, 2-photon), histology / microscopy, cellular / molecular, 1 reference

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.