""" Train the 2D U-Net for neurofilament segmentation. You need a small set of labeled slices: pairs of a grayscale image and a binary mask (fiber = 1, background = 0). Put images in --images and masks in --masks with matching filenames. This produces unet_weights.pt, which you then upload in the app's Advanced settings and select the U-Net method. Labels can come from the app's trainable classifier output, from Ilastik, or from hand tracing. Even 10 to 30 well-labeled slices are enough to start. Example: python train_unet.py --images ./train/img --masks ./train/mask --epochs 40 Requires: torch, numpy, scikit-image. """ import os import argparse import numpy as np import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader from skimage import io as skio, transform from unet_model import UNet def _norm01(a): a = a.astype(np.float32) p1, p99 = np.percentile(a, (1, 99.7)) return np.clip((a - p1) / max(p99 - p1, 1e-6), 0, 1) class SliceSet(Dataset): def __init__(self, img_dir, mask_dir, size=512): self.pairs = [] for fn in sorted(os.listdir(img_dir)): mp = os.path.join(mask_dir, fn) if os.path.exists(mp): self.pairs.append((os.path.join(img_dir, fn), mp)) self.size = size if not self.pairs: raise RuntimeError("No matching image/mask filename pairs found.") def __len__(self): return len(self.pairs) def __getitem__(self, i): ip, mp = self.pairs[i] img = _norm01(np.asarray(skio.imread(ip, as_gray=True))) msk = (np.asarray(skio.imread(mp, as_gray=True)) > 0).astype(np.float32) img = transform.resize(img, (self.size, self.size), preserve_range=True) msk = transform.resize(msk, (self.size, self.size), order=0, preserve_range=True) return (torch.from_numpy(img)[None].float(), torch.from_numpy(msk)[None].float()) def dice_bce(logits, target, eps=1.0): bce = nn.functional.binary_cross_entropy_with_logits(logits, target) p = torch.sigmoid(logits) dice = 1 - (2 * (p * target).sum() + eps) / (p.sum() + target.sum() + eps) return bce + dice def main(): ap = argparse.ArgumentParser() ap.add_argument("--images", required=True) ap.add_argument("--masks", required=True) ap.add_argument("--epochs", type=int, default=40) ap.add_argument("--batch", type=int, default=4) ap.add_argument("--lr", type=float, default=1e-3) ap.add_argument("--size", type=int, default=512) ap.add_argument("--out", default="unet_weights.pt") args = ap.parse_args() dev = "cuda" if torch.cuda.is_available() else "cpu" ds = SliceSet(args.images, args.masks, size=args.size) dl = DataLoader(ds, batch_size=args.batch, shuffle=True) net = UNet(1, 1).to(dev) opt = torch.optim.Adam(net.parameters(), lr=args.lr) for ep in range(args.epochs): net.train() tot = 0.0 for x, y in dl: x, y = x.to(dev), y.to(dev) opt.zero_grad() loss = dice_bce(net(x), y) loss.backward() opt.step() tot += loss.item() print(f"epoch {ep + 1}/{args.epochs} loss {tot / len(dl):.4f}") torch.save(net.state_dict(), args.out) print("saved", args.out) if __name__ == "__main__": main()