File size: 33,187 Bytes
a1bc931
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
663c79b
a1bc931
 
 
663c79b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a1bc931
 
 
663c79b
a1bc931
 
663c79b
a1bc931
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
663c79b
 
 
 
 
 
 
a1bc931
 
 
 
 
663c79b
 
 
 
 
 
 
a1bc931
 
 
 
 
663c79b
 
 
 
 
a1bc931
663c79b
a1bc931
 
 
 
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
import gradio as gr                                                                                                                                                                                                  
import numpy as np                                                                                                                                                                                                   
from PIL import Image                                                                                                                                                                                                
                                                                                                                                                                                                                     
# Import our custom layers and activations                                                                                                                                                                           
from layers import Conv, MaxPool, Flatten, Dense                                                                                                                                                                     
from activations import GELU                                                                                                                                                                                         
                                                                                                                                                                                                                     
# 1. Initialize the network structure                                                                                                                                                                                
conv = Conv(input_shape=(1, 28, 28), kernel_size=5, num_kernels=12)                                                                                                                                                  
gelu = GELU()                                                                                                                                                                                                        
pool = MaxPool(2, 2)                                                                                                                                                                                                 
flatten = Flatten()                                                                                                                                                                                                  
dense = Dense(1728, 10)                                                                                                                                                                                              
                                                                                                                                                                                                                     
# 2. Try loading the weights safely                                                                                                                                                                                  
weights_loaded = False                                                                                                                                                                                               
try:                                                                                                                                                                                                                 
    data = np.load("model_weights.npz")                                                                                                                                                                              
    conv.kernels = data["conv_kernels"]                                                                                                                                                                              
    conv.biases = data["conv_biases"]                                                                                                                                                                                
    dense.weights = data["dense_weights"]                                                                                                                                                                            
    dense.biases = data["dense_biases"]                                                                                                                                                                              
    weights_loaded = True                                                                                                                                                                                            
    print("Weights loaded successfully!")                                                                                                                                                                            
except FileNotFoundError:                                                                                                                                                                                            
    print("Warning: model_weights.npz not found.")                                                                                                                                                                   
except Exception as e:                                                                                                                                                                                               
    print(f"Warning: error loading weights: {e}")                                                                                                                                                                    
                                                                                                                                                                                                                     
# Helper to stitch feature maps into a nice grid                                                                                                                                                                     
def make_grid(feature_maps, cols=4):                                                                                                                                                                                 
    n, h, w = feature_maps.shape                                                                                                                                                                                     
    rows = (n + cols - 1) // cols                                                                                                                                                                                    
                                                                                                                                                                                                                     
    # Pad borders between filters so they look separate                                                                                                                                                              
    padding = 2                                                                                                                                                                                                      
    grid_h = rows * h + (rows - 1) * padding                                                                                                                                                                         
    grid_w = cols * w + (cols - 1) * padding                                                                                                                                                                         
    grid = np.zeros((grid_h, grid_w), dtype=np.uint8)                                                                                                                                                                
                                                                                                                                                                                                                     
    for idx in range(n):                                                                                                                                                                                             
        r = idx // cols                                                                                                                                                                                              
        c = idx % cols                                                                                                                                                                                               
                                                                                                                                                                                                                     
        f_map = feature_maps[idx]                                                                                                                                                                                    
        f_min, f_max = f_map.min(), f_map.max()                                                                                                                                                                      
        # Normalize to [0, 255] for image display                                                                                                                                                                    
        if f_max > f_min:                                                                                                                                                                                            
            f_map = 255.0 * (f_map - f_min) / (f_max - f_min)                                                                                                                                                        
        else:                                                                                                                                                                                                        
            f_map = f_map * 0                                                                                                                                                                                        
                                                                                                                                                                                                                     
        f_map = f_map.astype(np.uint8)                                                                                                                                                                               
                                                                                                                                                                                                                     
        y_start = r * (h + padding)                                                                                                                                                                                  
        x_start = c * (w + padding)                                                                                                                                                                                  
        grid[y_start:y_start+h, x_start:x_start+w] = f_map                                                                                                                                                           
                                                                                                                                                                                                                     
    img = Image.fromarray(grid)                                                                                                                                                                                      
    # Scale up using nearest-neighbor to keep pixels clean and sharp                                                                                                                                                 
    img = img.resize((grid_w * 12, grid_h * 12), Image.Resampling.NEAREST)                                                                                                                                           
    return img                                                                                                                                                                                                       
                                                                                                                                                                                                                     
# 3. Predict function                                                                                                                                                                                                
def predict(input_image):                                                                                                                                                                                            
    if not weights_loaded:                                                                                                                                                                                           
        return {"Error: please upload 'model_weights.npz'": 1.0}, None, None                                                                                                                                         
                                                                                                                                                                                                                     
    if input_image is None:                                                                                                                                                                                          
        return "No image drawn", None, None                                                                                                                                                                          
                                                                                                                                                                                                                     
    if isinstance(input_image, dict):                                                                                                                                                                                
        img = input_image['composite']                                                                                                                                                                               
    else:                                                                                                                                                                                                            
        img = input_image                                                                                                                                                                                            
                                                                                                                                                                                                                     
    # Resize and convert to grayscale                                                                                                                                                                                
    img = img.convert('L').resize((100, 100))                                                                                                                                                                        
    arr = np.array(img)                                                                                                                                                                                              
                                                                                                                                                                                                                     
    # Invert background to match MNIST (white text on black background)                                                                                                                                              
    if arr[0, 0] > 128:                                                                                                                                                                                              
        bg_noise = max(arr[0, 0], arr[-1, -1], arr[0, -1], arr[-1, 0])                                                                                                                                               
        arr[arr > bg_noise - 10] = bg_noise                                                                                                                                                                          
        arr = bg_noise - arr                                                                                                                                                                                         
    else:                                                                                                                                                                                                            
        bg_noise = min(arr[0, 0], arr[-1, -1], arr[0, -1], arr[-1, 0])                                                                                                                                               
        arr[arr < bg_noise + 10] = bg_noise                                                                                                                                                                          
        arr = arr - bg_noise                                                                                                                                                                                         
                                                                                                                                                                                                                     
    # Center and pad the digit exactly like MNIST processing                                                                                                                                                         
    non_zero = np.argwhere(arr > 35)                                                                                                                                                                                 
    if len(non_zero) > 0:                                                                                                                                                                                            
        min_y, min_x = non_zero.min(axis=0)                                                                                                                                                                          
        max_y, max_x = non_zero.max(axis=0)                                                                                                                                                                          
        cropped = arr[min_y:max_y+1, min_x:max_x+1]                                                                                                                                                                  
        cropped[cropped < 45] = 0                                                                                                                                                                                    
                                                                                                                                                                                                                     
        h, w = cropped.shape                                                                                                                                                                                         
        cropped_img = Image.fromarray(cropped)                                                                                                                                                                       
        if h > w:                                                                                                                                                                                                    
            new_h = 20                                                                                                                                                                                               
            new_w = int(20 * w / h)                                                                                                                                                                                  
        else:                                                                                                                                                                                                        
            new_w = 20                                                                                                                                                                                               
            new_h = int(20 * h / w)                                                                                                                                                                                  
        new_w = max(1, new_w)                                                                                                                                                                                        
        new_h = max(1, new_h)                                                                                                                                                                                        
                                                                                                                                                                                                                     
        resized = cropped_img.resize((new_w, new_h), Image.Resampling.LANCZOS)                                                                                                                                       
                                                                                                                                                                                                                     
        canvas = Image.new('L', (28, 28), 0)                                                                                                                                                                         
        offset_x = (28 - new_w) // 2                                                                                                                                                                                 
        offset_y = (28 - new_h) // 2                                                                                                                                                                                 
        canvas.paste(resized, (offset_x, offset_y))                                                                                                                                                                  
        arr = np.array(canvas)                                                                                                                                                                                       
                                                                                                                                                                                                                     
    x = (arr / 255.0) - 0.5                                                                                                                                                                                          
    x = x[np.newaxis, :, :] # (1, 28, 28)                                                                                                                                                                            
                                                                                                                                                                                                                     
    # Forward pass and record intermediate activations                                                                                                                                                               
    out_conv = conv.forward(x)                                                                                                                                                                                       
    out_gelu = gelu.forward(out_conv)                                                                                                                                                                                
    out_pool = pool.forward(out_gelu)                                                                                                                                                                                
                                                                                                                                                                                                                     
    out_flat = flatten.forward(out_pool[np.newaxis, :, :, :])                                                                                                                                                        
    logits = dense.forward(out_flat)[0]                                                                                                                                                                              
                                                                                                                                                                                                                     
    # Softmax probabilities                                                                                                                                                                                          
    probs = np.exp(logits - np.max(logits))                                                                                                                                                                          
    probs /= np.sum(probs)                                                                                                                                                                                           
                                                                                                                                                                                                                     
    # Generate feature map grids for visualization                                                                                                                                                                   
    conv_grid = make_grid(out_conv)                                                                                                                                                                                  
    pool_grid = make_grid(out_pool)                                                                                                                                                                                  
                                                                                                                                                                                                                     
    # Return dictionary of classes, plus the two grids                                                                                                                                                               
    class_probs = {str(i): float(probs[i]) for i in range(10)}                                                                                                                                                       
    return class_probs, conv_grid, pool_grid                                                                                                                                                                         
                                                                                                                                                                                                                     
# Gradio Interface layout                                                                                                                                                                                            
demo = gr.Interface(                                                                                                                                                                                                 
    fn=predict,                                                                                                                                                                                                      
    inputs=gr.Sketchpad(type="pil", image_mode="L"),                                                                                                                                                                 
    outputs=[                                                                                                                                                                                                        
        gr.Label(num_top_classes=3, label="Prediction"),                                                                                                                                                             
        gr.Image(label="Layer 1: Convolutional Activations (12 Filters)", type="pil"),                                                                                                                               
        gr.Image(label="Layer 2: Max Pooling Outputs (Downsampled Features)", type="pil")                                                                                                                            
    ],                                                                                                                                                                                                               
    title="CNN from Scratch - Digit Classifier",                                                                                                                                                                     
    description="Draw a digit in the box to predict its value and see inside the model's 'brain' in real-time!"                                                                                                      
)                                                                                                                                                                                                                    
                                                                                                                                                                                                                     
if __name__ == "__main__":                                                                                                                                                                                           
    demo.launch()