File size: 10,615 Bytes
28b08b8 a043ef3 28b08b8 b012e8c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 | ---
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.
demo:

## 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
```javascript
// 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:
1. Fetches image sprites and labels from Google Cloud Storage
2. Reshapes flat 784-element arrays into column vectors `[[val], [val], ...]`
3. Prepares one-hot encoded labels
4. 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
### 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
1. Open `index.html` in a modern web browser
2. Wait for MNIST dataset to load (automatic)
3. The initial network architecture [784, 16, 16, 10] will be visualized immediately
4. Click **"Next Epoch"** button to train for one epoch
5. 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:
```javascript
// 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 elements
```
- `y`: 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
## Class Reference
### Network
```javascript
// 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
```javascript
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
```javascript
addConnection(id, fromPos, toPos, config) // Add weight connection
removeConnection(id) // Remove connection
updateConnectionPositions(id, from, to) // Update endpoints
```
### NetworkToDisplay
```javascript
// 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:**
```javascript
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](http://yann.lecun.com/exdb/mnist/)
- Three.js: [Three.js Documentation](https://threejs.org/docs/)
- Neural Networks: [Neural Networks and Deep Learning](http://neuralnetworksanddeeplearning.com/)
## License
This is an educational project for learning machine learning and 3D visualization.
## Author
Created as part of a custom neural network learning project.
|