Tanusree88/ViT-MRI-FineTuning
0
1import os2import zipfile3import numpy as np4import torch5from transformers import SegformerForImageSegmentation, ResNetForImageClassification, AdamW6from PIL import Image7from torch.utils.data import Dataset, DataLoader8import streamlit as st9import gradio as gr10 11gr.load("models/nvidia/segformer-b0-finetuned-ade-512-512").launch()12 13# Function to extract zip files14def extract_zip(zip_file, extract_to):15 with zipfile.ZipFile(zip_file, 'r') as zip_ref:16 zip_ref.extractall(extract_to)17 18# Preprocess images19def preprocess_image(image_path):20 ext = os.path.splitext(image_path)[-1].lower()21 22 if ext == '.npy':23 image_data = np.load(image_path)24 image_tensor = torch.tensor(image_data).float()25 if len(image_tensor.shape) == 3:26 image_tensor = image_tensor.unsqueeze(0)27 28 elif ext in ['.jpg', '.jpeg']:29 img = Image.open(image_path).convert('RGB').resize((224, 224))30 img_np = np.array(img)31 image_tensor = torch.tensor(img_np).permute(2, 0, 1).float()32 33 else:34 raise ValueError(f"Unsupported format: {ext}")35 36 image_tensor /= 255.0 # Normalize to [0, 1]37 return image_tensor38 39# Prepare dataset40def prepare_dataset(extracted_folder):41 neuronii_path = os.path.join(extracted_folder, "neuroniiimages")42 43 if not os.path.exists(neuronii_path):44 raise FileNotFoundError(f"The folder neuroniiimages does not exist in the extracted folder: {neuronii_path}")45 46 image_paths = []47 labels = []48 49 for disease_folder in ['alzheimers_dataset', 'parkinsons_dataset', 'MSjpg']:50 folder_path = os.path.join(neuronii_path, disease_folder)51 52 if not os.path.exists(folder_path):53 print(f"Folder not found: {folder_path}")54 continue 55 label = {'alzheimers_dataset': 0, 'parkinsons_dataset': 1, 'MSjpg': 2}[disease_folder]56 57 for img_file in os.listdir(folder_path):58 if img_file.endswith(('.npy', '.jpg', '.jpeg')):59 image_paths.append(os.path.join(folder_path, img_file))60 labels.append(label)61 else:62 print(f"Unsupported file: {img_file}")63 print(f"Total images loaded: {len(image_paths)}")64 return image_paths, labels65 66# Custom Dataset class67class CustomImageDataset(Dataset):68 def __init__(self, image_paths, labels):69 self.image_paths = image_paths70 self.labels = labels71 72 def __len__(self):73 return len(self.image_paths)74 75 def __getitem__(self, idx):76 image = preprocess_image(self.image_paths[idx])77 label = self.labels[idx]78 return image, label79 80# Training function for classification81def fine_tune_classification_model(train_loader):82 model = ResNetForImageClassification.from_pretrained('microsoft/resnet-50', num_labels=3)83 model.train()84 optimizer = AdamW(model.parameters(), lr=1e-4)85 criterion = torch.nn.CrossEntropyLoss()86 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')87 model.to(device)88 89 for epoch in range(10):90 running_loss = 0.091 for images, labels in train_loader:92 images, labels = images.to(device), labels.to(device)93 optimizer.zero_grad()94 outputs = model(pixel_values=images).logits95 loss = criterion(outputs, labels)96 loss.backward()97 optimizer.step()98 running_loss += loss.item()99 return running_loss / len(train_loader)100 101# Streamlit UI for Fine-tuning102st.title("Fine-tune ResNet for MRI/CT Scans Classification")103 104zip_file_url = "https://huggingface.co/spaces/Tanusree88/ViT-MRI-FineTuning/resolve/main/neuroniiimages.zip"105 106if st.button("Start Training"):107 extraction_dir = "extracted_files"108 os.makedirs(extraction_dir, exist_ok=True)109 110 # Download the zip file (placeholder)111 zip_file = "neuroniiimages.zip" # Assuming you downloaded it with this name112 113 # Extract zip file114 extract_zip(zip_file, extraction_dir)115 116 # Prepare dataset117 image_paths, labels = prepare_dataset(extraction_dir)118 dataset = CustomImageDataset(image_paths, labels)119 train_loader = DataLoader(dataset, batch_size=32, shuffle=True)120 121 # Fine-tune the classification model122 final_loss = fine_tune_classification_model(train_loader)123 st.write(f"Training Complete with Final Loss: {final_loss}")124 125# Segmentation function (using SegFormer)126def fine_tune_segmentation_model(train_loader):127 model = SegformerForImageSegmentation.from_pretrained('nvidia/segformer-b0', num_labels=3)128 model.train()129 optimizer = AdamW(model.parameters(), lr=1e-4)130 criterion = torch.nn.CrossEntropyLoss()131 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')132 model.to(device)133 134 for epoch in range(10):135 running_loss = 0.0136 for images, labels in train_loader:137 images, labels = images.to(device), labels.to(device)138 optimizer.zero_grad()139 outputs = model(pixel_values=images).logits140 loss = criterion(outputs, labels)141 loss.backward()142 optimizer.step()143 running_loss += loss.item()144 return running_loss / len(train_loader)145 146# Add a button for segmentation training147if st.button("Start Segmentation Training"):148 # Assuming the dataset for segmentation is prepared similarly149 seg_train_loader = DataLoader(dataset, batch_size=32, shuffle=True)150 151 # Fine-tune the segmentation model152 final_loss_seg = fine_tune_segmentation_model(seg_train_loader)153 st.write(f"Segmentation Training Complete with Final Loss: {final_loss_seg}")154 155 