import os import argparse from glob import glob from PIL import Image import numpy as np import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader import torchvision.transforms as transforms from tqdm import tqdm # -------- Dataset -------- class RadiologyDataset(Dataset): def __init__(self, img_dir, img_size): self.files = glob(os.path.join(img_dir, '*.png')) + glob(os.path.join(img_dir, '*.jpg')) if not self.files: raise RuntimeError(f"پوشه '{img_dir}' خالی است؛ لطفاً تصاویر را قرار دهید.") self.transform = transforms.Compose([ transforms.Resize((img_size, img_size)), transforms.Grayscale(num_output_channels=1), transforms.ToTensor() ]) def __len__(self): return len(self.files) def __getitem__(self, idx): path = self.files[idx] img = Image.open(path).convert('L') x = self.transform(img) return x, idx # -------- Autoencoder -------- class SimpleAE(nn.Module): def __init__(self): super(SimpleAE, self).__init__() self.encoder = nn.Sequential( nn.Conv2d(1, 8, 3, stride=2, padding=1), nn.ReLU(True), # ->8x64x64 nn.Conv2d(8, 16, 3, stride=2, padding=1), nn.ReLU(True) # ->16x32x32 ) self.decoder = nn.Sequential( nn.ConvTranspose2d(16, 8, 3, stride=2, padding=1, output_padding=1), nn.ReLU(True), # ->8x64x64 nn.ConvTranspose2d(8, 1, 3, stride=2, padding=1, output_padding=1), nn.Sigmoid() # ->1x128x128 ) def forward(self, x): z = self.encoder(x) return self.decoder(z) # -------- Train AE -------- def train_autoencoder(dataset, device, epochs, batch_size, lr): loader = DataLoader(dataset, batch_size=batch_size, shuffle=True) model = SimpleAE().to(device) criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=lr) model.train() for epoch in range(epochs): epoch_loss = 0 for x, _ in loader: x = x.to(device) optimizer.zero_grad() recon = model(x) loss = criterion(recon, x) loss.backward() optimizer.step() epoch_loss += loss.item() * x.size(0) print(f"Epoch {epoch+1}/{epochs} - Loss: {epoch_loss/len(dataset):.6f}") return model # -------- Compute Errors & Score -------- def evaluate_quality(model, dataset, device, threshold): loader = DataLoader(dataset, batch_size=1, shuffle=False) errors = [] with torch.no_grad(): for x, idx in tqdm(loader, desc="Evaluating"): # single-batch for each image x = x.to(device) recon = model(x) err = ((recon - x)**2).mean().item() errors.append((idx.item(), err)) # determine threshold if not provided if threshold is None: errs = [e for _, e in errors] threshold = np.percentile(errs, 95) print(f"Error threshold set at 95th percentile: {threshold:.6f}") good = sum(1 for _, e in errors if e <= threshold) total = len(errors) score = good / total * 100 return score, threshold, errors # -------- Main -------- def main(args): device = torch.device('cpu') dataset = RadiologyDataset(args.data_dir, args.img_size) print("Training Autoencoder...") ae = train_autoencoder(dataset, device, args.epochs, args.batch_size, args.lr) ae.eval() print("Evaluating quality...") score, threshold, _ = evaluate_quality(ae, dataset, device, None) print(f"Quality Score: {score:.2f}% good images") if score < 95: print("هشدار: درصد تصاویر با کیفیت کمتر از 95% شد.") else: print("≥95% تصاویر با کیفیت هستند.") # save model and threshold os.makedirs(args.save_dir, exist_ok=True) torch.save({'model_state': ae.state_dict(), 'threshold': threshold}, os.path.join(args.save_dir, 'autoencoder_qc.pth')) print(f"Model and threshold saved to {args.save_dir}") # -------- CLI -------- if __name__ == '__main__': parser = argparse.ArgumentParser(description='Unsupervised QA via Autoencoder') parser.add_argument('--data_dir', type=str, default='./data', help='Path to images') parser.add_argument('--save_dir', type=str, default='./models', help='Save directory') parser.add_argument('--img_size', type=int, default=128, help='Resize images') parser.add_argument('--epochs', type=int, default=10, help='AE training epochs') parser.add_argument('--batch_size', type=int, default=16, help='Batch size') parser.add_argument('--lr', type=float, default=1e-3, help='Learning rate') args = parser.parse_args() main(args)