# Why do some Flux models train in parallel but not others?

**URL:** <https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453>\
**Category:** Machine Learning\
**Created:** [August 8, 2022, 1:33am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453 "2022-08-08T01:33:36Z")\
**Posts on this page:** 17\
**Page:** 1

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 8, 2022, 1:33am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/1 "2022-08-08T01:33:36Z")

</div>

Why does training this model use as many CPUs as I ask it to:

```julia
Chain(
             Conv((5, 5), imgsize[end]=>6, relu),
             MaxPool((2, 2)),
             Conv((5, 5), 6=>16, relu),
             MaxPool((2, 2)),
             flatten,
             Dense(prod(out_conv_size), 120, relu), 
             Dense(120, 84, relu), 
             Dense(84, nclasses)
)

```

And training this model use only 1?

```julia
Chain(
        #28x28 to 14x14
        Conv((5,5), 1=>8, pad = 2, stride = 2, relu),
        #14x14 to 7x7
        Conv((3,3), 8=>16, pad = 1, stride = 2, relu),
        #7x7 to 4x4
        Conv((3,3), 16=>32, pad = 1, stride = 2, relu),
    
        #Average pooling on each width x height feature map
        
        GlobalMeanPool(),
        Flux.flatten,
        Dense(32,10),
        softmax
    )

```

I paste both into the MNIST example from the Flux model zoo so all else should be equal.

I’m sure there is a reason; I’m not sure what it is

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [August 8, 2022, 5:08am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/2 "2022-08-08T05:08:49Z")

</div>

I gues this is because `Dense` layers hit Blas’ MatMul, which is multi-threaded whereas `Conv` is likely single-threaded.

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [August 8, 2022, 2:39pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/3 "2022-08-08T14:39:15Z")

</div>

Unless you are getting warnings about incompatible types from NNlib, convs should absolutely be multi-threaded. What does run single-threaded are (most) pooling operations, which the first model uses extensively and the second does not.

---

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 8, 2022, 2:51pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/4 "2022-08-08T14:51:37Z")

</div>

I’m not getting any warnings. What am I doing wrong that Conv layers are single-threaded?

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [August 8, 2022, 4:14pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/5 "2022-08-08T16:14:00Z")

</div>

Again, the conv layers are not. They likely just run fast enough that you only see the single threaded max-pooling layers (and/or other single-threaded parts of the training loop, like data loading) in whatever monitoring tool you’re using. You can confirm for yourself that the conv layers are indeed using multiple threads by removing said pooling layers and benchmarking on a fixed dummy input to remove any data loading overhead.

---

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 9, 2022, 1:33am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/6 "2022-08-09T01:33:11Z")

</div>

So if pooling is single-threaded and thereby masks the multi-threaded nature of the Conv layers, why does the model with more pooling layers use 8 cores most of the time and the one with only one pooling layer use 1 core almost all of the time?

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [August 9, 2022, 4:35am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/7 "2022-08-09T04:35:05Z")

</div>

After actually running these models locally, turns out it was the simplest answer and I was barking up the wrong tree 🙂

By default, Julia allocates but a single thread to the default thread pool (you can check with `Threads.nthreads()`). Because Flux conv layers use this thread pool, they end up running (mostly, more on that below) single-threaded. To make Julia use multiple threads, either pass `-t [nthreads]` or `-t auto` at startup. If you’re using VS Code, this is also exposed via the “Julia: Num Threads” option.

Now if conv layers are running single threaded, why does `model1` appear to use multiple? That’s because the matrix multiplication calls in `Dense` layers use a separate, BLAS threadpool which is \>1 by default (you can check this with `using LinearAlgebra; BLAS.get_num_threads()`). Because `model1` has two large dense layers to `model2`’s one small one, it spends a lot more time here and thus a lot more time in multi-threaded code. Conv layers also use matmuls under the hood, but these are generally smaller and need the aforementioned default thread pool for any significant parallelism.

---

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 9, 2022, 4:54am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/8 "2022-08-09T04:54:55Z")

</div>

Now I can shave a whole 8% of training time by running 8 threads.

The training of Conv layers is acting pretty memory bound, so having lots of threads seems pretty pointless. Empirically it mostly makes training slower

---

<div class="post-metadata">

**Author:** ![j\_u](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/j_u/32/219081_2.png) [@j\_u](https://discourse.julialang.org/u/j_u)\
**Post date:** [August 9, 2022, 11:30pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/9 "2022-08-09T23:30:11Z")

</div>

> [@Lewis\_Hein](#):
>
> Now I can shave a whole 8% of training time by running 8 threads.

Hi, out of curiosity, am I right that you have 12 cores on your machine? Are they real or Hyper Threaded? One or two sockets? Are you happy with this 8% increase if I may ask? Just wanted to mention `ThreadPinning.jl` package as well as `STREAMBenchmark.jl` and `BandwidthBenchmark.jl`. I believe they give additional insights, particularly useful for ML applications.

---

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 10, 2022, 4:11am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/10 "2022-08-10T04:11:40Z")

</div>

I should have been more specific. There are 12 real cores with hyperthreading, so 24 virtual cores.

I have never found much performance gain using hyperthreading in memory-bound applications so I seldom exceed the number of real cores

---

<div class="post-metadata">

**Author:** ![j\_u](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/j_u/32/219081_2.png) [@j\_u](https://discourse.julialang.org/u/j_u)\
**Post date:** [August 10, 2022, 10:45am UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/11 "2022-08-10T10:45:29Z")

</div>

My point here was rather related to dual socket systems and the way threads are affinitized by the OS or Julia. You are not providing any details about the architecture, the way how Julia was set and I have not run those examples by myself so its quite hard to refer. As for HT, AFAIK, you are probably very right; in case of dense computations its particularly visible due to the limited number of CPU vector units.

---

<div class="post-metadata">

**Author:** ![Lewis\_Hein](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lewis_hein/32/23540_2.png) [@Lewis\_Hein](https://discourse.julialang.org/u/Lewis_Hein)\
**Post date:** [August 10, 2022, 1:41pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/12 "2022-08-10T13:41:15Z")

</div>

I’m fine with the 8% gain for now; I will likely want more in the future. But for now I have the main thing I wanted, which is the ability to train with more than 1 thread if I want to.

If anyone cares, the architecture is an old Dell PowerEdge R610 with dual Intel Xeon E5-2630s. that have 6 threads/CPU plus hyperthreading. It has been a while since I checked, but I think the motherboard has 4 RAM channels per CPU, hence my choice of 8 threads

And yes, I know this hardware is old and slow by modern standards. I don’t make money with my machine learning / scientific computing/ data analysis projects so fancy new hardware is hard to justify over $32 servers

---

<div class="post-metadata">

**Author:** ![j\_u](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/j_u/32/219081_2.png) [@j\_u](https://discourse.julialang.org/u/j_u)\
**Post date:** [August 10, 2022, 2:07pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/13 "2022-08-10T14:07:25Z")

</div>

I use different computers as well, with my main being released around the same time as yours. In general, I was referring to the architecture and Julia setup and some of the packages I found interesting. I believe the setup in some cases may provide additional benefits / speedups. Tried to share some of my own experiences with Julia, BLAS and ML/AI models.

---

<div class="post-metadata">

**Author:** ![Rasmus\_Hoier](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rasmus_hoier/32/24036_2.png) [@Rasmus\_Hoier](https://discourse.julialang.org/u/Rasmus_Hoier)\
**Post date:** [August 10, 2022, 3:46pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/14 "2022-08-10T15:46:18Z")

</div>

[SimpleChains.jl](https://github.com/PumasAI/SimpleChains.jl) might be useful if you are limited to CPU. The authors show some benchmarks in [this blog post](https://julialang.org/blog/2022/04/simple-chains/#simplechainsjl_in_action_5x-ing_pytorch_in_small_examples).

---

<div class="post-metadata">

**Author:** ![j\_u](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/j_u/32/219081_2.png) [@j\_u](https://discourse.julialang.org/u/j_u)\
**Post date:** [August 10, 2022, 10:38pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/15 "2022-08-10T22:38:34Z")

</div>

Interesting article about `SimpleChains.jl`. I was not aware about it. Thanks. I am wondering: a) Do you think that such or similar techniques could be used for networks like `AlphaZero.jl`? and b) As for `PyTorch` comparison, was it with `IPEX` / `oneDNN` or it is not applicable at all in this case?

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [August 10, 2022, 10:40pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/16 "2022-08-10T22:40:46Z")

</div>

NNlib is definitely oversubscribing threads. I’m not sure how much it has an impact in your case because the default im2col algorithm is memory-intensive and GC isn’t great at keeping up with heavily allocating multi-threaded code, but [https://github.com/FluxML/NNlib.jl/pull/395](https://github.com/FluxML/NNlib.jl/pull/395) suggests some impact.

As noted in that PR, the biggest blocker for toying around with threading in NNlib is proper benchmarking code + infrastructure. If anyone is interested in that, please reach out! Until then, I second the suggestion to check out SimpleChains if your model and inputs are sufficiently small (e.g. MNIST-sized).

---

<div class="post-metadata">

**Author:** ![j\_u](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/j_u/32/219081_2.png) [@j\_u](https://discourse.julialang.org/u/j_u)\
**Post date:** [August 10, 2022, 11:50pm UTC](https://discourse.julialang.org/t/why-do-some-flux-models-train-in-parallel-but-not-others/85453/17 "2022-08-10T23:50:55Z")

</div>

I did some additional reading. Sorry about my previous questions … too focused on one area.
