# Flux Transformer Out of Memory

**URL:** <https://discourse.julialang.org/t/flux-transformer-out-of-memory/95922>\
**Category:** Machine Learning\
**Created:** [March 11, 2023, 4:38pm UTC](https://discourse.julialang.org/t/flux-transformer-out-of-memory/95922 "2023-03-11T16:38:55Z")\
**Posts on this page:** 1\
**Showing post:** 6

<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:** [March 12, 2023, 4:17pm UTC](https://discourse.julialang.org/t/flux-transformer-out-of-memory/95922/6 "2023-03-12T16:17:55Z")

</div>

> [@Gadersd](#):
>
> I am confused as to how Flux could have overlooked crucial functions such as a transpose that doesn’t make copies. Is there some limitation that precludes the use of functions that PyTorch users take for granted in the Flux ecosystem or has no one gotten around to implementing them?

No and no. To resolve your confusion, we should clarify _where_ `permutedims` is defined. Unlike with PyTorch where the framework defines and provides implementations for each operator, many functions you call in Flux models are defined elsewhere. In the case of `permutedims`, that elsewhere is actually the Julia standard library (Base). A good analogy would be if you could use [numpy.transpose — NumPy v1.26 Manual](https://numpy.org/doc/stable/reference/generated/numpy.transpose.html) instead of `torch.Tensor.transpose` in PyTorch and have it just work.

However, that just pushes the question upstream: why does the stdlib `permutedims` copy? I don’t know the correct answer, but I’ve asked around for a historical record of this decision and will update this thread if/when I receive one. The more relevant answer is that you can perform a non-copying dim transpose by using [`PermutedDimsArray`](https://docs.julialang.org/en/v1/base/arrays/#Base.PermutedDimsArrays.PermutedDimsArray). This is essentially what PyTorch is doing under the hood, and is also part of the Base stdlib. The biggest caveat of `PermutedDimsArray` is that not all user-defined functions may understand the wrapper and take the most [efficient](https://discourse.julialang.org/t/multiplication-after-transpose-much-faster-than-multiplication-after-permuteddimsarray/22997) [codepath](https://discourse.julialang.org/t/permuteddimsarray-slower-than-permutedims/46401/4). Relevant to this thread’s example, note how the last post on the second linked thread mentions NNlib’s [`batched_mul`](https://fluxml.ai/NNlib.jl/dev/reference/#NNlib.batched_mul) routines. What’s that in the docs page? `PermutedDimsArray`. This is precisely why I linked the definition of `dot_product_attention` in NNlib: it shows you how to use these tools to write an efficient attention operation which works much like the PyTorch one does.

All that said, I think your follow-up question is more or less addressed:

> [@Gadersd](#):
>
> It seems that Flux isn’t there yet. I commend the Flux team for what they have accomplished and I would jump at Flux the moment it ever becomes competitive with PyTorch. Is Flux on track to achieve this, or does Flux make trade-offs that will hinder it from ever having the raw performance of PyTorch?

1. Most of the “Flux” performance here is actually Base Julia performance and should be discussed accordingly.
2. Direct translations are often not apples-to-apples and knowing what the idiomatic patterns are in each language (e.g. `PermutedDimsArray` to avoid copies) can make a big difference.
3. Because DL frameworks have such large API surface areas and NN models vary greatly, “competitiveness” is always context sensitive. PyTorch definitely gets the most engineering effort towards optimizing its operations, but that [doesn’t always translate](https://julialang.org/blog/2022/04/simple-chains/) to a performance win.

---

_[View the full topic](https://discourse.julialang.org/t/flux-transformer-out-of-memory/95922)._
