USER
import os
import glob
import pandas as pd
import pydicom as dcm
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
from torchvision import transforms, models
from PIL import Image
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import LabelEncoder
from sklearn.metrics import f1_score, roc_auc_score, confusion_matrix, balanced_accuracy_score
import wandb
from torch.optim.lr_scheduler import StepLR
from torchvision.transforms import RandomAffine, InterpolationMode
import numpy as np
from imblearn.over_sampling import SMOTE
class FocalLoss(nn.Module):
def __init__(self, alpha=1, gamma=2):
super().__init__()
self.alpha = alpha
self.gamma = gamma
def forward(self, inputs, targets):
BCE_loss = nn.CrossEntropyLoss()(inputs, targets)
pt = torch.exp(-BCE_loss)
F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss
return F_loss
# Initialize wandb
wandb.init(project="mamo-clinicalfeatures-resnetest", name="teste7")
wandb.config.update({"num_epochs": 15, "batch_size": 128, "learning_rate": 0.001})
# Constants
root_dir = 'CHData/manifest-1616439774456/CMMD'
clinical_data_path = 'clinical.csv' # Update this path
num_classes = 2
clinical_data = pd.read_csv(clinical_data_path, sep=';')
unique_ids = clinical_data['ID1'].unique()
train_ids, temp_ids = train_test_split(unique_ids, test_size=0.3, stratify=clinical_data.groupby('ID1').first()['classification'], random_state=42)
val_ids, test_ids = train_test_split(temp_ids, test_size=0.5, stratify=clinical_data[clinical_data['ID1'].isin(temp_ids)].groupby('ID1').first()['classification'], random_state=42)
# Remove duplicates
clinical_data = clinical_data.drop_duplicates(subset=['ID1'])
# Create datasets with only 'classification' column
train_data = clinical_data[clinical_data['ID1'].isin(train_ids)][['ID1', 'classification']]
val_data = clinical_data[clinical_data['ID1'].isin(val_ids)][['ID1', 'classification']]
test_data = clinical_data[clinical_data['ID1'].isin(test_ids)][['ID1', 'classification']]
# Label Encoding for 'classification'
le_classification = LabelEncoder()
le_classification.fit(train_data['classification'])
train_data['classification'] = le_classification.transform(train_data['classification'])
val_data['classification'] = le_classification.transform(val_data['classification'])
test_data['classification'] = le_classification.transform(test_data['classification'])
# Improved Transforms
def define_transforms():
return transforms.Compose([
transforms.Lambda(lambda x: x.convert("RGB")), # Convert to RGB
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# Custom dataset class
class MammographyDataset(Dataset):
def __init__(self, data, transform=None):
self.data = data # Initialize data here
self.transform = transform
self.samples = []
for id1 in data['ID1'].unique():
id1_data = data[data['ID1'] == id1]
label = id1_data.iloc[0]['classification']
dcm_files = glob.glob(os.path.join(root_dir, f"{id1}", '*', '*', '*.dcm'))
for dcm_file in dcm_files:
dicom_data = dcm.read_file(dcm_file)
self.samples.append((dcm_file, label))
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
dcm_file, label = self.samples[idx]
image = dcm.read_file(dcm_file).pixel_array
image = image[np.newaxis, :, :] # Add a channel dimension
image = Image.fromarray(image[0]) # Convert to PIL Image for transformations
if self.transform:
image = self.transform(image)
if isinstance(label, str):
label = le_classification.transform([label])[0]
label = torch.tensor(label, dtype=torch.long)
return image, label
# Custom model class
class CustomModel(nn.Module):
def __init__(self, num_classes):
#def __init__(self, num_features, num_classes):
super().__init__()
self.base_model = models.resnet18(pretrained=True)
self.classifier = nn.Sequential(
nn.Linear(512, 64),
nn.Dropout(0.5),
nn.ReLU(),
nn.Linear(64, num_classes)
)
self.base_model.fc = nn.Identity()
def forward(self, x):
x = self.base_model(x)
x = self.classifier(x)
return x
# Function to calculate metrics
def calculate_extended_metrics(y_true, y_pred):
tn, fp, fn, tp = confusion_matrix(y_true, y_pred).ravel()
sensitivity = tp / (tp + fn)
specificity = tn / (tn + fp)
f1 = f1_score(y_true, y_pred)
auc = roc_auc_score(y_true, y_pred)
return sensitivity, specificity, f1, auc
# Split data
train_data, temp_data = train_test_split(clinical_data, test_size=0.3, stratify=clinical_data['classification'], random_state=42)
val_data, test_data = train_test_split(temp_data, test_size=0.5, stratify=temp_data['classification'], random_state=42)
# Define transformations
transform = define_transforms()
# Now you can use train_data_balanced
train_dataset = MammographyDataset(train_data, transform)
#Debug remover
sample = train_dataset[0]
val_dataset = MammographyDataset(val_data, transform)
test_dataset = MammographyDataset(test_data, transform)
# Create dataloaders
batch_size = wandb.config.batch_size
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=24)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=24)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=24)
# Define model and move to device
#model = CustomModel(num_features=2, num_classes=num_classes)
model = CustomModel(num_classes=num_classes)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
model = model.to(device)
weights = torch.tensor([0.3, 0.7], dtype=torch.float32) # Ajuste esses valores com base na distribuição da sua classe
weights = weights.to(device)
criterion = torch.nn.CrossEntropyLoss(weight=weights)
optimizer = torch.optim.Adam(model.parameters(), lr=wandb.config.learning_rate)
# Learning Rate Scheduler
scheduler = StepLR(optimizer, step_size=10, gamma=0.7)
# Training loop with Early Stopping, Logging, and Metrics
best_val_loss = float('inf')
early_stop_count = 0
print("Starting the training loop...")
for epoch in range(wandb.config.num_epochs):
# Initialize metrics for this epoch
running_loss = 0.0
running_corrects = 0
# Training Phase
model.train()
for inputs, labels in train_loader:
# Move data to device and zero the gradients
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
# Forward pass and loss computation
outputs = model(inputs)
loss = criterion(outputs, labels)
# Backward pass and optimization
loss.backward()
optimizer.step()
# Update running metrics
running_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
running_corrects += torch.sum(preds == labels.data)
# Step the learning rate scheduler
scheduler.step()
# Calculate epoch level metrics for training set
epoch_loss = running_loss / len(train_loader.dataset)
epoch_acc = running_corrects.double() / len(train_loader.dataset)
# Validation Phase
model.eval()
val_loss = 0.0
val_corrects = 0
all_preds = []
all_labels = []
total_tp = 0
total_tn = 0
total_fp = 0
total_fn = 0
with torch.no_grad():
for inputs, labels in val_loader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = model(inputs)
loss = criterion(outputs, labels)
val_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
val_corrects += torch.sum(preds == labels.data)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.cpu().numpy())
# Calculate epoch level metrics for validation set
val_epoch_loss = val_loss / len(val_loader.dataset)
val_epoch_acc = val_corrects.double() / len(val_loader.dataset)
val_sensitivity, val_specificity, val_f1, val_auc = calculate_extended_metrics(all_labels, all_preds)
val_balanced_accuracy = balanced_accuracy_score(all_labels, all_preds)
print("Confusion Matrix:", confusion_matrix(all_labels, all_preds))
# Early Stopping based on validation loss
if val_epoch_loss < best_val_loss:
best_val_loss = val_epoch_loss
early_stop_count = 0
else:
early_stop_count += 1
if early_stop_count >= 5:
print(f"Early stopping. Best validation loss: {best_val_loss}")
break
# Log metrics to wandb
wandb.log({
"epoch_loss": epoch_loss,
"epoch_acc": epoch_acc,
"val_loss": val_epoch_loss,
"val_acc": val_epoch_acc,
"Sensitivity": val_sensitivity,
"Specificity": val_specificity,
"F1_Score": val_f1,
"AUC_ROC": val_auc,
"Balanced_Accuracy": val_balanced_accuracy
})
# Save the model and update wandb config
wandb.save("model.pth")
wandb.config.update({
"model_architecture": str(model),
"optimizer": str(optimizer),
"criterion": str(criterion),
"scheduler": str(scheduler)
})
Meu codigo está com erro que o Está saindo com specificity 0. Auc roc 0.5. F1 Score 1 Sensitivity 1.0 Vall acc 0.79. Vall loss 0.12. Epoch looss 0.051 epoch acc em 0.27
O problema parece estar nos dados, mas não tenho certeza