You are browsing as a guest. Sign up (or log in) to start making projects!

Neurorust

  • 12 Devlogs
  • 60 Total hours

GPU accelerated neural network from scratch in rust

Ship #1 Changes requested

Neurorust lets you train a simple neural network in your browser, using the GPU. The network predicts handwritten digits. After training you can draw your own digits and see its predictions.

This project has tought me a lot about neural networks, Rust, GPU programming and WebAssembly. It was really challenging to produce a good result when just learning about everything. Even though the code is not perfect and there are a lot of things I would do different next time I feel like the end result has turned out quite well.

  • 12 devlogs
  • 60h
Try project → See source code →
Open comments for this post

4h 8m 36s logged

Done

I added a README, improved the website defaults and did some final debugging. Now the project is ready to ship.

0
0
16
Open comments for this post

8h 38m 52s logged

Website finished

I finished the frontend. You can now create a network with specific hyperparameters and then train it. After training you can draw a number and see what it predicts.

Now I just need a good README and the project is ready to ship.

0
0
11
Open comments for this post

6h 9m 53s logged

Does it ever get easier?

I thought using the network in the web would be a piece of cake. I was wrong.

After fixing some minor issues I wrote some quick code to test training. It worked… technically. It was really slow and froze my whole pc. A quick glance into my system monitor revealed the secret. All CPU cores were almost maxed out while the GPU sat idle. So after spending a bunch of time to use the GPU everything was back to the CPU again.

WebGPU is still experimental. Especially on Linux which is why instead of throwing an error when it couldn’t access my graphics card it just used all my CPU cores as a fallback. To get the GPU back into the game I started my browser with some flags and compatibility layers which makes training still slower than on native hardware but faster than with the CPU.

Hopefully this was the last major challenge and implementing the frontend will be easier.

0
0
27
Open comments for this post

6h 50m 43s logged

It acutally learns (again)

After some more pain with the loss function, buffer sizes and shape mismatches I got a neural network running on the gpu. The GPU is about 10x faster than the CPU version although some hyperparameters can affect this result.

Now I want to build a small wasm interface and build a small web page around that to allow anyone to run a local neural net in the browser.

1
0
117
Open comments for this post

4h 21m 57s logged

GPU Layers

I initially thought porting the forward and backward passes to the GPU would be as simple as replacing Matrix with GpuMatrix everywhere. It wasn’t.

Unlike CPU operations, GPU operations can’t allocate memory themselves. On the CPU, x + 2 simply creates a new value. On the GPU, every operation needs a pre-existing buffer to write its result into. So x += 2 works, but y = x + 2 requires a buffer for y to exist first. And x = x + 2 isn’t supported on the GPU at all.

The solution is to pre-allocate buffers and keep them around, resizing them only when the batch size changes. It works, but it also means every struct starts accumulating buffer fields and scaling that approach to an entire network only makes the problem worse.

Next: some kind of global buffer manager, before tackling the loss function.

0
0
8
Open comments for this post

4h 40m 19s logged

Forward pass

After adding all necessary operations on the GpuMatrix I implemented the forward pass. It was surprisingly hard to get the caching for the backward pass right because reallocating buffers on the gpu should be avoided. As always I let AI write some test cases and everything passed.

Next I’ll implement the backward pass which will hopefully let the network learn with incredible speeds.

0
0
11
Open comments for this post

4h 39m 6s logged

GPU Matmul

I thought beating the CPU at matrix multiplication would be easy. It turns out that hand-optimized assembly code is a tough opponent.

My naive GPU version isn’t bad, though: it’s roughly on par with NumPy, and about 13x faster than my CPU code but still worse than I initially expected. There’s plenty of room to improve. Right now the GPU sits idle about 80% of the time, just waiting on memory reads.

For now, I’ll keep building out the rest of the neural net on the GPU and come back to optimizing matmul later.

0
0
5
Open comments for this post

1h 46m 39s logged

Sidequest: Doubling numbers

In my last devlog I stated that doubling numbers on the GPU is only faster if you want to double a HUUUGE amount of numbers. Of course I wanted to see at what point the GPU really is faster. I found out that the CPU is actually quite good at doubling numbers so the GPU wasn’t faster at all. I then measured just the time the GPU takes to double the numbers ignoring the time it needs to transfer the data. With that the GPU eventually overtakes the CPU.

0
0
23
Open comments for this post

2h 4m 55s logged

Doubling numbers

Turns out you have to go back to square one to run a neural net on the GPU. For me that meant writing a program that multiplies every number in an array by 2, on the GPU.

Limited learning resources

Rust has a great library called wgpu for talking to the GPU. It’s cross-platform and can even run in WebAssembly if you do it right. The problem: most tutorials only cover rendering, and my neural net doesn’t care about triangles on a screen, it needs compute. I never found a tutorial on compute shaders from scratch, so I read the rendering material, stripped out what didn’t apply (windows, vertices, …), and added what did (workgroup size, how a compute pass works, …). Not trivial, but the official wgpu compute example got me on track. I mostly just copied it.

Why it’s hard

GPU programming isn’t straightforward. Here’s everything required just to multiply an array by 2:

  1. Get a wgpu instance, adapter, device, and queue
  2. Create a shader module
  3. Create the input, output, and download buffers
  4. Create the bind group layout and bind group
  5. Create the pipeline layout
  6. Create an encoder and a compute pass
  7. Copy data CPU → GPU, run the shader, then copy data GPU → CPU

After 190 lines and some unexpected debugging (on code copied straight from an example, no less), I could finally double numbers fast, but only with huge arrays. Below a certain size, shuttling data to the GPU and back is slower than just doing the math on the CPU.

I don’t fully understand every step yet. Hopefully I won’t need to, or I’ll pick it up along the way.

Up next

Matrix multiplication is next, and most of the boilerplate above should carry over. Big questions remain about the final network, but taking it one step at a time should get me there eventually… or to a complete surrender, who knows.

0
0
8
Open comments for this post

5h 28m 30s logged

It actually learns

After getting XOR working, I tried my neural net on MNIST (handwritten digit recognition). Loading the data was easier than expected, though getting it into the right shape for my network took some fiddling. To make sure I loaded it correctly, I drew a few of the digits to check they looked right before training.

Once that was sorted, training worked well, and the results were better than I hoped: 97.8% accuracy, in about 5 minutes.

I also experimented with running the whole thing in a browser, but that turned out to be more involved than expected, so I’m parking it for now and focusing on other improvements first.

0
0
27
Open comments for this post

7h 22m 38s logged

It learns (a bit)

After implementing loss functions, activation functions, and the forward and backward pass for the layers, I tried training a simple network to predict the XOR function.

After running the network, I found that it didn’t learn at all. The loss just stayed at 0.25 without any changes. Not only that, the whole neural network didn’t update anything.

After some initial digging, I had a feeling this was one of those dumb, simple bugs you can’t find for hours. Unfortunately, my feeling was right: the error was multiplying the bias instead of adding it. This is really bad because the bias initially has a value of 0, which means all outputs turn out to be 0 (not good).

Of course, the fix was a one character change: * to +.

Now that the network can predict XOR, I’ll see if it can also handle MNIST. I’m still a bit scared about performance issues with matrix multiplication.

0
0
6
Open comments for this post

3h 51m 45s logged

Rust Learning Project: Building a Neural Network

I recently read a few chapters of The Rust Programming Language and wanted to put my newly gained knowledge to use. Naturally, the first thing that came to mind was building a neural network in Rust. I’d already tried this in Python and failed, so I figured maybe, for some reason, it would go better this time around.

Matmul

The first step was implementing matrices and their operations in Rust, the most important of course being matrix multiplication. Since this would likely become the biggest bottleneck for the whole model, I got curious about its performance and benchmarked multiplying two 1024x1024 matrices.

The naive implementation took about 20 seconds. One simple optimization brought that down to 14 seconds.

Then I built with compiler optimizations enabled, and the naive implementation suddenly ran in 63 nanoseconds. I was floored, briefly convinced I’d stumbled onto some incredible compiler magic, before realizing what actually happened: the compiler had simply skipped the calculation entirely, since the result was never used. Classic dead code elimination.

After forcing the compiler to keep the computation, the naive implementation came in at a more believable 2 seconds, and the optimized version dropped to 200ms.

For comparison, Numpy does the same multiplication in about 4ms. Whether 200ms is fast enough for training neural nets in reasonable time remains to be seen.

0
0
5

Delete project?

Are you sure you want to permanently delete this project? This action cannot be undone.

All devlogs, followers, and associated data will be removed.

Followers

Loading…