""" Minimal 2D U-Net for neurofilament segmentation. Used by app.py's optional U-Net backend and by train_unet.py. Requires PyTorch. The app applies this per z-slice; input is a single-channel image, output is a one-channel logit map (apply sigmoid for probability). """ import torch import torch.nn as nn def _block(cin, cout): return nn.Sequential( nn.Conv2d(cin, cout, 3, padding=1), nn.BatchNorm2d(cout), nn.ReLU(inplace=True), nn.Conv2d(cout, cout, 3, padding=1), nn.BatchNorm2d(cout), nn.ReLU(inplace=True), ) class UNet(nn.Module): def __init__(self, in_ch=1, out_ch=1, base=32): super().__init__() self.d1 = _block(in_ch, base) self.d2 = _block(base, base * 2) self.d3 = _block(base * 2, base * 4) self.d4 = _block(base * 4, base * 8) self.pool = nn.MaxPool2d(2) self.bott = _block(base * 8, base * 16) self.up4 = nn.ConvTranspose2d(base * 16, base * 8, 2, stride=2) self.u4 = _block(base * 16, base * 8) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2) self.u3 = _block(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2) self.u2 = _block(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) self.u1 = _block(base * 2, base) self.out = nn.Conv2d(base, out_ch, 1) def _pad_to(self, x, ref): dy = ref.shape[-2] - x.shape[-2] dx = ref.shape[-1] - x.shape[-1] if dy or dx: x = nn.functional.pad(x, [0, dx, 0, dy]) return x def forward(self, x): c1 = self.d1(x) c2 = self.d2(self.pool(c1)) c3 = self.d3(self.pool(c2)) c4 = self.d4(self.pool(c3)) b = self.bott(self.pool(c4)) x = self.u4(torch.cat([self._pad_to(self.up4(b), c4), c4], 1)) x = self.u3(torch.cat([self._pad_to(self.up3(x), c3), c3], 1)) x = self.u2(torch.cat([self._pad_to(self.up2(x), c2), c2], 1)) x = self.u1(torch.cat([self._pad_to(self.up1(x), c1), c1], 1)) return self.out(x)