Radiology-Quality-Control / TrainModel.py
TahaGorji's picture
Upload 8 files
05e9fbe verified
Raw
History Blame Contribute Delete
4.97 kB
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)