title: Simple Neural Network Visualizer
short_description: interactive 3D model nntraining on MNIST
emoji: π§
colorFrom: indigo
colorTo: yellow
sdk: static
pinned: false
app_file: index.html
Future Enhancements
- Real-time training animation via epochs -> change in colors.
- Layer-specific zoom focus
- Neuron activation heatmap
- Weight magnitude visualization
- input / output bars
Custom Neural Network Visualizer
A 3D interactive visualization of a neural network trained on MNIST handwritten digit recognition. This project combines a custom-built neural network implementation with Three.js to create an immersive visualization of network architecture, weights, and connections.
Purpose: Visualization of actual model architecture is often lacking in machine learning. This project demonstrates what real-time neural network training looks like in 3D space, with dynamic weight updates and color-coded connections.
Stopped now: Entire simple model of sigmoid neural networks is in js. turns out people don't visualize this because it's a pain in the ass. so much pain.
Features
- Custom Neural Network: Fully implemented backpropagation neural network from scratch in pure JavaScript (no ML frameworks)
- MNIST Training: Load and preprocess MNIST dataset (55,000 training samples, 10,000 test samples)
- 3D Visualization: Interactive Three.js scene showing network architecture in 3D space
- Input Layer Grid: Input layer (784 neurons) displayed as a 28Γ28 grid (matching MNIST image dimensions)
- Dynamic Weight Visualization: Connections colored by weight sign and magnitude
- Green: Positive weights (excitatory connections)
- Red: Negative weights (inhibitory connections)
- White: Neutral/weak weights
- Real-time Training Updates: Weights update color dynamically as network trains
- Smart Connection Management: Weak weights below threshold are automatically hidden; strong weights are shown/created dynamically
- Interactive HUD: Top bar displays:
- Epoch counter (starts at 0)
- Mini-batch progress counter
- Real-time accuracy percentage
- Training control buttons
- Interactive Camera: OrbitControls for smooth 3D navigation with reset button
- Async Training: Non-blocking training with UI updates every 10 mini-batches
Project Structure
βββ CustomNeuralNetwork.js # Neural network implementation with backpropagation
βββ MNIST_dataset.js # MNIST data loading from Google Cloud
βββ preprocessing.js # Data preprocessing and reshaping
βββ NetworkToDisplay.js # Converts network to 3D scene configuration
βββ SceneManager.js # Three.js scene, camera, and lighting setup
βββ SphereManager.js # Creates and manages neuron spheres
βββ ConnectionManager.js # Creates and manages weight connections
βββ main.js # Entry point and training orchestration
βββ index.html # HTML page with canvas
βββ README.md # This file
Architecture
Neural Network (CustomNeuralNetwork.js)
Custom implementation featuring:
- Matrix operations: Dot product, element-wise operations, transpose
- Activation functions: Sigmoid activation with derivatives
- Training: Stochastic Gradient Descent (SGD) with mini-batches
- Backpropagation: Full backpropagation algorithm for weight updates
- Evaluation: Accuracy testing on test data
// Example usage
const network = new Network([784, 16, 16, 10]); // Input, hidden, hidden, output
network.SGD(trainingData, epochs, batchSize, learningRate, testData);
Data Pipeline (preprocessing.js)
Loads MNIST data and formats it for the network:
- Fetches image sprites and labels from Google Cloud Storage
- Reshapes flat 784-element arrays into column vectors
[[val], [val], ...] - Prepares one-hot encoded labels
- Returns
{trainingData, testData}ready for training
Visualization (NetworkToDisplay.js)
Instance-based class that creates and updates the 3D scene directly:
- Neurons: Represented as spheres with layer-specific colors
- Red: Input layer (28Γ28 grid)
- Yellow: Hidden layers (vertical line)
- Cyan: Output layer (vertical line for 10 classes)
- Connections: Dynamically managed based on weight magnitude
- Color coding by weight sign:
- Green: Positive weights (excitatory)
- Red: Negative weights (inhibitory)
- White: Neutral weights
- Smart visibility: Connections below
weightThreshold(default 0.1) are automatically hidden - Dynamic updates:
updateConnections()creates new strong connections and removes weak ones in real-time
- Color coding by weight sign:
3D Scene Management (SceneManager.js)
Provides Three.js scene setup with:
- Perspective camera with orbit controls
- Ambient and directional lighting with atmospheric point lights
- Fog effect for depth perception
- Real-time rendering loop with connection position updates
- Window resize handling
- Camera reset functionality (removes momentum/velocity)
Usage
Basic Setup
- Open
index.htmlin a modern web browser - Wait for MNIST dataset to load (automatic)
- The initial network architecture [784, 16, 16, 10] will be visualized immediately
- Click "Next Epoch" button to train for one epoch
- Watch as:
- Mini-batch counter updates in real-time
- Connections change color based on weight updates
- Accuracy percentage updates after each epoch
- Epoch counter increments
Interactive Controls
Mouse Controls:
- Left-click + drag: Rotate view
- Right-click + drag: Pan camera
- Mouse scroll: Zoom in/out
HUD Controls:
- Next Epoch button: Train network for one epoch (async, non-blocking)
- Reset Camera button: Return camera to initial position and stop momentum
Custom Network Configuration
Edit main.js to modify network architecture and training parameters:
// Change network architecture
const network = new Network([784, 32, 32, 16, 10]);
// Adjust visualization parameters
const visualizer = new NetworkToDisplay(network, sceneManager, {
layerSpacing: 10,
neuronSpacing: 2,
neuronRadius: 0.5,
inputColor: '#ff6b6b',
hiddenColor: '#ffd93d',
outputColor: '#4ecdc4',
weightThreshold: 0.1, // Hide connections below this weight magnitude
lineWidth: 0.1,
opacity: 0.1
});
// Modify training parameters in the button click handler
network.SGD_single_epoch(trainingData, miniBatchSize=100, learningRate=0.5, testData, visualizer_function);
Data Format
Training Data
Each sample is a pair [x, y]:
x: Column vector of 784 pixel values (normalized 0-1)[[0.5], [0.3], [0.8], ...] // 784 elementsy: One-hot encoded label (10 elements for digits 0-9)[[0], [0], [1], [0], ...] // For digit 2
Network Weights
weights[i]: 2D array from layer i to layer i+1- Dimensions:
[numNeuronsInLayerI+1, numNeuronsInLayerI] weights[i][j][k]: Weight from neuron k in layer i to neuron j in layer i+1
- Dimensions:
Class Reference
Network
// Constructor
new Network(neuronCounts) // e.g., [784, 16, 16, 10]
// Methods
feedforward(inputVector) // Get network prediction
SGD_single_epoch(trainingData, miniBatchSize, lr, testData, cb) // Train one epoch (async)
evaluate(testData) // Count correct predictions
backpropagation(x, y) // Compute gradients
SphereManager
addSphere(name, config) // Add neuron sphere
removeSphere(name) // Remove sphere
updateSphere(name, config) // Update sphere properties
getSphere(name) // Get sphere mesh
changeColorSmooth(name, color) // Animate color change
ConnectionManager
addConnection(id, fromPos, toPos, config) // Add weight connection
removeConnection(id) // Remove connection
updateConnectionPositions(id, from, to) // Update endpoints
NetworkToDisplay
// Constructor (instance-based, not static)
new NetworkToDisplay(network, sceneManager, options)
// Methods
initialGeneration() // Create all spheres and initial connections in scene
updateConnections() // Update/create/delete connections based on current weights
highlightNeuron(layer, neuron, activation) // Highlight active neuron (instance method)
Static Methods:
NetworkToDisplay.weightToColor(weight) // Convert weight value to color (red/white/green gradient)
Performance Considerations
- Input layer grid: 28Γ28 = 784 neurons (no sampling)
- Dynamic connections: Only connections with weights above threshold are rendered
- Initial: ~12,960 connections (784β16 + 16β16 + 16β10)
- Runtime: Varies based on weight magnitudes (weak connections hidden)
- Async training: Mini-batches update UI every 10 iterations to prevent blocking
- Color updates: Real-time weight-to-color conversion with interpolation
- Rendering: Optimized for dynamic object creation/deletion with orbit controls
Dependencies
- Three.js r128: 3D graphics library (loaded from CDN)
- OrbitControls: Camera control extension for Three.js (loaded from CDN)
- TensorFlow.js 1.0.0: Used only for MNIST data loading from Google Cloud Storage
All dependencies are loaded via CDN - no installation required.
Future Enhancements
- Layer-specific zoom focus
- Neuron activation heatmap (color neurons by activation value during inference)
- Input/output visualization bars showing sample images and predictions
- Weight magnitude histogram
- Training loss curve overlay
- Export/import trained model weights
- Multiple activation functions (ReLU, Tanh, etc.)
- Adjustable learning rate during training
References
- MNIST Dataset: Yann LeCun's MNIST Database
- Three.js: Three.js Documentation
- Neural Networks: Neural Networks and Deep Learning
License
This is an educational project for learning machine learning and 3D visualization.
Author
Created as part of a custom neural network learning project.
