Skip to main content

🌐 t-SNE: Visualizing High-Dimensional Data

šŸ“š What You'll Learn

By the end of this lesson, you will be able to:

  • Explain how t-SNE preserves local neighborhood structure to reveal clusters in 2D or 3D
  • Tune the key hyperparameters — perplexity, learning rate, and number of iterations
  • Understand why distances between clusters and cluster sizes in a t-SNE map should not be over-interpreted
  • Apply TSNE in scikit-learn, often running PCA first for speed and stability
  • Recognize that t-SNE is a visualization tool, not a general transform for new, unseen data

ā±ļø Estimated Time: 45–60 minutes

šŸŽÆ Project: Visualize a high-dimensional dataset (such as handwritten digits) in 2D with t-SNE and experiment with perplexity to see how the map changes.

Introduction

t-SNE (t-Distributed Stochastic Neighbor Embedding) is a powerful non-linear dimensionality reduction technique particularly well-suited for visualizing high-dimensional data. Unlike PCA which preserves global structure, t-SNE excels at preserving local neighborhoods, making it ideal for exploring clusters and patterns in complex datasets. It's become the go-to method for visualizing everything from word embeddings to single-cell RNA sequences.

Core Concepts and Theory

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.manifold import TSNE
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
from sklearn.datasets import load_digits, load_iris, fetch_openml
from sklearn.metrics import pairwise_distances
from scipy.spatial.distance import pdist, squareform
import time
import warnings
warnings.filterwarnings('ignore')

# Set style
plt.style.use('seaborn-v0_8-darkgrid')
sns.set_palette("husl")

print("="*60)
print("t-SNE FUNDAMENTALS")
print("="*60)

# Core concepts
tsne_concepts = """
t-SNE KEY CONCEPTS:

1. ALGORITHM OVERVIEW:
   • Maps high-dimensional points to low dimensions (usually 2D)
   • Preserves local structure (nearby points stay nearby)
   • Non-linear transformation
   • Probabilistic approach

2. HOW IT WORKS:
   Step 1: Calculate pairwise similarities in high-D space (Gaussian)
   Step 2: Calculate pairwise similarities in low-D space (t-distribution)
   Step 3: Minimize KL divergence between distributions
   Step 4: Use gradient descent to optimize

3. KEY PARAMETERS:
   • perplexity: Balance between local and global aspects (5-50)
   • learning_rate: Step size for gradient descent (10-1000)
   • n_iter: Number of iterations (250-5000)
   • metric: Distance metric for high-D space

4. PERPLEXITY:
   • Roughly the number of neighbors considered
   • Low values: Focus on local structure
   • High values: More global structure
   • Rule of thumb: 5-50, dataset_size/100

5. ADVANTAGES:
   • Excellent for visualization
   • Reveals clusters and patterns
   • Handles non-linear relationships
   • Works well with many data types

6. LIMITATIONS:
   • Computational complexity O(n²)
   • Non-deterministic (random initialization)
   • Cannot transform new points
   • Hyperparameter sensitive
   • Preserves neighborhoods, not distances
"""

print(tsne_concepts)

Basic t-SNE Implementation and Visualization

class TSNEVisualizer:
    """Comprehensive t-SNE visualization and analysis"""
    
    def __init__(self):
        self.embeddings = {}
        self.models = {}
        
    def compare_perplexities(self, X, y=None, perplexities=[5, 30, 50, 100]):
        """Compare t-SNE with different perplexity values"""
        
        fig, axes = plt.subplots(2, len(perplexities)//2, figsize=(12, 10))
        axes = axes.ravel()
        
        for idx, perp in enumerate(perplexities):
            print(f"Running t-SNE with perplexity={perp}...")
            
            # Run t-SNE
            tsne = TSNE(n_components=2, perplexity=perp, 
                       random_state=42, n_iter=1000)
            X_embedded = tsne.fit_transform(X)
            
            # Store results
            self.embeddings[f'perp_{perp}'] = X_embedded
            self.models[f'perp_{perp}'] = tsne
            
            # Plot
            if y is not None:
                scatter = axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                          c=y, cmap='viridis', alpha=0.6, s=30)
                plt.colorbar(scatter, ax=axes[idx])
            else:
                axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                alpha=0.6, s=30)
            
            axes[idx].set_title(f'Perplexity = {perp}')
            axes[idx].set_xlabel('t-SNE 1')
            axes[idx].set_ylabel('t-SNE 2')
            axes[idx].grid(True, alpha=0.3)
            
            # Add KL divergence if available
            kl_div = tsne.kl_divergence_
            axes[idx].text(0.02, 0.98, f'KL div: {kl_div:.2f}',
                         transform=axes[idx].transAxes,
                         fontsize=9, verticalalignment='top',
                         bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))
        
        plt.suptitle('t-SNE: Effect of Perplexity Parameter', fontsize=14, y=1.02)
        plt.tight_layout()
        plt.show()
        
        return self.embeddings
    
    def compare_learning_rates(self, X, y=None, learning_rates=[10, 50, 200, 1000]):
        """Compare different learning rates"""
        
        fig, axes = plt.subplots(2, 2, figsize=(12, 10))
        axes = axes.ravel()
        
        for idx, lr in enumerate(learning_rates):
            print(f"Running t-SNE with learning_rate={lr}...")
            
            # Run t-SNE
            tsne = TSNE(n_components=2, learning_rate=lr,
                       perplexity=30, random_state=42, n_iter=1000)
            X_embedded = tsne.fit_transform(X)
            
            # Plot
            if y is not None:
                axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                c=y, cmap='viridis', alpha=0.6, s=30)
            else:
                axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                alpha=0.6, s=30)
            
            axes[idx].set_title(f'Learning Rate = {lr}')
            axes[idx].set_xlabel('t-SNE 1')
            axes[idx].set_ylabel('t-SNE 2')
            axes[idx].grid(True, alpha=0.3)
            
            # Add KL divergence
            kl_div = tsne.kl_divergence_
            axes[idx].text(0.02, 0.98, f'KL div: {kl_div:.2f}',
                         transform=axes[idx].transAxes,
                         fontsize=9, verticalalignment='top',
                         bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))
        
        plt.suptitle('t-SNE: Effect of Learning Rate', fontsize=14, y=1.02)
        plt.tight_layout()
        plt.show()
    
    def visualize_convergence(self, X, y=None, n_iter_steps=[250, 500, 1000, 5000]):
        """Visualize t-SNE convergence over iterations"""
        
        fig, axes = plt.subplots(2, 2, figsize=(12, 10))
        axes = axes.ravel()
        
        for idx, n_iter in enumerate(n_iter_steps):
            print(f"Running t-SNE with n_iter={n_iter}...")
            
            # Run t-SNE
            tsne = TSNE(n_components=2, n_iter=n_iter,
                       perplexity=30, random_state=42)
            X_embedded = tsne.fit_transform(X)
            
            # Plot
            if y is not None:
                axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                c=y, cmap='viridis', alpha=0.6, s=30)
            else:
                axes[idx].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                alpha=0.6, s=30)
            
            axes[idx].set_title(f'Iterations = {n_iter}')
            axes[idx].set_xlabel('t-SNE 1')
            axes[idx].set_ylabel('t-SNE 2')
            axes[idx].grid(True, alpha=0.3)
            
            # Add KL divergence
            kl_div = tsne.kl_divergence_
            axes[idx].text(0.02, 0.98, f'KL div: {kl_div:.2f}',
                         transform=axes[idx].transAxes,
                         fontsize=9, verticalalignment='top',
                         bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))
        
        plt.suptitle('t-SNE: Convergence Over Iterations', fontsize=14, y=1.02)
        plt.tight_layout()
        plt.show()
    
    def stability_analysis(self, X, y=None, n_runs=5):
        """Analyze stability across multiple runs"""
        
        fig, axes = plt.subplots(2, 3, figsize=(15, 10))
        axes = axes.ravel()
        
        embeddings = []
        
        for run in range(n_runs):
            print(f"Run {run+1}/{n_runs}...")
            
            # Run t-SNE with different random state
            tsne = TSNE(n_components=2, perplexity=30,
                       random_state=run*42, n_iter=1000)
            X_embedded = tsne.fit_transform(X)
            embeddings.append(X_embedded)
            
            if run < 5:  # Plot first 5 runs
                if y is not None:
                    axes[run].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                    c=y, cmap='viridis', alpha=0.6, s=30)
                else:
                    axes[run].scatter(X_embedded[:, 0], X_embedded[:, 1],
                                    alpha=0.6, s=30)
                
                axes[run].set_title(f'Run {run+1}')
                axes[run].set_xlabel('t-SNE 1')
                axes[run].set_ylabel('t-SNE 2')
                axes[run].grid(True, alpha=0.3)
        
        # Calculate pairwise correlations between embeddings
        correlations = []
        for i in range(n_runs):
            for j in range(i+1, n_runs):
                # Calculate correlation between flattened embeddings
                corr = np.corrcoef(embeddings[i].flatten(), 
                                  embeddings[j].flatten())[0, 1]
                correlations.append(abs(corr))
        
        # Plot correlation distribution
        axes[-1].hist(correlations, bins=20, edgecolor='black', alpha=0.7)
        axes[-1].set_xlabel('Absolute Correlation')
        axes[-1].set_ylabel('Frequency')
        axes[-1].set_title(f'Stability: Mean Corr = {np.mean(correlations):.3f}')
        axes[-1].grid(True, alpha=0.3)
        
        plt.suptitle('t-SNE Stability Analysis', fontsize=14, y=1.02)
        plt.tight_layout()
        plt.show()
        
        print(f"\nStability Metrics:")
        print(f"  Mean correlation: {np.mean(correlations):.3f}")
        print(f"  Std correlation: {np.std(correlations):.3f}")
        
        return embeddings

# Load sample data
digits = load_digits()
X_digits = digits.data
y_digits = digits.target

# Sample for faster computation
sample_idx = np.random.choice(len(X_digits), 500, replace=False)
X_sample = X_digits[sample_idx]
y_sample = y_digits[sample_idx]

# Initialize visualizer
viz = TSNEVisualizer()

print("\n" + "="*60)
print("t-SNE PARAMETER EXPLORATION")
print("="*60)

# Compare perplexities
print("\nComparing different perplexity values...")
perp_embeddings = viz.compare_perplexities(X_sample, y_sample)

# Compare learning rates
print("\nComparing different learning rates...")
viz.compare_learning_rates(X_sample, y_sample)

# Visualize convergence
print("\nVisualizing convergence...")
viz.visualize_convergence(X_sample, y_sample)

# Stability analysis
print("\nAnalyzing stability...")
stability_embeddings = viz.stability_analysis(X_sample[:200], y_sample[:200])

Practice Exercises

Exercise 1: Interactive t-SNE Explorer

Build an interactive t-SNE visualization tool:

  1. Create slider widgets for parameters
  2. Real-time parameter updates
  3. Show KL divergence convergence
  4. Compare multiple perplexities side-by-side
  5. Export results and parameters

Exercise 2: t-SNE Quality Metrics

Develop metrics to evaluate t-SNE quality:

  1. Implement neighborhood preservation metric
  2. Calculate trustworthiness and continuity
  3. Measure cluster separation
  4. Compare with other DR methods
  5. Create quality report

Exercise 3: Parametric t-SNE

Implement parametric t-SNE using neural networks:

  1. Build neural network architecture
  2. Implement t-SNE loss function
  3. Train on sample dataset
  4. Apply to new data points
  5. Compare with standard t-SNE

Summary and Key Takeaways

šŸŽÆ Key Points to Remember

šŸ““ Learning Journal

Keep a learning journal — digital or physical. After this lesson, take a few minutes to write down:

  • Key concepts you learned
  • Techniques that clicked for you
  • Questions or confusion points to revisit
  • Ideas you want to try
  • Your progress and feelings about learning this

āœļø This lesson's prompt: A t-SNE map can look convincing but can also mislead. What did changing the perplexity teach you about how much to trust the clusters you see?

šŸ“ Lesson Summary

šŸŽ“ Key Takeaways

  • t-SNE preserves local structure, making clusters pop out in a 2D or 3D picture.
  • Perplexity, learning rate, and iteration count strongly shape the result — always try several settings.
  • Distances between clusters and the sizes of clusters in a t-SNE plot are not meaningful.
  • t-SNE is for visualization only: it has no reliable transform for new data and is computationally heavy.

šŸŽ‰ What You've Accomplished

You can now turn thousands of dimensions into an intuitive 2D picture that reveals hidden groupings — a powerful way to build intuition about a dataset before modeling it.

ā“ Common Questions at This Stage

What perplexity value should I use?

Typically somewhere between 5 and 50. Think of it loosely as the number of neighbors each point considers, and try a few values to see which reveals structure most clearly.

Can I use t-SNE outputs as features for a model?

Not reliably — t-SNE is a visualization technique with no stable transform for unseen data. Use PCA or UMAP, or supervised features, when you need embeddings for modeling.

Why do my results change on every run?

t-SNE is stochastic. Set random_state for reproducibility, and avoid reading too much into any single run of the algorithm.

šŸ”­ Looking Ahead

Next you'll meet UMAP, which often produces similar visualizations much faster, preserves more global structure, and can even transform new data.

āœ… Before the Next Lesson

🌟 Encouragement for the Journey

You can now peer into high-dimensional data and see its shape — an almost magical way to understand your datasets. Enjoy the view, and keep experimenting!