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.")