散模型合成病理圖像評(píng)估實(shí)戰(zhàn))
之前在做病理圖像分類項(xiàng)目時(shí)最頭疼的不是模型結(jié)構(gòu)怎么選而是數(shù)據(jù)根本不夠用。醫(yī)院病理科收集的切片數(shù)據(jù)本身就需要專家逐張標(biāo)注標(biāo)注一塊 Tiles 往往就要耗費(fèi)大量時(shí)間再加上罕見(jiàn)病樣本稀缺、患者隱私保護(hù)嚴(yán)格想湊齊一個(gè)類別均衡的訓(xùn)練集非常困難。后來(lái)我們?cè)诜桨咐镆肓藯l件擴(kuò)散模型Conditional Diffusion Model做合成組織病理學(xué)圖像生成把類別標(biāo)簽當(dāng)成生成條件成功補(bǔ)充了低樣本類別下游模型效果也有明顯提升。這篇文章就把這套評(píng)估流程完整記錄下來(lái)包含原理講解、可運(yùn)行代碼、量化評(píng)估方法和常見(jiàn)踩坑總結(jié)希望對(duì)做病理AI、醫(yī)學(xué)圖像生成的朋友有幫助。1. 為什么需要生成合成組織病理學(xué)圖像1.1 病理AI的數(shù)據(jù)瓶頸組織病理學(xué)圖像是病理醫(yī)生診斷腫瘤類型、分級(jí)、判斷預(yù)后最重要的依據(jù)。在進(jìn)行全玻片掃描成像WSI后一張切片往往能達(dá)到數(shù)萬(wàn)像素甚至十億像素級(jí)別直接訓(xùn)練深度學(xué)習(xí)模型并不現(xiàn)實(shí)常規(guī)做法是先將其切分為 256×256 或 512×512 的 Tiles再針對(duì)這些 Tiles 做分類、分割或特征提取。但數(shù)據(jù)層面存在幾個(gè)難以繞開(kāi)的瓶頸。首先是隱私合規(guī)病理圖像涉及患者診斷信息直接跨機(jī)構(gòu)共享樣本需要經(jīng)過(guò)嚴(yán)格的倫理審批和數(shù)據(jù)脫敏流程。其次是標(biāo)注成本病理結(jié)構(gòu)與自然圖像差異極大判定一個(gè)組織區(qū)域是良性、惡性還是腫瘤浸潤(rùn)前沿往往需要多年經(jīng)驗(yàn)的病理醫(yī)生參與標(biāo)注費(fèi)用高、周期長(zhǎng)。第三是類別不均衡某些罕見(jiàn)病變類型在真實(shí)臨床數(shù)據(jù)里占比很低模型在訓(xùn)練時(shí)很容易被多數(shù)類主導(dǎo)對(duì)少樣本類別幾乎學(xué)不到有效特征。合成圖像生成技術(shù)成為緩解上述問(wèn)題的關(guān)鍵手段。通過(guò)生成模型構(gòu)造符合目標(biāo)類別分布的新樣本可以在不直接復(fù)用患者原始數(shù)據(jù)的前提下擴(kuò)充訓(xùn)練集為分類、分割、檢測(cè)等下游任務(wù)提供額外樣本。這類合成樣本不是簡(jiǎn)單的裁剪、翻轉(zhuǎn)、顏色抖動(dòng)而是從數(shù)據(jù)分布層面產(chǎn)生了全新圖像對(duì)提升模型泛化能力更有價(jià)值。1.2 為什么是條件擴(kuò)散模型過(guò)去幾年生成對(duì)抗網(wǎng)絡(luò)GAN是醫(yī)學(xué)圖像生成的主流方案尤其是 StyleGAN 系列在人臉和自然圖像上取得了很好的效果。但在病理圖像場(chǎng)景中GAN 存在訓(xùn)練不穩(wěn)定、模式坍塌、生成圖像紋理重復(fù)等明顯問(wèn)題。組織病理圖像有極強(qiáng)的形態(tài)學(xué)特征比如腺管結(jié)構(gòu)、核異型性、間質(zhì)纖維化一旦模型坍塌到少數(shù)幾種模式生成結(jié)果對(duì)整個(gè)訓(xùn)練集的補(bǔ)充意義就非常有限。擴(kuò)散模型Diffusion Model提供了一個(gè)更穩(wěn)定的生成范式。它的思路不是像 GAN 那樣直接讓生成器與判別器對(duì)抗而是先給真實(shí)圖像逐步加入高斯噪聲直到圖像幾乎完全變成噪聲再訓(xùn)練一個(gè)神經(jīng)網(wǎng)絡(luò)學(xué)習(xí)逐步去噪從而恢復(fù)原始圖像。這個(gè)過(guò)程可以看作從目標(biāo)數(shù)據(jù)分布驅(qū)動(dòng)的噪聲還原訓(xùn)練過(guò)程相對(duì)穩(wěn)定生成質(zhì)量也更容易通過(guò)增加去噪步數(shù)來(lái)提升。無(wú)條件擴(kuò)散模型雖然能生成逼真圖像但無(wú)法控制生成類別這在醫(yī)學(xué)場(chǎng)景中很難直接用。條件擴(kuò)散模型則在噪聲預(yù)測(cè)網(wǎng)絡(luò)中引入額外條件信息比如類別標(biāo)簽、文本描述、圖像引導(dǎo)等。它解決了病理 AI 中非常核心的訴求我們不僅需要生成一張圖像更需要生成一張指定類別、指定病理特征的圖像。這也就是為什么在合成組織病理學(xué)圖像生成任務(wù)中條件擴(kuò)散模型逐步成為主流研究方向。1.3 本文評(píng)估方案與閱讀路線本文圍繞條件擴(kuò)散模型在組織病理學(xué)圖像生成中的評(píng)估展開(kāi)完整流程包括核心原理、數(shù)據(jù)預(yù)處理、基于 Diffusers 庫(kù)的最小實(shí)現(xiàn)、FID、IS、MS-SSIM 評(píng)估方法以及訓(xùn)練穩(wěn)定性和工程落地建議。如果你是剛開(kāi)始接觸擴(kuò)散模型建議先完整閱讀第 3 章原理部分再對(duì)照代碼運(yùn)行。如果已經(jīng)跑過(guò)相關(guān)實(shí)驗(yàn)可以直接跳到第 5 章看評(píng)估指標(biāo)再對(duì)照第 6 章常見(jiàn)問(wèn)題排查。整篇文章的代碼以 PyTorch 生態(tài)為基礎(chǔ)可以按你自己的數(shù)據(jù)集替換數(shù)據(jù)路徑和類別配置。2. 環(huán)境準(zhǔn)備與實(shí)驗(yàn)設(shè)計(jì)2.1 硬件與依賴環(huán)境訓(xùn)練擴(kuò)散模型對(duì)算力有一定要求。本文示例使用 128×128 分辨率和較淺的 UNet顯存占用約 6GB 到 12GB一張 NVIDIA GTX 3060 或更高顯存的顯卡可以完成訓(xùn)練。如果只有普通 CPU 環(huán)境也可以通過(guò)減小圖像尺寸、降低 batch size 跑通流程但生成質(zhì)量會(huì)受限制。跨設(shè)備訓(xùn)練時(shí)建議使用顯存 16GB 以上的 GPU或使用云 GPU 平臺(tái)。軟件環(huán)境以 Python 3.9 以上版本為基準(zhǔn)依賴庫(kù)包括 PyTorch、Diffusers、Torchvision、Accelerate、Tqdm、Pillow、Scikit-learn、OpenCV、Pytorch-FID 和 Torchmetrics。這里不固定具體版本號(hào)因?yàn)?PyTorch 和 Diffusers 迭代較快建議安裝時(shí)使用當(dāng)前穩(wěn)定版本。以下命令可以創(chuàng)建基礎(chǔ)環(huán)境pip install torch torchvision diffusers accelerate tqdm pillow pip install scikit-learn opencv-python pytorch-fid torchmetrics實(shí)際項(xiàng)目中版本需要根據(jù)你的項(xiàng)目環(huán)境調(diào)整。如果使用 Conda也可以先創(chuàng)建虛擬環(huán)境再安裝依賴避免與系統(tǒng) Python 環(huán)境沖突。2.2 項(xiàng)目結(jié)構(gòu)設(shè)計(jì)開(kāi)始寫代碼前先把項(xiàng)目結(jié)構(gòu)規(guī)劃清楚。本文采用以下結(jié)構(gòu)histo_diffusion_eval/ ├── config.py # 全局配置數(shù)據(jù)路徑、訓(xùn)練輪數(shù)、圖像尺寸 ├── dataset.py # 病理 Tiles 數(shù)據(jù)集加載與增強(qiáng) ├── model.py # 條件 UNet 構(gòu)建 ├── train.py # 訓(xùn)練入口 ├── sample.py # 條件生成采樣 └── evaluate.py # 評(píng)估腳本FID、IS、MS-SSIM這種按功能拆分的結(jié)構(gòu)便于復(fù)現(xiàn)實(shí)驗(yàn)也方便后續(xù)更換數(shù)據(jù)集或調(diào)整模型。配置集中在config.py中可以避免在多個(gè)文件里硬編碼參數(shù)。3. 條件擴(kuò)散模型原理與條件注入方式3.1 擴(kuò)散模型的核心過(guò)程擴(kuò)散模型由前向過(guò)程和逆向過(guò)程組成。前向過(guò)程是一個(gè)固定的加噪過(guò)程每一時(shí)間步都向圖像中添加少量高斯噪聲經(jīng)過(guò)足夠多步之后圖像近似變成標(biāo)準(zhǔn)高斯噪聲。若用 T 表示總時(shí)間步數(shù)通常取 1000則前向過(guò)程可以寫成從原始圖像 x? 出發(fā)逐步得到 x?, x?, ..., x_T。訓(xùn)練階段并不需要逐步迭代采樣擴(kuò)散模型的數(shù)學(xué)性質(zhì)允許直接根據(jù)任意時(shí)間步 t 計(jì)算出帶噪圖像。設(shè) α?_t 是噪聲調(diào)度器的累計(jì)系數(shù)隨機(jī)噪聲為 ε則帶噪圖像 x_t 可以表示為x_t sqrt(α?_t) * x_0 sqrt(1 - α?_t) * ε神經(jīng)網(wǎng)絡(luò)的任務(wù)是預(yù)測(cè)噪聲 ε。只要模型能準(zhǔn)確預(yù)測(cè)出當(dāng)前時(shí)刻添加的噪聲逆向過(guò)程就可以從 x_t 中減去預(yù)測(cè)噪聲逐步得到更接近原始圖像的 x_{t-1}最終從純?cè)肼曋羞€原出清晰圖像。這個(gè)方法被驗(yàn)證為穩(wěn)定的生成方案也是近年擴(kuò)散模型在圖像生成領(lǐng)域快速發(fā)展的基礎(chǔ)。3.2 條件信息的三種注入方式條件擴(kuò)散模型與無(wú)條件擴(kuò)散模型最大的區(qū)別在于去噪網(wǎng)絡(luò)需要額外接收條件信息。不同任務(wù)的數(shù)據(jù)形態(tài)不同條件注入方式也有區(qū)別。第一類是類別條件最典型的做法是把類別標(biāo)簽通過(guò)nn.Embedding映射成類別嵌入向量再在 UNet 的殘差塊中與時(shí)間嵌入向量相加引導(dǎo)每個(gè)特征層在去噪時(shí)保持對(duì)應(yīng)類別的語(yǔ)義。本文后續(xù)代碼使用的UNet2DModel支持通過(guò)class_labels參數(shù)傳入類別標(biāo)簽屬于這一類。第二類是圖像條件常用于圖像修復(fù)、超分辨率、分割引導(dǎo)等任務(wù)比如把低分辨率圖或掩碼圖與噪聲圖像在通道維度拼接或者通過(guò)交叉注意力機(jī)制讓網(wǎng)絡(luò)參考條件圖像的特征。這種方法在病理場(chǎng)景中也可以用于指定生成區(qū)域的形態(tài)結(jié)構(gòu)。第三類是文本條件常見(jiàn)做法是使用 CLIP 文本編碼器提取文本特征再通過(guò)交叉注意力層與 UNet 內(nèi)部圖像特征交互。這類方法在自然圖像生成中很流行但病理文本描述標(biāo)注成本高因此目前醫(yī)學(xué)圖像生成研究中使用類別條件和圖像條件的場(chǎng)景更多。本文的病理圖像生成屬于多類別組織分類場(chǎng)景適合使用類別條件。實(shí)際項(xiàng)目中如果數(shù)據(jù)集中包含病變區(qū)域分割掩碼也可以進(jìn)一步改成圖像條件讓模型生成指定區(qū)域的病理結(jié)構(gòu)。3.3 訓(xùn)練目標(biāo)與采樣要點(diǎn)條件擴(kuò)散模型的訓(xùn)練目標(biāo)非常簡(jiǎn)潔。給定干凈圖像 x?、類別標(biāo)簽 c、隨機(jī)時(shí)間步 t 和噪聲 ε模型輸出其對(duì)噪聲的預(yù)測(cè) ε_(tái)θ(x_t, t, c)損失函數(shù)采用噪聲預(yù)測(cè)與真實(shí)噪聲的均方誤差L E[ || ε - ε_(tái)θ(x_t, t, c) ||2 ]這個(gè)目標(biāo)函數(shù)不依賴對(duì)抗訓(xùn)練因此訓(xùn)練過(guò)程相對(duì)穩(wěn)定。需要注意的是時(shí)間步 t 應(yīng)該隨機(jī)均勻采樣讓模型學(xué)會(huì)在所有噪聲強(qiáng)度下都能正確去噪而不是只擅長(zhǎng)某幾個(gè)時(shí)間步。條件信息在訓(xùn)練時(shí)也不能總是參與否則模型會(huì)過(guò)度依賴條件導(dǎo)致無(wú)條件采樣時(shí)效果明顯退化。采樣階段可以使用 DDPM 調(diào)度器逐步去噪。DDPM 的采樣步數(shù)與訓(xùn)練步數(shù)一致速度較慢如果需要更快生成可以使用 DDIM 采樣器用更少的采樣步數(shù)達(dá)到接近的效果。實(shí)際項(xiàng)目中我通常先用 DDPM 完整采樣確認(rèn)質(zhì)量再調(diào)整為 DDIM 加速實(shí)驗(yàn)迭代。4. 完整實(shí)戰(zhàn)條件擴(kuò)散模型生成病理圖像4.1 數(shù)據(jù)集準(zhǔn)備與預(yù)處理本文代碼假設(shè)數(shù)據(jù)集目錄按類別組織每個(gè)類別一個(gè)子文件夾文件夾內(nèi)是已經(jīng)切好的病理 Tiles。Camelyon16、TCGA 等公開(kāi)組織病理數(shù)據(jù)集都可以作為實(shí)驗(yàn)來(lái)源但使用時(shí)需要核對(duì)數(shù)據(jù)授權(quán)協(xié)議按自己的科研或業(yè)務(wù)場(chǎng)景合規(guī)使用。切塊預(yù)處理通常包含以下幾個(gè)步驟從 WSI 中讀取組織區(qū)域過(guò)濾掉純白色背景和玻璃雜質(zhì)區(qū)域?qū)⒂行ЫM織區(qū)域切分為固定尺寸的 Tiles最后人工或基于已有標(biāo)簽完成類別標(biāo)注。對(duì)于快速?gòu)?fù)現(xiàn)實(shí)驗(yàn)可以先收集每類 200 到 500 張 Tiles數(shù)量不多但足夠驗(yàn)證完整流程。dataset.py中實(shí)現(xiàn)一個(gè)讀取本地病理 Tiles 數(shù)據(jù)集的Dataset類。它掃描根目錄下的子文件夾把類別名稱轉(zhuǎn)換成數(shù)字標(biāo)簽并返回圖像和標(biāo)簽。數(shù)據(jù)增強(qiáng)部分使用了隨機(jī)水平和垂直翻轉(zhuǎn)保持病理圖像的結(jié)構(gòu)語(yǔ)義不變。import os from PIL import Image import torch from torch.utils.data import Dataset from torchvision import transforms class HistologyTileDataset(Dataset): def __init__(self, root_dir, image_size128): self.image_paths [] self.labels [] self.class_names sorted(os.listdir(root_dir)) self.class_to_idx {name: i for i, name in enumerate(self.class_names)} for cls_name in self.class_names: cls_dir os.path.join(root_dir, cls_name) if not os.path.isdir(cls_dir): continue for fname in os.listdir(cls_dir): if fname.lower().endswith((.png, .jpg, .jpeg)): self.image_paths.append(os.path.join(cls_dir, fname)) self.labels.append(self.class_to_idx[cls_name]) self.transform transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ]) def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img Image.open(self.image_paths[idx]).convert(RGB) label self.labels[idx] img self.transform(img) return img, label這里把圖像像素歸一化到 -1 到 1 范圍與擴(kuò)散模型噪聲調(diào)度器的輸出范圍保持一致。如果你的數(shù)據(jù)集圖像不是正方形Resize會(huì)統(tǒng)一拉伸到指定尺寸實(shí)際項(xiàng)目中也可以考慮中心裁剪后再縮放減少形變影響。4.2 構(gòu)建條件UNet模型部分直接使用 Hugging Face Diffusers 庫(kù)提供的UNet2DModel它內(nèi)部已經(jīng)實(shí)現(xiàn)了時(shí)間嵌入、類別嵌入、殘差塊、注意力機(jī)制和上下采樣路徑。這樣可以在保持代碼簡(jiǎn)潔的同時(shí)使用經(jīng)過(guò)大規(guī)模實(shí)驗(yàn)驗(yàn)證的模型結(jié)構(gòu)。model.py中的構(gòu)建函數(shù)接收Config對(duì)象返回一個(gè)支持類別條件的 UNet。num_class_embeds對(duì)應(yīng)類別數(shù)量block_out_channels控制每一層的通道數(shù)sample_size需要與數(shù)據(jù)集中圖像尺寸一致。from diffusers import UNet2DModel def build_unet(config): return UNet2DModel( sample_sizeconfig.image_size, in_channels3, out_channels3, layers_per_block2, block_out_channelsconfig.block_out_channels, num_class_embedsconfig.num_class_embeds, dropout0.1, )如果想更深入理解條件注入機(jī)制可以在UNet2DModel的源碼中看到它把類別標(biāo)簽映射為嵌入向量并在多個(gè)殘差塊中與時(shí)間嵌入相加。這種做法可以有效引導(dǎo)生成過(guò)程讓不同類別的圖像在去噪階段逐漸分離開(kāi)來(lái)。config.py中統(tǒng)一管理所有參數(shù)這里給出一個(gè)可運(yùn)行的默認(rèn)配置import torch class Config: # 數(shù)據(jù) data_dir data/tiles image_size 128 num_classes 2 class_names [benign, malignant] # 訓(xùn)練 batch_size 16 num_epochs 100 lr 2e-4 weight_decay 1e-4 grad_clip 1.0 ema_decay 0.995 device cuda if torch.cuda.is_available() else cpu # 模型 block_out_channels (64, 128, 128, 256) time_emb_dim 256 class_emb_dim 128 num_class_embeds 2 # 擴(kuò)散 timesteps 1000 beta_start 1e-4 beta_end 0.02 # 采樣與評(píng)估 ddim_steps 100 sample_batch_size 16 ckpt_dir checkpoints4.3 訓(xùn)練循環(huán)配置訓(xùn)練過(guò)程包括加噪、噪聲預(yù)測(cè)、損失計(jì)算和參數(shù)更新四個(gè)核心步驟。train.py使用DDPMScheduler管理噪聲調(diào)度調(diào)用add_noise方法一步生成帶噪圖像然后用模型預(yù)測(cè)噪聲計(jì)算 MSE 損失。import os import torch import torch.nn.functional as F from torch.utils.data import DataLoader from torchvision import transforms from diffusers import DDPMScheduler from tqdm.auto import tqdm from config import Config from dataset import HistologyTileDataset from model import build_unet def train(config): device torch.device(config.device) dataset HistologyTileDataset(config.data_dir, config.image_size) loader DataLoader( dataset, batch_sizeconfig.batch_size, shuffleTrue, num_workers4, drop_lastTrue, ) noise_scheduler DDPMScheduler( num_train_timestepsconfig.timesteps, beta_startconfig.beta_start, beta_endconfig.beta_end, ) model build_unet(config).to(device) optimizer torch.optim.AdamW(model.parameters(), lrconfig.lr, weight_decayconfig.weight_decay) lr_scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxlen(loader) * config.num_epochs ) os.makedirs(config.ckpt_dir, exist_okTrue) global_step 0 for epoch in range(config.num_epochs): model.train() pbar tqdm(loader, descfEpoch {epoch 1}/{config.num_epochs}) for images, labels in pbar: images images.to(device) labels labels.to(device) noise torch.randn_like(images) timesteps torch.randint( 0, config.timesteps, (images.shape[0],), devicedevice ).long() noisy_images noise_scheduler.add_noise(images, noise, timesteps) noise_pred model(noisy_images, timesteps, class_labelslabels).sample loss F.mse_loss(noise_pred, noise) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), config.grad_clip) optimizer.step() lr_scheduler.step() optimizer.zero_grad() global_step 1 pbar.set_postfix(lossloss.item()) # 每個(gè) epoch 結(jié)束后保存一次 torch.save(model.state_dict(), os.path.join(config.ckpt_dir, fmodel_epoch{epoch 1}.pt)) if __name__ __main__: train(Config())訓(xùn)練中同時(shí)使用了余弦退火學(xué)習(xí)率調(diào)度和梯度裁剪這兩項(xiàng)對(duì)穩(wěn)定訓(xùn)練很有幫助。noise_scheduler.add_noise的第三個(gè)參數(shù)是隨機(jī)采樣的時(shí)間步不要固定成一個(gè)值。模型輸出的.sample字段才是預(yù)測(cè)噪聲這一點(diǎn)在使用UNet2DModel時(shí)需要注意。如果顯存較小可以調(diào)低image_size到 64 或 96或者減小batch_size。如果 GPU 利用率過(guò)低可以增加num_workers來(lái)加速數(shù)據(jù)加載。4.4 生成采樣與保存模型訓(xùn)練完成后使用DDPMScheduler的timesteps從 T 到 0 逐步去噪。每步將當(dāng)前時(shí)刻的帶噪圖像和類別標(biāo)簽輸入模型得到預(yù)測(cè)噪聲再調(diào)用調(diào)度器的step方法獲取去噪后的prev_sample。循環(huán)結(jié)束后即可得到生成圖像。import os import torch from torchvision.utils import save_image from diffusers import DDPMScheduler from tqdm.auto import tqdm from config import Config from model import build_unet def sample(config, class_label0, ckpt_namemodel_epoch100.pt): device torch.device(config.device) noise_scheduler DDPMScheduler( num_train_timestepsconfig.timesteps, beta_startconfig.beta_start, beta_endconfig.beta_end, ) model build_unet(config).to(device) ckpt_path os.path.join(config.ckpt_dir, ckpt_name) model.load_state_dict(torch.load(ckpt_path, map_locationdevice)) model.eval() labels torch.full( (config.sample_batch_size,), fill_valueclass_label, dtypetorch.long, devicedevice, ) x torch.randn( config.sample_batch_size, 3, config.image_size, config.image_size, devicedevice, ) for t in tqdm(noise_scheduler.timesteps): with torch.no_grad(): noise_pred model(x, t, class_labelslabels).sample x noise_scheduler.step(noise_pred, t, x).prev_sample # 從 [-1, 1] 轉(zhuǎn)回 [0, 1] 并保存 x (x 1) / 2 x torch.clamp(x, 0.0, 1.0) save_image(x, fgenerated_class{class_label}.png, nrow4) return x if __name__ __main__: config Config() sample(config, class_label0)生成的圖像可以保存為一個(gè)網(wǎng)格圖片肉眼檢查生成結(jié)果是否具備目標(biāo)類別的基本組織形態(tài)。后續(xù)定量評(píng)估時(shí)則需要把生成圖像批量導(dǎo)出到文件夾中作為 FID 等指標(biāo)的輸入。5. 量化評(píng)估與對(duì)比分析5.1 FID評(píng)估FIDFréchet Inception Distance是圖像生成任務(wù)中最常用的評(píng)估指標(biāo)之一。它先使用特征提取網(wǎng)絡(luò)提取真實(shí)圖像和生成圖像的高維特征再計(jì)算兩個(gè)特征分布之間的 Wasserstein-2 距離。FID 越低說(shuō)明生成圖像與真實(shí)圖像的分布越接近。對(duì)于病理圖像直接用 ImageNet 預(yù)訓(xùn)練的 InceptionV3 提取特征并不是最優(yōu)選擇因?yàn)?ImageNet 的自然圖像特征與病理圖像的形態(tài)特征差異很大。更合理的做法是使用病理圖像預(yù)訓(xùn)練模型作為特征提取器例如在 WSI 數(shù)據(jù)上訓(xùn)練的病理基礎(chǔ)模型。如果只是為了橫向?qū)Ρ炔煌赡P偷南鄬?duì)好壞使用通用的pytorch_fid實(shí)現(xiàn)也能得到一個(gè)有效的參考指標(biāo)。from pytorch_fid import fid_score real_dir data/real_images_class0 gen_dir data/generated_images_class0 fid_value fid_score.calculate_fid_given_paths( [real_dir, gen_dir], batch_size32, devicecuda, dims2048, ) print(fFID: {fid_value:.4f})計(jì)算 FID 時(shí)真實(shí)圖像和生成圖像最好保持相同的數(shù)量和預(yù)處理方式避免因分辨率不一致導(dǎo)致指標(biāo)偏差。生成圖像數(shù)量太少會(huì)帶來(lái)較大方差實(shí)際評(píng)估時(shí)建議每類生成 1000 張以上。5.2 IS評(píng)估ISInception Score從兩個(gè)維度衡量生成質(zhì)量清晰度和多樣性。它使用 InceptionV3 對(duì)生成圖像進(jìn)行分類如果每張圖像的類別預(yù)測(cè)置信度很高同時(shí)整體預(yù)測(cè)分布足夠分散IS 就高。IS 并不需要真實(shí)圖像作為參考因此計(jì)算簡(jiǎn)單但它對(duì)病理圖像的指導(dǎo)意義有限。病理圖像類別的定義與 ImageNet 類別完全不同高 IS 只能說(shuō)明生成圖像在自然圖像特征空間中可分和清晰不能說(shuō)明其病理學(xué)特征是否真實(shí)。所以建議在病理場(chǎng)景中將 IS 作為輔助指標(biāo)重點(diǎn)仍然看 FID 和下游任務(wù)性能。使用torchmetrics可以快速計(jì)算 ISimport torch from torchmetrics.image.inception import InceptionScore inception InceptionScore(splits10) # gen_tensors 是歸一化到 [0, 1] 的生成圖像 Tensor形狀為 [N, C, H, W] inception.update(gen_tensors) score, std inception.compute() print(fIS: {score:.4f} ± {std:.4f})5.3 MS-SSIM與下游任務(wù)評(píng)估FID 和 IS 主要從感知分布上評(píng)估生成質(zhì)量無(wú)法直接反映生成圖像內(nèi)部結(jié)構(gòu)是否合理。組織病理圖像有腺管、細(xì)胞核、間質(zhì)等結(jié)構(gòu)特征因此結(jié)構(gòu)相似性指標(biāo)也有一定參考價(jià)值。MS-SSIMMulti-Scale Structural Similarity Index Measure通過(guò)多尺度比較亮度、對(duì)比度和結(jié)構(gòu)信息衡量生成圖像與真實(shí)圖像之間的結(jié)構(gòu)相似程度。需要強(qiáng)調(diào)MS-SSIM 衡量的是兩幅圖像逐像素級(jí)別的結(jié)構(gòu)相似性它天然適合圖像修復(fù)、超分辨率類任務(wù)。在無(wú)條件生成任務(wù)中生成圖像和真實(shí)圖像本來(lái)就不應(yīng)該完全一致因此 MS-SSIM 更適合作為生成樣本與真實(shí)樣本之間是否出現(xiàn)大面積結(jié)構(gòu)崩壞的參考而不能作為唯一的生成效果指標(biāo)。實(shí)際使用建議分兩類分別計(jì)算比如良性和惡性 Tiles 各自比較避免混合類別導(dǎo)致指標(biāo)失真。更貼近業(yè)務(wù)價(jià)值的評(píng)估方式是下游任務(wù)評(píng)估。將真實(shí)數(shù)據(jù)加上生成數(shù)據(jù)混合訓(xùn)練一個(gè)病理圖像分類模型在獨(dú)立測(cè)試集上評(píng)估分類準(zhǔn)確率或 AUC。如果加入合成數(shù)據(jù)后分類效果有提升說(shuō)明生成樣本確實(shí)能夠補(bǔ)充有效信息。這也是很多醫(yī)學(xué)圖像生成論文使用的評(píng)估思路。最終評(píng)估報(bào)告建議同時(shí)包含生成質(zhì)量指標(biāo)和下游任務(wù)指標(biāo)結(jié)論更有說(shuō)服力。6. 常見(jiàn)問(wèn)題與排查清單6.1 高頻問(wèn)題匯總條件擴(kuò)散模型訓(xùn)練和評(píng)估過(guò)程中會(huì)遇到一些高頻問(wèn)題下面整理成表格方便對(duì)照排查。問(wèn)題現(xiàn)象常見(jiàn)原因解決思路訓(xùn)練損失不下降學(xué)習(xí)率過(guò)大或過(guò)小、數(shù)據(jù)歸一化不一致調(diào)整學(xué)習(xí)率檢查圖像是否歸一化到 [-1,1]生成圖像模糊訓(xùn)練輪數(shù)不足、模型容量偏小增加訓(xùn)練輪數(shù)適當(dāng)增大 UNet 通道數(shù)類別條件失效生成結(jié)果與標(biāo)簽無(wú)關(guān)標(biāo)簽沒(méi)有傳入模型、類別嵌入維度過(guò)大導(dǎo)致過(guò)擬合檢查class_labels參數(shù)考慮類別條件 dropout顯存溢出batch size 過(guò)大、圖像分辨率過(guò)高減小 batch size降低分辨率使用梯度累積FID 偏高生成數(shù)據(jù)量少、評(píng)估特征提取器不匹配增加采樣量使用病理預(yù)訓(xùn)練特征提取器采樣時(shí)出現(xiàn) NaN學(xué)習(xí)率過(guò)高導(dǎo)致模型發(fā)散降低學(xué)習(xí)率使用梯度裁剪檢查 beta 配置訓(xùn)練速度非常慢UNet 注意力層計(jì)算量大使用小分辨率跑通減少block_out_channels6.2 訓(xùn)練不穩(wěn)定排查訓(xùn)練不穩(wěn)定是擴(kuò)散模型最需要注意的問(wèn)題。如果損失曲線出現(xiàn)劇烈抖動(dòng)先檢查學(xué)習(xí)率。擴(kuò)散模型一般使用 1e-4 到 3e-4 的 AdamW 學(xué)習(xí)率過(guò)大會(huì)導(dǎo)致噪聲預(yù)測(cè)目標(biāo)震蕩。其次是檢查beta_start和beta_end配置是否合理默認(rèn)值適合絕大多數(shù)自然圖像任務(wù)自定義數(shù)據(jù)集也可以適當(dāng)調(diào)整。另一個(gè)常見(jiàn)問(wèn)題是條件信息在訓(xùn)練時(shí)分布過(guò)分集中。病理數(shù)據(jù)往往類別樣本量差異很大如果某個(gè)類別只有幾十張圖模型很難學(xué)會(huì)該類別的條件映射。一個(gè)簡(jiǎn)單做法是在訓(xùn)練時(shí)以一定概率比如 10%將類別標(biāo)簽隨機(jī)替換成其他類別或者使用條件 dropout讓模型即使丟掉條件信息也能保持一定生成能力這也有助于避免類別條件過(guò)擬合。采樣階段如果發(fā)現(xiàn)生成圖像中混有非目標(biāo)類別的結(jié)構(gòu)可以檢查采樣時(shí)傳入的class_labels是否與訓(xùn)練時(shí)的標(biāo)簽編號(hào)一致。num_class_embeds的編號(hào)是從 0 開(kāi)始的類別名稱排序過(guò)后的索引必須保持一致否則會(huì)生成錯(cuò)誤類別。7. 最佳實(shí)踐與工程建議7.1 數(shù)據(jù)工程建議病理圖像生成首先要重視數(shù)據(jù)質(zhì)量。原始 WSI 中大量區(qū)域是背景、玻璃、氣泡或邊緣陰影這些區(qū)域如果進(jìn)入訓(xùn)練集模型會(huì)把無(wú)意義紋理當(dāng)作病理結(jié)構(gòu)生成圖像就會(huì)包含大量無(wú)用區(qū)域。預(yù)處理時(shí)建議先分割組織區(qū)域過(guò)濾低對(duì)比度的空白 Tiles再用顏色歸一化降低不同染色方案帶來(lái)的色彩差異。類別劃分需要基于真實(shí)病理標(biāo)注不能只靠文件名約定。如果使用弱標(biāo)簽數(shù)據(jù)還需要額外處理標(biāo)簽噪聲。擴(kuò)展樣本時(shí)要避免同一張 WSI 的近鄰 Tiles 同時(shí)出現(xiàn)在訓(xùn)練集和測(cè)試集防止數(shù)據(jù)泄漏導(dǎo)致評(píng)估虛高。對(duì)于小數(shù)據(jù)集先不要追求分辨率。可以用 64×64 跑通流程確認(rèn)模型能夠擬合訓(xùn)練集后再逐步提高到 128×128 或 256×256。過(guò)早使用高分辨率不僅訓(xùn)練慢排查問(wèn)題也會(huì)更困難