simoyaman/Ct_anomaly_detector
0
1"""OrganAMNIST(正常画像のみ)でConvolutional Autoencoderを学習するスクリプト。2 3使い方:4 python train_autoencoder.py --epochs 155 python train_autoencoder.py --limit-samples 2000 --epochs 1 # 動作確認用の高速実行6"""7from __future__ import annotations8 9import argparse10import json11import time12from pathlib import Path13 14import numpy as np15import torch16from torch import nn17from torch.utils.data import DataLoader, Dataset18 19from model_def import IMG_SIZE, LATENT_DIM, ConvAutoencoder20 21MODEL_DIR = Path(__file__).parent / "model"22WEIGHTS_PATH = MODEL_DIR / "autoencoder.pth"23META_PATH = MODEL_DIR / "autoencoder_meta.json"24 25DATA_FLAG = "organamnist"26 27 28def _resize_stack(imgs: np.ndarray, size: int) -> np.ndarray:29 """(N,H,W) uint8 配列を size x size にリサイズする。"""30 if imgs.shape[1] == size and imgs.shape[2] == size:31 return imgs32 from PIL import Image33 34 resized = np.zeros((imgs.shape[0], size, size), dtype=np.uint8)35 for i in range(imgs.shape[0]):36 resized[i] = np.array(37 Image.fromarray(imgs[i]).resize((size, size), Image.BILINEAR)38 )39 return resized40 41 42def load_organamnist_split(43 split: str, img_size: int, limit: int | None = None44) -> np.ndarray:45 """MedMNIST OrganAMNISTの指定splitを (N, img_size, img_size) uint8 配列で返す。"""46 import medmnist47 from medmnist import INFO48 49 info = INFO[DATA_FLAG]50 data_class = getattr(medmnist, info["python_class"])51 52 try:53 dataset = data_class(split=split, download=True, size=img_size)54 imgs = dataset.imgs55 if imgs.shape[1] != img_size:56 imgs = _resize_stack(imgs, img_size)57 except TypeError:58 # 古いmedmnistバージョンは size 引数未対応 -> 28x28で取得しリサイズする59 dataset = data_class(split=split, download=True)60 imgs = _resize_stack(dataset.imgs, img_size)61 62 if limit is not None:63 imgs = imgs[:limit]64 return imgs65 66 67class NormalImageDataset(Dataset):68 """正常画像(全画素0-1に正規化)のDataset。"""69 70 def __init__(self, imgs: np.ndarray):71 self.imgs = imgs.astype(np.float32) / 255.072 73 def __len__(self) -> int:74 return len(self.imgs)75 76 def __getitem__(self, idx: int) -> torch.Tensor:77 return torch.from_numpy(self.imgs[idx]).unsqueeze(0)78 79 80def train(args: argparse.Namespace) -> None:81 device = "cuda" if torch.cuda.is_available() else "cpu"82 print(f"device: {device}")83 84 print("Loading OrganAMNIST train split (normal images)...")85 train_imgs = load_organamnist_split("train", args.img_size, args.limit_samples)86 print(f"train images: {len(train_imgs)}")87 88 print("Loading OrganAMNIST val split (calibration)...")89 val_limit = args.limit_samples // 4 if args.limit_samples else None90 val_imgs = load_organamnist_split("val", args.img_size, val_limit)91 print(f"val images: {len(val_imgs)}")92 93 train_ds = NormalImageDataset(train_imgs)94 val_ds = NormalImageDataset(val_imgs)95 96 train_loader = DataLoader(97 train_ds, batch_size=args.batch_size, shuffle=True, num_workers=098 )99 val_loader = DataLoader(100 val_ds, batch_size=args.batch_size, shuffle=False, num_workers=0101 )102 103 model = ConvAutoencoder(img_size=args.img_size, latent_dim=args.latent_dim).to(104 device105 )106 optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)107 criterion = nn.MSELoss()108 109 MODEL_DIR.mkdir(exist_ok=True)110 111 start = time.time()112 for epoch in range(1, args.epochs + 1):113 model.train()114 epoch_start = time.time()115 running_loss = 0.0116 for batch in train_loader:117 batch = batch.to(device)118 optimizer.zero_grad()119 recon = model(batch)120 loss = criterion(recon, batch)121 loss.backward()122 optimizer.step()123 running_loss += loss.item() * batch.size(0)124 train_loss = running_loss / len(train_ds)125 126 model.eval()127 val_loss = 0.0128 with torch.no_grad():129 for batch in val_loader:130 batch = batch.to(device)131 recon = model(batch)132 loss = criterion(recon, batch)133 val_loss += loss.item() * batch.size(0)134 val_loss /= len(val_ds)135 136 epoch_time = time.time() - epoch_start137 total_elapsed = (time.time() - start) / 60138 print(139 f"epoch {epoch}/{args.epochs} train_mse={train_loss:.5f} "140 f"val_mse={val_loss:.5f} epoch_time={epoch_time:.1f}s "141 f"total={total_elapsed:.1f}min"142 )143 144 torch.save(model.state_dict(), WEIGHTS_PATH)145 print(f"saved weights to {WEIGHTS_PATH}")146 147 # 校正: 検証(正常)画像ごとの再構成誤差(MSE)分布からスコア化の基準値を求める148 model.eval()149 per_image_mse = []150 with torch.no_grad():151 for batch in val_loader:152 batch = batch.to(device)153 recon = model(batch)154 mse = torch.mean((recon - batch) ** 2, dim=(1, 2, 3))155 per_image_mse.extend(mse.cpu().numpy().tolist())156 per_image_mse = np.array(per_image_mse)157 158 meta = {159 "img_size": args.img_size,160 "latent_dim": args.latent_dim,161 "normal_mse_mean": float(per_image_mse.mean()),162 "normal_mse_std": float(per_image_mse.std()),163 "normal_mse_p50": float(np.percentile(per_image_mse, 50)),164 "normal_mse_p95": float(np.percentile(per_image_mse, 95)),165 "normal_mse_p99": float(np.percentile(per_image_mse, 99)),166 "threshold_mse": float(np.percentile(per_image_mse, 99)),167 "epochs": args.epochs,168 "train_images": len(train_ds),169 "val_images": len(val_ds),170 }171 META_PATH.write_text(172 json.dumps(meta, ensure_ascii=False, indent=2), encoding="utf-8"173 )174 print(f"saved calibration metadata to {META_PATH}")175 print(json.dumps(meta, ensure_ascii=False, indent=2))176 177 178def parse_args() -> argparse.Namespace:179 parser = argparse.ArgumentParser(180 description="OrganAMNIST正常画像のみでConv Autoencoderを学習する"181 )182 parser.add_argument("--epochs", type=int, default=15)183 parser.add_argument("--batch-size", type=int, default=128)184 parser.add_argument("--lr", type=float, default=1e-3)185 parser.add_argument("--img-size", type=int, default=IMG_SIZE)186 parser.add_argument("--latent-dim", type=int, default=LATENT_DIM)187 parser.add_argument(188 "--limit-samples",189 type=int,190 default=None,191 help="動作確認用に学習画像数を制限する(例: 2000)。val集合はその1/4に制限される。",192 )193 return parser.parse_args()194 195 196if __name__ == "__main__":197 train(parse_args())198 