118 lines
3.9 KiB
Python
118 lines
3.9 KiB
Python
import os
|
|
import torch
|
|
import torch.nn as nn
|
|
from torch.utils.data import DataLoader, Dataset
|
|
from torchvision import models, transforms
|
|
from PIL import Image
|
|
import numpy as np
|
|
from sklearn.manifold import TSNE
|
|
import pandas as pd
|
|
import plotly.express as px
|
|
|
|
# 1. CONFIGURACIÓN
|
|
IMG_DIR = "/home/sierra/Documentos/Prueba teams/Julia_msi_GUI-main/datos_para_ai"
|
|
ENCODER_PATH = "Reporte/Model/msi_encoder_trained.pth"
|
|
CSV_PATH = "msi_clusters_results.csv" # Fuente de verdad
|
|
N_CLUSTERS = 10
|
|
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
|
|
# 2. CARGAR EL MODELO ENTRENADO
|
|
def load_trained_encoder(path):
|
|
if not os.path.exists(path):
|
|
raise FileNotFoundError(f"No se encontró el modelo en {path}")
|
|
|
|
model = models.efficientnet_b0()
|
|
model.classifier = nn.Identity()
|
|
|
|
state_dict = torch.load(path, map_location=DEVICE)
|
|
model.load_state_dict(state_dict)
|
|
model.to(DEVICE)
|
|
model.eval()
|
|
return model
|
|
|
|
# 3. DATASET BASADO EN EL CSV
|
|
class MSIDatasetFromCSV(Dataset):
|
|
def __init__(self, img_dir, img_names):
|
|
self.img_dir = img_dir
|
|
self.img_names = img_names
|
|
self.transform = transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.ToTensor(),
|
|
])
|
|
|
|
def __len__(self):
|
|
return len(self.img_names)
|
|
|
|
def __getitem__(self, idx):
|
|
img_path = os.path.join(self.img_dir, self.img_names[idx])
|
|
image = Image.open(img_path).convert("RGB")
|
|
return self.transform(image)
|
|
|
|
# --- PROCESO ---
|
|
if __name__ == "__main__":
|
|
# 4. CARGAR EL CSV DE RESULTADOS
|
|
print(f"Cargando datos desde {CSV_PATH}...")
|
|
df = pd.read_csv(CSV_PATH)
|
|
ion_names = df['ion_image'].tolist()
|
|
|
|
# Asegurarnos de que cluster sea categórico para Plotly
|
|
df['cluster'] = df['cluster'].astype(str)
|
|
|
|
print("Extrayendo características con el encoder...")
|
|
encoder = load_trained_encoder(ENCODER_PATH)
|
|
dataset = MSIDatasetFromCSV(IMG_DIR, ion_names)
|
|
loader = DataLoader(dataset, batch_size=1, shuffle=False)
|
|
|
|
features = []
|
|
with torch.no_grad():
|
|
for i, img in enumerate(loader):
|
|
img = img.to(DEVICE)
|
|
feat = encoder(img)
|
|
# Normalización L2
|
|
feat = feat / feat.norm(p=2, dim=1, keepdim=True)
|
|
features.append(feat.cpu().numpy().flatten())
|
|
if (i+1) % 100 == 0:
|
|
print(f"Iones procesados: {i+1}/{len(ion_names)}")
|
|
|
|
features = np.array(features)
|
|
|
|
# 5. t-SNE
|
|
print("Calculando t-SNE...")
|
|
tsne = TSNE(n_components=2, perplexity=30, random_state=42)
|
|
embeddings_2d = tsne.fit_transform(features)
|
|
|
|
# Agregar coordenadas t-SNE al DataFrame original
|
|
df['tsne_1'] = embeddings_2d[:, 0]
|
|
df['tsne_2'] = embeddings_2d[:, 1]
|
|
|
|
# 6. VISUALIZACIÓN INTERACTIVA CON PLOTLY
|
|
print("Generando gráfica interactiva con Plotly...")
|
|
|
|
# Usar escala de colores similar a tab10 o cualitativa estándar
|
|
fig = px.scatter(
|
|
df,
|
|
x='tsne_1',
|
|
y='tsne_2',
|
|
color='cluster',
|
|
hover_data={'ion_image': True, 'mz': True, 'cluster': True, 'tsne_1': False, 'tsne_2': False},
|
|
title=f"t-SNE Interactivo de Iones MSI (Consistente con {CSV_PATH})",
|
|
labels={'tsne_1': 't-SNE 1', 'tsne_2': 't-SNE 2'},
|
|
category_orders={"cluster": [str(i) for i in range(N_CLUSTERS)]},
|
|
color_discrete_sequence=px.colors.qualitative.Plotly # O T10 si está disponible
|
|
)
|
|
|
|
# Mejorar el diseño
|
|
fig.update_traces(marker=dict(size=10, opacity=0.8, line=dict(width=1, color='DarkSlateGrey')))
|
|
fig.update_layout(
|
|
legend_title_text='Cluster ID',
|
|
font=dict(size=14),
|
|
hoverlabel=dict(bgcolor="white", font_size=16)
|
|
)
|
|
|
|
# Guardar como HTML
|
|
output_html = "msi_tsne_interactivo.html"
|
|
fig.write_html(output_html)
|
|
|
|
print(f"\n¡Éxito! Gráfica interactiva guardada en '{output_html}'.")
|
|
print("Puedes abrir este archivo en cualquier navegador para explorar los clusters.")
|