Unsupervised deep learning enables blur-free resolution enhancement in two-photon microscopy.
The 5 matches
- [1] § STAR★Methods › Methods details › Hyperparameters for training ↔ model.py, lines 496–553 · score 0.64 · Gibson Lanni model, tg, wavelength, NA
- [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] § STAR★Methods › Methods details › Network architecture ↔ model.py, lines 403–471 · score 0.60 · Gibson Lanni model, neural implicit, PSF, trained
- [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] § 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
- import numpy as np
- import scipy
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- import torch.distributions as dist
- from torch.utils.checkpoint import checkpoint
- import torch.special as S
- from scipy.stats import lognorm
- from fft_conv_pytorch import fft_conv
- import matplotlib.pyplot as plt
- import time
- class JNetBlock0(nn.Module):
- def __init__(self, in_channels, out_channels):
- super().__init__()
- self.conv = nn.Conv3d(in_channels = in_channels ,
- out_channels = out_channels,
- kernel_size = 7 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- def forward(self, x):
- x = self.conv(x)
- return x
- class JNetBlock(nn.Module):
- def __init__(self, in_channels, hidden_channels, dropout):
- super().__init__()
- self.bn1 = nn.BatchNorm3d(num_features = in_channels)
- self.relu1 = nn.ReLU(inplace=True)
- self.conv1 = nn.Conv3d(in_channels = in_channels ,
- out_channels = hidden_channels,
- kernel_size = 3 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- self.bn2 = nn.BatchNorm3d(num_features = hidden_channels)
- self.relu2 = nn.ReLU(inplace=True)
- self.dropout1 = nn.Dropout(p = dropout)
- self.conv2 = nn.Conv3d(in_channels = hidden_channels,
- out_channels = in_channels ,
- kernel_size = 3 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- def forward(self, x):
- d = self.bn1(x)
- d = self.relu1(d)
- d = self.conv1(d)
- d = self.bn2(d)
- d = self.relu2(d)
- d = self.dropout1(d)
- d = self.conv2(d)
- x = x + d
- return x
- class JNetBlockN(nn.Module):
- def __init__(self, in_channels, out_channels):
- super().__init__()
- self.conv = nn.Conv3d(in_channels = in_channels ,
- out_channels = out_channels,
- kernel_size = 3 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- self.sigm = nn.Sigmoid()
- def forward(self, x):
- x = self.conv(x)
- x = self.sigm(x)
- return x
- class JNetPooling(nn.Module):
- def __init__(self, in_channels, out_channels):
- super().__init__()
- self.maxpool = nn.MaxPool3d(kernel_size = 2)
- self.conv = nn.Conv3d(in_channels = in_channels ,
- out_channels = out_channels ,
- kernel_size = 1 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- self.relu = nn.ReLU(inplace=True)
- def forward(self, x):
- x = self.maxpool(x)
- x = self.conv(x)
- x = self.relu(x)
- return x
- class JNetUnpooling(nn.Module):
- def __init__(self, in_channels, out_channels):
- super().__init__()
- self.upsample = nn.Upsample(scale_factor = 2 ,
- mode = 'trilinear' ,)
- self.conv = nn.Conv3d(in_channels = in_channels ,
- out_channels = out_channels ,
- kernel_size = 1 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- self.relu = nn.ReLU(inplace=True)
- def forward(self, x):
- x = self.upsample(x)
- x = self.conv(x)
- x = self.relu(x)
- return x
- class JNetUpsample(nn.Module):
- def __init__(self, scale_factor):
- super().__init__()
- self.upsample = nn.Upsample(scale_factor = scale_factor ,
- mode = 'trilinear' ,)
- def forward(self, x):
- return self.upsample(x)
- class SuperResolutionBlock(nn.Module):
- def __init__(self, scale_factor, in_channels, nblocks, dropout):
- super().__init__()
- self.upsample = nn.Upsample(scale_factor = scale_factor ,
- mode = 'trilinear' ,)
- self.post = nn.ModuleList([JNetBlock(in_channels = in_channels ,
- hidden_channels = in_channels ,
- dropout = dropout ,
- ) for _ in range(nblocks)])
- def forward(self, x):
- x = self.upsample(x)
- for f in self.post:
- x = f(x)
- return x
- class CrossAttentionBlock(nn.Module):
- """
- ### Transformer Layer
- """
- def __init__(self, channels: int, n_heads: int, d_cond: int):
- """
- :param d_model: is the input embedding size
- :param n_heads: is the number of attention heads
- :param d_head: is the size of a attention head
- :param d_cond: is the size of the conditional embeddings
- """
- super().__init__()
- self.attn = CrossAttention(d_model = channels,
- d_cond = d_cond,
- n_heads = n_heads,
- d_head = channels // n_heads,)
- self.norm = nn.LayerNorm(normalized_shape = channels,)
- def forward(self, x: torch.Tensor):
- """
- :param x: are the input embeddings of shape `[batch_size, height * width, d_model]`
- :param cond: is the conditional embeddings of shape `[batch_size, n_cond, d_cond]`
- """
- b, c, d, h, w = x.shape
- x = x.permute(0, 2, 3, 4, 1).view(b, d * h * w, c)
- x = self.attn(self.norm(x)) + x
- x = x.view(b, d, h, w, c).permute(0, 4, 1, 2, 3)
- return x
- class CrossAttention(nn.Module):
- """
- ### Cross Attention Layer
- This falls-back to self-attention when conditional embeddings are not specified.
- """
- def __init__(self, d_model: int, d_cond: int, n_heads: int, d_head: int, is_inplace: bool = True):
- """
- :param d_model: is the input embedding size
- :param n_heads: is the number of attention heads
- :param d_head: is the size of a attention head
- :param d_cond: is the size of the conditional embeddings
- :param is_inplace: specifies whether to perform the attention softmax computation inplace to
- save memory
- """
- super().__init__()
- self.is_inplace = is_inplace
- self.n_heads = n_heads
- self.d_head = d_head
- # Attention scaling factor
- self.scale = d_head ** -0.5
- # Query, key and value mappings
- d_attn = d_head * n_heads
- self.to_q = nn.Linear(in_features = d_model ,
- out_features = d_attn ,
- bias = False ,)
- self.to_k = nn.Linear(in_features = d_cond ,
- out_features = d_attn ,
- bias = False ,)
- self.to_v = nn.Linear(in_features = d_cond ,
- out_features = d_attn ,
- bias = False ,)
- # Final linear layer
- self.to_out = nn.Sequential(
- nn.Linear(in_features = d_attn ,
- out_features = d_model,),
- )
- def forward(self, x: torch.Tensor, cond=None):
- """
- :param x: are the input embeddings of shape `[batch_size, height * width, d_model]`
- :param cond: is the conditional embeddings of shape `[batch_size, n_cond, d_cond]`
- """
- has_cond = cond is not None
- if not has_cond:
- cond = x
- q = self.to_q(x)
- k = self.to_k(cond)
- v = self.to_v(cond)
- return self.normal_attention(q, k, v)
- def normal_attention(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor):
- """
- #### Normal Attention
- :param q: `[batch_size, seq, d_attn]`
- :param k: `[batch_size, seq, d_attn]`
- :param v: `[batch_size, seq, d_attn]`
- """
- q = q.view(*q.shape[:2], self.n_heads, -1)
- k = k.view(*k.shape[:2], self.n_heads, -1)
- v = v.view(*v.shape[:2], self.n_heads, -1)
- attn = torch.einsum('bihd,bjhd->bhij', q, k) * self.scale
- if self.is_inplace:
- half = attn.shape[0] // 2
- attn[half:] = attn[half:].softmax(dim=-1)
- attn[:half] = attn[:half].softmax(dim=-1)
- else:
- attn = attn.softmax(dim=-1)
- out = torch.einsum('bhij,bjhd->bihd', attn, v)
- # Reshape to `[batch_size, height * width * depth, n_heads * d_head]`
- out = out.reshape(*out.shape[:2], -1)
- # Map to `[batch_size, height * width, d_model]` with a linear layer
- return self.to_out(out)
- class VectorQuantizer(nn.Module):
- def __init__(self, threshold, device):
- super().__init__()
- self.t = threshold
- self.device = device
- def forward(self, x):
- x_quantized = (x >= self.t).to(self.device).float()
- x_quantized = x + (x_quantized - x).detach()
- quantize_loss = F.mse_loss(x_quantized.detach(), x)
- return x_quantized, quantize_loss
- class JNetLayer(nn.Module):
- def __init__(self, in_channels, hidden_channels_list,
- attn_list, nblocks, dropout):
- super().__init__()
- is_attn = attn_list.pop(0)
- hidden_channels = hidden_channels_list.pop(0)
- self.hidden_channels = hidden_channels
- self.pool = JNetPooling(in_channels = in_channels ,
- out_channels = hidden_channels,)
- self.conv = nn.Conv3d(in_channels = hidden_channels,
- out_channels = hidden_channels,
- kernel_size = 1 ,
- padding = 'same' ,
- padding_mode = 'replicate' ,)
- self.prev = nn.ModuleList([JNetBlock(in_channels = hidden_channels,
- hidden_channels = hidden_channels,
- dropout = dropout ,
- ) for _ in range(nblocks)])
- self.mid = JNetLayer(in_channels = hidden_channels ,
- hidden_channels_list = hidden_channels_list ,
- attn_list = attn_list ,
- nblocks = nblocks ,
- dropout = dropout ,
- ) if hidden_channels_list else nn.Identity()
- self.attn = CrossAttentionBlock(channels = hidden_channels ,
- n_heads = 8 ,
- d_cond = hidden_channels ,)\
- if is_attn else nn.Identity()
- self.post = nn.ModuleList([JNetBlock(in_channels = hidden_channels,
- hidden_channels = hidden_channels,
- dropout = dropout ,
- ) for _ in range(nblocks)])
- self.unpool = JNetUnpooling(in_channels = hidden_channels,
- out_channels = in_channels ,)
- def forward(self, x):
- d = self.pool(x)
- d = self.conv(d) # checkpoint
- for f in self.prev:
- d = f(d)
- d = self.mid(d)
- d = self.attn(d)
- for f in self.post:
- d = f(d)
- d = self.unpool(d) # checkpoint
- x = x + d
- return x
- class Emission(nn.Module):
- def __init__(self):
- super().__init__()
- def sample(self, x, params):
- b = x.shape[0]
- pz0 = dist.LogNormal(loc = torch.tensor(float(params["mu_z"])).view( b,1,1,1,1).expand(*x.shape),
- scale = torch.tensor(float(params["sig_z"])).view(b,1,1,1,1).expand(*x.shape),)
- x = x * pz0.sample().to(x.device)
- x = torch.clip(x, min=0., max=1.)
- return x
- class NeuralImplicitPSF(nn.Module):
- def __init__(self, config, psf):
- super().__init__()
- mid = config["mid"]
- self.layers = nn.Sequential(
- nn.BatchNorm1d(2) ,
- nn.Linear(2, mid) ,
- nn.Sigmoid() ,
- nn.BatchNorm1d(mid),
- nn.Linear(mid, 1) ,
- nn.Sigmoid()
- )
- self.config = config
- self._gen_coord(psf)
- self._gen_label(psf)
- def _init_weights(self):
- for m in self.parameters():
- torch.nn.init.normal_(m)
- def forward(self, coords):
- return self.layers(coords)
- def trainer(self):
- self.to(self.config["device"])
- self._train()
- def _gen_coord(self, psf):
- self.psf_shape = psf.shape
- z, r = self.psf_shape
- zs = torch.linspace(-1, 1, steps=z).abs()
- rs = torch.linspace(-1, 1, steps=r)
- grid_z, grid_r = torch.meshgrid(zs, rs, indexing='ij')
- self.coord = torch.stack((
- grid_z.flatten(), grid_r.flatten()), -1).to(self.config["device"])
- def _gen_label(self, psf):
- self.label = psf.flatten()[:, None]
- def _train(self):
- label = self.label.to(self.config["device"])
- coord = self.coord.to(self.config["device"])
- loss_fn = eval(self.config["loss_fn"])
- loss = torch.tensor(1.)
- while loss.item() >= self.config["nipsf_loss_target"]:
- self._init_weights()
- count = 0
- label = label.to(self.config["device"])
- coord = coord.to(self.config["device"])
- optim = torch.optim.Rprop(self.parameters(), lr=self.config["lr"])
- averaged_model = torch.optim.swa_utils.AveragedModel(
- self,
- multi_avg_fn=torch.optim.swa_utils.get_ema_multi_avg_fn(0.999))
- sched = torch.optim.lr_scheduler.ExponentialLR(
- optim,
- gamma=0.999,)
- while loss.item() >= self.config["nipsf_loss_target"]:
- optim.zero_grad()
- loss = loss_fn(self(coord), label)
- loss.backward()
- optim.step()
- averaged_model.update_parameters(self)
- sched.step()
- count += 1
- if count == self.config["num_iter_psf_pretrain"]:
- print("Neural Implicit PSF is not good enough. " \
- "Initializing NeuriPSF and trying again...")
- break
- print(f"NeuriPSF train done with loss target"\
- f" {self.config['nipsf_loss_target']}.")
- def array(self):
- return self(self.coord).view(self.psf_shape)
- class Blur(nn.Module):
- def __init__(self, params):
- super().__init__()
- self.device = params["device"]
- self.use_fftconv = params["use_fftconv"]
- if params["blur_mode"] == "gaussian":
- psf_model = GaussianModel(params)
- elif params["blur_mode"] == "gibsonlanni":
- psf_model = GibsonLanniModel(params)
- else:
- raise(NotImplementedError(
- f'blur_mode {params["blur_mode"]} is not implemented. Try "gaussian" or "gibsonlanni".'))
- self.init_psf_rz = torch.tensor(psf_model.PSF_rz, requires_grad=False).float().to(self.device)
- self.neuripsf = NeuralImplicitPSF(params, self.init_psf_rz)
- # self.neuripsf.trainer(self.init_psf_rz) <- moved to train_runner.py
- self.psf_rz_s0 = self.init_psf_rz.shape[0]
- self.size_z = params["size_z"]
- xy = torch.meshgrid(torch.arange(params["size_y"]),
- torch.arange(params["size_x"]),
- indexing='ij')
- r = torch.tensor(psf_model.r)
- x0 = (params["size_x"] - 1) / 2
- y0 = (params["size_y"] - 1) / 2
- r_pixel = torch.sqrt((xy[1] - x0) ** 2 + (xy[0] - y0) ** 2) * params["res_lateral"]
- rs0, = r.shape
- self.rps0, self.rps1 = r_pixel.shape
- r_e = r[:, None, None].expand(rs0, self.rps0, self.rps1)
- r_pixel_e = r_pixel[None].expand(rs0, self.rps0, self.rps1)
- r_index = torch.argmin(torch.abs(r_e- r_pixel_e), dim=0)
- r_index_fe = r_index.flatten().expand(self.psf_rz_s0, -1)
- self.r_index_fe = r_index_fe.to(self.device)
- self.z_pad = int((params["size_z"] - params["res_axial"] // params["res_lateral"] + 1) // 2)
- self.x_pad = (params["size_x"]) // 2
- self.y_pad = (params["size_y"]) // 2
- self.stride = (params["scale"], 1, 1)
- def forward(self, x):
- psf_rz = self.neuripsf.array()
- l2_psf_rz = torch.mean((psf_rz - self.init_psf_rz) ** 2)
- psf = torch.gather(psf_rz, 1, self.r_index_fe)
- psf = psf.reshape(self.psf_rz_s0, self.rps0, self.rps1)
- psf = F.interpolate(
- input = psf[None, None, :] ,
- size = (self.size_z, self.rps0, self.rps1) ,
- mode = "nearest" ,
- )[0, 0]
- psf = psf / torch.sum(psf)#;print("sum: ", torch.sum(psf));print("max: ", torch.max(psf))#/ torch.sum(psf)
- if self.use_fftconv:
- _x = fft_conv(signal = x ,
- kernel = psf ,
- stride = self.stride ,
- padding = (self.z_pad, self.x_pad, self.y_pad,),
- )
- else:
- _x = F.conv3d(input = x ,
- weight = psf ,
- stride = self.stride ,
- padding = (self.z_pad, self.x_pad, self.y_pad,),
- )
- return {"out" : _x ,
- "psf_loss": l2_psf_rz}
- def show_psf_3d(self):
- # with torch.no_grad():
- psf_rz = self.neuripsf.array()
- psf = torch.gather(psf_rz, 1, self.r_index_fe)
- psf = psf / torch.sum(psf)
- psf = psf.reshape(self.psf_rz_s0, self.rps0, self.rps1)
- return psf
- class GaussianModel():
- def __init__(self, params):
- oversampling = 1 # Defines the upsampling ratio on the image space grid for computations
- size_x = params["size_x"]
- size_y = params["size_y"]
- size_z = params["size_z"] // params["scale"]
- bet_xy = params["bet_xy"]
- bet_z = params["bet_z" ]
- x0 = (size_x - 1) / 2
- y0 = (size_y - 1) / 2
- z0 = (size_z - 1) / 2
- res_lateral = params["res_lateral"]#0.05 # microns # # # # param # # # #
- max_radius = round(np.sqrt((size_x - x0) * (size_x - x0) + (size_y - y0) * (size_y - y0)))
- self.r = res_lateral * np.arange(0, oversampling * max_radius) / oversampling
- xy = np.meshgrid(np.arange(size_z), np.arange(max_radius), indexing="ij")
- distance = np.sqrt((xy[1] / bet_xy) ** 2 + ((xy[0] - z0) / bet_z) ** 2)
- self.PSF_rz = np.exp(- distance ** 2)
- def __call__(self):
- return self.PSF_rz
- class GibsonLanniModel():
- def __init__(self, params):
- size_x = params["size_x"]#256 # # # # param # # # #
- size_y = params["size_y"]#256 # # # # param # # # #
- size_z = params["size_z"] // params["scale"]#128 # # # # param # # # #
- # Precision control
- num_basis = 100 # Number of rescaled Bessels that approximate the phase function
- num_samples = 1000 # Number of pupil samples along radial direction
- oversampling = 1 # Defines the upsampling ratio on the image space grid for computations
- # Microscope parameters
- NA = params["NA"] #1.1 # # # # param # # # #
- wavelength = params["wavelength"]#0.910 # microns # # # # param # # # #
- M = params["M"] #25 # magnification # # # # param # # # #
- ns = params["ns"] #1.33 # specimen refractive index (RI)
- ng0 = params["ng0"]#1.5 # coverslip RI design value
- ng = params["ng"] #1.5 # coverslip RI experimental value
- ni0 = params["ni0"]#1.5 # immersion medium RI design value
- ni = params["ni"] #1.5 # immersion medium RI experimental value
- ti0 = params["ti0"]#150 # microns, working distance (immersion medium thickness) design value
- tg0 = params["tg0"]#170 # microns, coverslip thickness design value
- tg = params["tg"] #170 # microns, coverslip thickness experimental value
- res_lateral = params["res_lateral"]#0.05 # microns # # # # param # # # #
- res_axial = params["res_axial"]#0.5 # microns # # # # param # # # #
- pZ = params["pZ"] # 2 microns, particle distance from coverslip
- # Scaling factors for the Fourier-Bessel series expansion
- min_wavelength = 0.436 # microns
- scaling_factor = NA * (3 * np.arange(1, num_basis + 1) - 2) * min_wavelength / wavelength
- x0 = (size_x - 1) / 2
- y0 = (size_y - 1) / 2
- max_radius = round(np.sqrt((size_x - x0) * (size_x - x0) + (size_y - y0) * (size_y - y0)))
- r = res_lateral * np.arange(0, oversampling * max_radius) / oversampling
- self.r = r
- a = min([NA, ns, ni, ni0, ng, ng0]) / NA
- rho = np.linspace(0, a, num_samples)
- z = res_axial * np.arange(-size_z / 2, size_z /2) + res_axial / 2
- OPDs = pZ * np.sqrt(ns * ns - NA * NA * rho * rho) # OPD in the sample
- 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
- OPDg = tg * np.sqrt(ng * ng - NA * NA * rho * rho) - tg0 * np.sqrt(ng0 * ng0 - NA * NA * rho * rho) # OPD in the coverslip
- W = 2 * np.pi / wavelength * (OPDs + OPDi + OPDg)
- phase = np.cos(W) + 1j * np.sin(W)
- J = scipy.special.jv(0, scaling_factor.reshape(-1, 1) * rho)
- C, residuals, _, _ = np.linalg.lstsq(J.T, phase.T, rcond=-1)
- b = 2 * np.pi * r.reshape(-1, 1) * NA / wavelength
- J0 = lambda x: scipy.special.j0(x)
- J1 = lambda x: scipy.special.j1(x)
- denom = scaling_factor * scaling_factor - b * b
- R = scaling_factor * J1(scaling_factor * a) * J0(b * a) * a - b * J0(scaling_factor * a) * J1(b * a) * a
- R /= denom
- PSF_rz = (np.abs(R.dot(C))**2).T
- self.PSF_rz = PSF_rz / np.max(PSF_rz)
- def __call__(self):
- return self.PSF_rz
- class Noise(nn.Module):
- def __init__(self, params):
- super().__init__()
- self.sig_eps = params["sig_eps"]
- self.a = params["poisson_weight"]
- def forward(self, x):
- x = x + torch.randn_like(x) * (x * self.a + self.sig_eps)
- return x
- class PreProcess(nn.Module):
- def __init__(self, min, max, params):
- super().__init__()
- self.min = min
- self.max = max
- self.background = torch.tensor(params["background"])
- self.gamma = dist.Gamma(torch.tensor([1.0]),
- torch.tensor(params["background"]))
- def forward(self, x):
- x = torch.clip(x, min=self.min, max=self.max)
- #x = (x - self.min) / (self.max - self.min)
- return x
- def sample(self, x):
- x = (x - self.min) / (self.max - self.min)
- # x = x + self.background
- # max_value = torch.quantile(x.flatten(), self.max)
- x = torch.clip(x, min=self.min, max=self.max) #max_value.item())
- return x
- class Hill(nn.Module):
- def __init__(self, n, ka, params):
- super().__init__()
- self.n_init = n
- self.ka_init = ka
- self.n = n
- self.k = ka
- self.x1 = torch.ones(1).to(device=params["device"]) * 0.5
- self.x2 = torch.ones(1).to(device=params["device"]) * 0.8
- def forward(self, x):
- n = F.sigmoid(self.n)
- ka = torch.clip(self.ka, min=0.)
- x = self.hill(x, n, ka)
- return x
- def hill_with_best_value(self, x):
- v = self.solve_hill(
- self.find_y(x, self.x1),
- self.find_y(x, self.x2),
- self.hill_ideal(self.x1),
- self.hill_ideal(self.x2),)
- x = self.hill(x, v['n'], v["ka"])
- return x
- def sample(self, x):
- x = self.hill(x, self.n_init, self.ka_init)
- return x
- def hill(self, x, n, ka):
- return (ka ** n + 1) * x ** n / (ka ** n + x ** n)
- def hill_ideal(self, x):
- return x ** 0.5 / (1 + x ** 0.5)
- def find_y(self, x, x1):
- return torch.quantile(x.flatten(), x1)
- def solve_hill(self, x1, x2, y1, y2):
- a = torch.log(x2) * (torch.log(1 - y1) - torch.log(y1))
- b = torch.log(x1) * (torch.log(1 - y2) - torch.log(y2))
- c = torch.log(y2) + torch.log(1 - y1)
- d = torch.log(y1) + torch.log(1 - y2)
- self.k = torch.exp((a - b) / (c - d))
- self.n = (torch.log(1 - y1) - torch.log(y1)) \
- / (torch.log(self.k) - torch.log(x1))
- return {"n" : self.n,
- "ka" : self.k }
- def inverse_hill(self, x):
- x = x.clip(min = 0., max = 1.)
- return ((self.k - x + 1) / (self.k * x)) ** (- 1 / self.n)
- class ImagingProcess(nn.Module):
- def __init__(self, params):
- super().__init__()
- self.device = params["device"]
- self.mu_z = params["mu_z"]
- self.sig_z = params["sig_z"]
- self.log_ez0 = nn.Parameter(
- (torch.tensor(params["mu_z"] + 0.5 \
- * params["sig_z"] ** 2)).to(self.device),
- requires_grad=True)
- self.emission = Emission()
- self.blur = Blur(params = params)
- self.noise = Noise(params)
- self.preprocess = PreProcess(min=0., max=1., params=params)
- self.hill = Hill(n=0.5, ka=1., params=params)
- def forward(self, x):
- out = self.blur(x)
- x = out["out"]
- x = self.preprocess(x)
- # x = self.hill.sample(x)
- out = {"out" : x ,
- "psf_loss" : out["psf_loss"] }
- return out
- class JNet(nn.Module):
- def __init__(self, params):
- super().__init__()
- t1 = time.time()
- print('initializing JNet model...')
- scale_factor = (params["scale"], 1, 1)
- hidden_channels_list = params["hidden_channels_list"].copy()
- attn_list = params["attn_list"].copy()
- hidden_channels = hidden_channels_list.pop(0)
- attn_list.pop(0)
- self.prev0 = JNetBlock0(in_channels = 1 ,
- out_channels = hidden_channels,)
- self.prev = nn.ModuleList(
- [JNetBlock(
- in_channels = hidden_channels,
- hidden_channels = hidden_channels,
- dropout = params["dropout"],
- ) for _ in range(params["nblocks"])
- ])
- self.mid = JNetLayer(
- in_channels = hidden_channels ,
- hidden_channels_list = hidden_channels_list ,
- attn_list = attn_list ,
- nblocks = params["nblocks"] ,
- dropout = params["dropout"] ,
- ) if hidden_channels_list else nn.Identity()
- self.postx = nn.ModuleList(
- [JNetBlock(in_channels = hidden_channels,
- hidden_channels = hidden_channels,
- dropout = params["dropout"],
- ) for _ in range(params["nblocks"])
- ])
- self.postx.append(JNetBlockN(
- in_channels = hidden_channels ,
- out_channels = 1 ,
- ))
- self.postz = nn.ModuleList(
- [JNetBlock(in_channels = hidden_channels,
- hidden_channels = hidden_channels,
- dropout = params["dropout"],
- ) for _ in range(params["nblocks"])
- ])
- self.postz.append(JNetBlockN(
- in_channels = hidden_channels ,
- out_channels = 1 ,
- ))
- self.image = ImagingProcess(params = params)
- self.upsample = JNetUpsample(scale_factor = scale_factor)
- self.activation = params["activation"]
- self.superres = params["superres"]
- self.reconstruct = params["reconstruct"]
- self.apply_vq = params["apply_vq"]
- self.vq = VectorQuantizer(threshold=params["threshold"],
- device=params["device"])
- t2 = time.time()
- print(f'JNet init done ({t2-t1:.2f} s)')
- self.use_x_quantized = params["use_x_quantized"]
- self.tau = 1.
- a = params["poisson_weight"]
- a_inv_sig = np.log(a / (1 - a))
- self.a = nn.Parameter(torch.tensor(a_inv_sig)
- .to(device=params["device"]))
- s = params["sig_eps"]
- s_inv_sig = np.log(s / (1 - s))
- self.sigma = nn.Parameter(torch.tensor(s_inv_sig)
- .to(device=params["device"]))
- def forward(self, x):
- if self.superres:
- x = self.upsample(x)
- x = self.prev0(x)
- for f in self.prev:
- x = f(x)
- _x = self.mid(x)
- x = _x
- z = _x
- for f in self.postx:
- x = f(x)
- for f in self.postz:
- z = f(z)
- if self.apply_vq:
- if self.use_x_quantized:
- x, qloss = self.vq(x)
- else:
- _, qloss = self.vq(x)
- if self.reconstruct:
- lu = x * z
- out = self.image(lu)
- r = out["out"]
- psf_loss = out["psf_loss"]
- else:
- r = x
- out = {"enhanced_image" : x ,
- "reconstruction" : r ,
- "mid" : _x ,
- "estim_luminance" : z ,
- "poisson_weight" : F.sigmoid(self.a) ,
- "gaussian_sigma" : F.sigmoid(self.sigma),
- }
- vqd = {"quantized_loss" : qloss} if self.apply_vq\
- else {"quantized_loss" : None}
- out = dict(**out, **vqd)
- psl = {"psf_loss": psf_loss} if self.reconstruct\
- else {"psf_loss": None}
- out = dict(**out, **psl)
- return out
model.py at commit 4d6bc0e, no license · at the source
Overview
- 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
- Innovation Center of NanoMedicine, Kawasaki Institute of Industrial Promotion, 3-25-14 Tonomachi, Kawasaki-ku, Kawasaki 210-0821, Japan
- Nagoya University Graduate School of Medicine, Department of Anatomy and Molecular Cell Biology, Tsurumaicho, Showa-ku, Nagoya, Aichi 466-8560, Japan
- Department of Physiology, Nippon Medical School Graduate School of Medicine, 1-25-16 Nezu, Bunkyo-ku, Tokyo 113-0031, Japan
- Department of Anatomy and Molecular Cell Biology, Nagoya University Graduate School of Medicine, 65 Tsurumai-cho, Showa-ku, Nagoya 466-8550, Japan
- Division of Multicellular Circuit Dynamics, National Institute for Physiological Sciences, National Institute of Natural Sciences, 38 Nishigonaka, Myodaiji-cho, Okazaki, Japan
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-regulari
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
4d6bc0e988969e3ca01c3f4fece7f8d96d6d9867, 27 May 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
24 files
- apply.py, Python, 84 lines
- blur.py, Python, 73 lines
- data_generation/
fig2d_data.ipynb , Jupyter, 47 lines - dataset.py, Python, 1,118 lines
- experiments/
bash/ , Shell, 1 linebeads_experiment.sh - experiments/
bash/ , Shell, 1 linesimulation_experiment.sh - figures/
paper_figures2b_adv.ipyn , Jupyter, 508 linesb - figures/
paper_figures2b_our.ipyn , Jupyter, 928 linesb - figures/
paper_figures3.ipynb , Jupyter, 928 lines - figures/
paper_figures4.ipynb , Jupyter, 712 lines - figures/
paper_figures_suppli.ipy , Jupyter, 1,318 lines, 1 matchnb - finetuning.py, Python, 211 lines
- finetuning_with_simulati
on.py , Python, 177 lines - inference.py, Python, 819 lines
- makedata.py, Python, 260 lines
- model.py, Python, 775 lines, 4 matches
- randomdataset_blurer.py, Python, 13 lines
- randomdataset_maker.py, Python, 52 lines
- report_beads_cv.py, Python, 43 lines
- reporter.py, Python, 270 lines
- train_loop.py, Python, 1,288 lines
- train_runner.py, Python, 126 lines
- utils.py, Python, 952 lines
- readme.md, Text, 253 lines
Zenodo 19663648
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
- 27 September 2026: the link answers (HTTP 200)
23 files
- apply.py, Python, 84 lines
- data_generation/
fig2d_data.ipynb , Jupyter, 47 lines - dataset.py, Python, 1,118 lines
- experiments/
bash/ , Shell, 1 linebeads_experiment.sh - experiments/
bash/ , Shell, 1 linesimulation_experiment.sh - figures/
paper_figures2b_adv.ipyn , Jupyter, 508 linesb - figures/
paper_figures2b_our.ipyn , Jupyter, 928 linesb - figures/
paper_figures3.ipynb , Jupyter, 928 lines - figures/
paper_figures4.ipynb , Jupyter, 712 lines - figures/
paper_figures_suppli.ipy , Jupyter, 1,318 linesnb - finetuning.py, Python, 211 lines
- finetuning_with_simulati
on.py , Python, 177 lines - inference.py, Python, 819 lines
- makedata.py, Python, 260 lines
- model.py, Python, 775 lines
- randomdataset_blurer.py, Python, 13 lines
- randomdataset_maker.py, Python, 52 lines
- report_beads_cv.py, Python, 43 lines
- reporter.py, Python, 270 lines
- train_loop.py, Python, 1,288 lines
- train_runner.py, Python, 126 lines
- utils.py, Python, 955 lines
- readme.md, Text, 253 lines
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
- zenodo:15545449, at Zenodo; found in “Data and code availability”
- zenodo:18869333, at Zenodo; found in “Data and code availability”
Data and code availability
The code for reproducing the tables and figures is available at https://
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://
BibTeX
@article{morita2026unsup
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/
url = {https://
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/
VL - 6
IS - 8
SP - 101476
SN - 2667-2375
PB - Elsevier
DO - 10.1016/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1016/
"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":
"volume": "6",
"issue": "8",
"page": "101476",
"DOI": "10.1016/
"PMID": "42248144",
"PMCID": "PMC13494536",
"ISSN": "2667-2375",
"publisher": "Elsevier",
"URL": "https://
"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: iScienceIn 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: eLifeIn 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 reportsIn 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 communicationsIn 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 communicationsIn 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 communicationsIn 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 expressIn 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 biologyIn 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: iScienceIn 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 biophotonicsIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 2 repositories of the authors' code, each at its verified commit and with its license, 45 scripts, and 5 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:241629136277168a…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
