CoolFace
Apppublic

simoyaman/Ct_anomaly_detector

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes
train_autoencoder.py198 linesDownload Raw Back to root
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