MSI_Julia_CNN/scripts_python/legacy/graficar_tsne_plotly.py

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