Spaces:
Build error
Build error
| import torch | |
| import torch.nn as nn | |
| import torch.optim as optim | |
| from torchvision import transforms | |
| from torch.utils.data import Dataset, DataLoader | |
| from PIL import Image | |
| import requests | |
| from io import BytesIO | |
| from pymongo import MongoClient | |
| from sklearn.model_selection import train_test_split | |
| from sklearn.metrics import accuracy_score, f1_score | |
| import segmentation_models_pytorch as smp | |
| from torchvision.models import densenet121 | |
| import gradio as gr | |
| import json | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| client = MongoClient("mongodb+srv://xcovidinsight:FYP25S108@cluster0.vqesy.mongodb.net/") | |
| db = client["covid_system"] | |
| collection = db["validated_xrays"] | |
| class UNet(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.model = smp.Unet( | |
| encoder_name="resnet34", | |
| encoder_weights="imagenet", | |
| in_channels=1, | |
| classes=1 | |
| ) | |
| def forward(self, x): | |
| return self.model(x) | |
| class MyClassifier(nn.Module): | |
| def __init__(self, num_classes=4): | |
| super().__init__() | |
| from torchvision.models import DenseNet121_Weights | |
| base_model = densenet121(weights=DenseNet121_Weights.IMAGENET1K_V1) | |
| in_features = base_model.classifier.in_features | |
| base_model.classifier = nn.Linear(in_features, num_classes) | |
| self.model = base_model | |
| def forward(self, x): | |
| return self.model(x) | |
| def fetch_image(url): | |
| try: | |
| response = requests.get(url) | |
| image = Image.open(BytesIO(response.content)).convert("L").resize((256, 256)) | |
| return image | |
| except Exception as e: | |
| print(f"Failed to fetch image from {url}: {e}") | |
| return None | |
| def trigger_incremental_train(): | |
| samples = list(collection.find()) | |
| samples = [s for s in samples if "image_path" in s and "true_label" in s] | |
| if not samples or len(samples) < 100: | |
| return {"error": "Not enough validated samples (minimum 100 required)."} | |
| m1 = UNet() | |
| m1.model.load_state_dict(torch.load("Segmentation Model.pth")) | |
| m1.to(device).eval() | |
| m2_old = MyClassifier() | |
| m2_old.load_state_dict(torch.load("Classification Model.pth", map_location=device)) | |
| m2_old.to(device).eval() | |
| m2 = MyClassifier() | |
| m2.load_state_dict(torch.load("Classification Model.pth", map_location=device)) | |
| m2.to(device) | |
| class LungXrayDataset(Dataset): | |
| def __init__(self, entries, transform): | |
| self.entries = entries | |
| self.transform = transform | |
| self.label_map = {label: i for i, label in enumerate(sorted(set(s["true_label"] for s in entries)))} | |
| def __len__(self): | |
| return len(self.entries) | |
| def __getitem__(self, idx): | |
| entry = self.entries[idx] | |
| image = fetch_image(entry["image_path"]) | |
| if image is None: | |
| raise RuntimeError(f"Image at {entry['image_path']} could not be loaded.") | |
| image_tensor = transforms.ToTensor()(image).unsqueeze(0).to(device) | |
| with torch.no_grad(): | |
| mask = m1(image_tensor).sigmoid() | |
| mask = (mask > 0.5).float() | |
| masked = image_tensor * mask | |
| masked_rgb = masked.squeeze(0).repeat(3, 1, 1).cpu() | |
| if self.transform: | |
| masked_rgb = self.transform(masked_rgb) | |
| label = self.label_map[entry["true_label"]] | |
| return masked_rgb, label | |
| transform = transforms.Compose([ | |
| transforms.Resize((224, 224)), | |
| transforms.Normalize(mean=[0.5]*3, std=[0.5]*3) | |
| ]) | |
| train_entries, val_entries = train_test_split( | |
| samples, test_size=0.2, stratify=[s["true_label"] for s in samples], random_state=42 | |
| ) | |
| train_ds = LungXrayDataset(train_entries, transform) | |
| val_ds = LungXrayDataset(val_entries, transform) | |
| train_loader = DataLoader(train_ds, batch_size=16, shuffle=True) | |
| val_loader = DataLoader(val_ds, batch_size=16) | |
| criterion = nn.CrossEntropyLoss() | |
| optimizer = optim.Adam(m2.parameters(), lr=1e-4) | |
| for epoch in range(5): | |
| m2.train() | |
| total_loss, correct = 0, 0 | |
| for imgs, labels in train_loader: | |
| imgs, labels = imgs.to(device), labels.to(device) | |
| out = m2(imgs) | |
| loss = criterion(out, labels) | |
| optimizer.zero_grad() | |
| loss.backward() | |
| optimizer.step() | |
| total_loss += loss.item() | |
| correct += (out.argmax(1) == labels).sum().item() | |
| acc = correct / len(train_ds) | |
| print(f"Epoch {epoch+1} ➤ Loss: {total_loss:.4f} | Accuracy: {acc:.4f}") | |
| def evaluate(model, loader, version): | |
| model.eval() | |
| y_true, y_pred = [], [] | |
| with torch.no_grad(): | |
| for imgs, labels in loader: | |
| imgs = imgs.to(device) | |
| outputs = model(imgs) | |
| y_pred.extend(outputs.argmax(1).cpu().numpy()) | |
| y_true.extend(labels.numpy()) | |
| return { | |
| "version": version, | |
| "accuracy": round(accuracy_score(y_true, y_pred), 4), | |
| "f1_macro": round(f1_score(y_true, y_pred, average="macro"), 4), | |
| } | |
| eval_old = evaluate(m2_old, val_loader, "Old") | |
| eval_new = evaluate(m2, val_loader, "New") | |
| torch.save(m2.state_dict(), "New Classification Model.pth") | |
| return { | |
| "old_model": eval_old, | |
| "new_model": eval_new | |
| } | |
| with gr.Blocks() as demo: | |
| with gr.Row(): | |
| train_button = gr.Button("Start Incremental Training") | |
| output_json = gr.JSON(label="Training Result") | |
| train_button.click(fn=trigger_incremental_train, inputs=[], outputs=output_json) | |
| demo.launch(server_name="0.0.0.0", server_port=7860) | |