# Significantly Higher VRAM Usage and Slower Training on Flux Compared to PyTorch

**URL:** <https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124>\
**Category:** Machine Learning\
**Created:** [May 14, 2026, 9:29pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124 "2026-05-14T21:29:09Z")\
**Posts on this page:** 20\
**Page:** 1

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 14, 2026, 9:29pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/1 "2026-05-14T21:29:09Z")

</div>

I’ve been working on a project where I’m converting [Timm](https://github.com/huggingface/pytorch-image-models) models to Flux along with their pre-trained weights. Everything is going well so far, but I noticed that Flux has significantly higher VRAM consumption and takes about twice as long to train compared to PyTorch.

To confirm this, I wrote a simple PyTorch and Flux script that trains a ResNet-34 model on 8,000 samples from the [AID](https://www.kaggle.com/datasets/jiayuanchengala/aid-scene-classification-datasets) dataset. To make sure the difference wasn’t down to my own mistakes, I used the ResNet implementation from [Metalhead](https://github.com/FluxML/Metalhead.jl).

Here are the respective scripts:

### PyTorch

```python
class TimmClassifier(pl.LightningModule):

    def __init__ (self, model:str, num_classes: int, learning_rate: float = 1e-3, pretrained: bool = True):
        super(). __init__ ()
        self.save_hyperparameters()
        self.model = timm.create_model(
            model, pretrained=pretrained, num_classes=num_classes
        )
        self.loss_fn = nn.CrossEntropyLoss()
        self.learning_rate = learning_rate

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.model(x)

    def _shared_step(self, batch, stage: str):
        images, labels = batch
        logits = self(images)
        loss = self.loss_fn(logits, labels)
        preds = logits.argmax(dim=1)
        acc = (preds == labels).float().mean()
        self.log(f"{stage}_loss", loss, prog_bar=True)
        self.log(f"{stage}_acc", acc, prog_bar=True)
        return loss

    def training_step(self, batch, batch_idx):
        return self._shared_step(batch, "train")

    def validation_step(self, batch, batch_idx):
        self._shared_step(batch, "val")

    def test_step(self, batch, batch_idx):
        self._shared_step(batch, "test")

    def configure_optimizers(self):
        optimizer = torch.optim.Adam(self.parameters(), lr=self.learning_rate)
        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)
        return [optimizer], [scheduler]

def train_resnet34(
    dataset,
    num_classes: int,
    val_split: float = 0.2,
    batch_size: int = 16,
    max_epochs: int = 20,
    learning_rate: float = 1e-4,
    pretrained: bool = True,
    num_workers: int = 4,
    accelerator: str = "auto",
):
    # --- Split dataset ---
    val_size = int(len(dataset) * val_split)
    train_size = len(dataset) - val_size
    train_ds, val_ds = random_split(dataset, [train_size, val_size], generator=torch.Generator().manual_seed(42))

    train_loader = DataLoader(
        train_ds, batch_size=batch_size, shuffle=True,
        num_workers=num_workers, pin_memory=True,
    )

    val_loader = DataLoader(
        val_ds, batch_size=batch_size, shuffle=False,
        num_workers=num_workers, pin_memory=True,
    )

    # --- Model ---
    model = TimmClassifier(
        model="resnet34",
        num_classes=num_classes,
        learning_rate=learning_rate,
        pretrained=pretrained,
    )

    # --- Callbacks ---
    checkpoint_cb = ModelCheckpoint(
        monitor="val_loss", mode="min", save_top_k=1, filename="best"
    )
    early_stop_cb = EarlyStopping(monitor="val_loss", patience=5, mode="min")

    # --- Trainer ---
    trainer = pl.Trainer(
        max_epochs=max_epochs,
        accelerator=accelerator,
        callbacks=[checkpoint_cb, early_stop_cb],
        log_every_n_steps=10,
        precision="32-true", 
    )
    trainer.fit(model, train_loader, val_loader)

def main():
    dataset = AID(path="../../Data/AID")
    train_resnet34(dataset, num_classes=len(dataset.classes), pretrained=False, max_epochs=5)

```

### Flux

```julia
struct FineTuneEncoder{E} <: FluxModule
    encoder::E
end

function FineTuneEncoder(num_classes::Int)
    encoder = Metalhead.ResNet(34, pretrain=false, nclasses=num_classes)
    return FineTuneEncoder(encoder)
end

Flux.Optimisers.trainable(x::FineTuneEncoder) = (; x.encoder)

(model::FineTuneEncoder)(x) = model.encoder(x)

function loss_and_accuracy(model::FineTuneEncoder, batch)
    x, y = batch
    ŷ = model(x)
    return Flux.logitcrossentropy(ŷ, y), Tsunami.accuracy(ŷ, y)
end

function Tsunami.train_step(model::FineTuneEncoder, trainer, batch)
    loss, acc = loss_and_accuracy(model, batch)
    Tsunami.log(trainer, "loss/train", loss, prog_bar=true)
    Tsunami.log(trainer, "accuracy/train", acc, prog_bar=true)
    return loss
end

function Tsunami.val_step(model::FineTuneEncoder, trainer, batch)
    loss, acc = loss_and_accuracy(model, batch)
    Tsunami.log(trainer, "loss/val", loss)
    Tsunami.log(trainer, "accuracy/val", acc)
end

function Tsunami.configure_optimisers(m::FineTuneEncoder, trainer)
    opt_rule = Flux.Optimisers.Adam(1e-4)
    opt_state = Flux.Optimisers.setup(opt_rule, m)
    return opt_state
end

function run_train()
    # Prepare the dataset
    dataset = ImageDataset("../../Data/AID", imsize=(224, 224))
    train_data, test_data = Flux.MLUtils.splitobs(dataset, at=0.8, shuffle=true)
    train_loader = Flux.DataLoader(train_data, batchsize=16, collate=true, parallel=true, shuffle=true)
    test_loader = Flux.DataLoader(test_data, batchsize=16, collate=true, parallel=true)

    # Create and train the model
    model = FineTuneEncoder(length(dataset.labels))
    trainer = Trainer(max_epochs=5)
    Tsunami.fit!(model, trainer, train_loader, test_loader)
end

```

These are the results for ResNet-34 on a machine with 64 GB of RAM and an NVIDIA RTX 5090 with 32 GB of VRAM. VRAM usage is reported by [nvtop](https://github.com/Syllo/nvtop).

**Flux:** 19.1 GB VRAM - 11 seconds / epoch

**PyTorch:** 2.6 GB VRAM - 7 seconds / epoch

I also tried another experiment with a ViT base model, which produced the following.

**Flux:** 16.2 GB VRAM - 38 seconds / epoch

**PyTorch:** 5.1 GB VRAM - 23 seconds / epoch

To determine if Julia was perhaps just allocating more buffer space without actually using it, I tried to increase the batch size until I got an OOM. PyTorch was able to handle a batch size of 128 with ViT base, while Flux could only go up to 24 after throwing OOM errors at 32.

Does anyone have any idea what might be causing this discrepancy?

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [May 15, 2026, 7:31am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/2 "2026-05-15T07:31:37Z")

</div>

I would test two things :

- use [Interface · GPUArrays.jl](https://juliagpu.github.io/GPUArrays.jl/dev/interface/#Caching-Allocator) to only allocate once
- use [GitHub - EnzymeAD/Reactant.jl: Optimize Julia Functions With MLIR and XLA for High-Performance Execution on CPU, GPU, TPU and more. · GitHub](https://github.com/EnzymeAD/Reactant.jl) to enable kernel fusion reducing memory by a ton.

---

<div class="post-metadata">

**Author:** ![CarloLucibello](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carlolucibello/32/3278_2.png) [@CarloLucibello](https://discourse.julialang.org/u/CarloLucibello)\
**Post date:** [May 16, 2026, 8:04am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/3 "2026-05-16T08:04:21Z")

</div>

> - use [Interface · GPUArrays.jl](https://juliagpu.github.io/GPUArrays.jl/dev/interface/#Caching-Allocator) to only allocate once

Se how this is done in [Flux.train!](https://github.com/FluxML/Flux.jl/blob/master/src/train.jl#L111)

---

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 16, 2026, 5:41pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/4 "2026-05-16T17:41:48Z")

</div>

It looks like Tsunami.jl already caches with GPUArrays, which is what I’m using for training:

```julia
## SINGLE EPOCH TRAINING LOOP
for (batch_idx, batch) in enumerate(train_dataloader)
    fit_state.step += 1
    fit_state.batchsize = MLUtils.numobs(batch)

    hook(on_train_batch_start, model, trainer, batch, batch_idx)
        
    GPUArrays.@cached trainer.cache begin
        out, grad = gradient_train_step(model, trainer, batch, batch_idx)
    end
        
    hook(on_before_update, model, trainer, out, grad)
        
    GPUArrays.@cached trainer.cache begin 
        update!(trainer.optimisers, model, grad)
    end

    if fit_state.step == trainer.max_steps
        fit_state.should_stop = true
    end

    hook(on_train_batch_end, model, trainer, out, batch, batch_idx)
        
    ProgressMeter.next!(train_progbar,
        showvalues = values_for_train_progbar(trainer.metalogger),
        valuecolor = :yellow, 
        final = fit_state.should_stop || batch_idx == _length(train_dataloader),
        keep = fit_state.should_stop || fit_state.epoch == trainer.max_epochs
    )

    fit_state.should_stop && break
end

```

I’ll try re-writing the training loop to use Reactant to see if there’s any improvement.

---

<div class="post-metadata">

**Author:** ![csvance](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/csvance/32/218927_2.png) [@csvance](https://discourse.julialang.org/u/csvance)\
**Post date:** [May 16, 2026, 7:56pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/5 "2026-05-16T19:56:43Z")

</div>

In my experience w/ Lux.jl, it seems to allocate more than it actually uses. A PyTorch model that I know takes around 3GB for forwards/backwards pass w/ the same batch size shows 24GB when training it with Lux.jl. When I increase the batch size by 8x nvidia-smi then reports its using 48GB 🤔

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [May 16, 2026, 8:01pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/6 "2026-05-16T20:01:28Z")

</div>

Yes that’s why we’re so happy to have Reactant now, I wonder how much respecting Zygote way forced this at the time those layer were written

---

<div class="post-metadata">

**Author:** ![csvance](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/csvance/32/218927_2.png) [@csvance](https://discourse.julialang.org/u/csvance)\
**Post date:** [May 16, 2026, 8:39pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/7 "2026-05-16T20:39:22Z")

</div>

> [@yolhan\_mannes](#):
>
> Yes that’s why we’re so happy to have Reactant now,

The thing that’s likely keeping many people from using Reactant.jl is the documentation. It’s not totally clear how to do things like DDP style training, how to use Reactant in practice when also using SciMLSensitivity.jl, etc.

That being said it’s a very promising direction for ML in Julia. Some very cool stuff is possible like being able to directly load PyTorch models via StableHLO: [PyTorch StableHLO Support · Issue #2065 · EnzymeAD/Reactant.jl · GitHub](https://github.com/EnzymeAD/Reactant.jl/issues/2065)

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [May 16, 2026, 8:50pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/8 "2026-05-16T20:50:13Z")

</div>

Yes but with SciMLSensitivity it’s pretty rare you need GPU and for CPU small network Julia + Enzyme or Mooncake is largely enough they are some cases when you need a big network (PDE + spectral solver is one ) but they are pretty rare actually.

For now the best thing to do is to look at all the already implemented Lux+Reactant cases learning by example.

I think Reactant may change a lot before 1.0 though which explains not to hurry building a big doc for now, for instance

- If the @trace can be remove
- if switching option become more user friendly
- if the backend could be chosen on the fly so that is simplify other repo using reactant (edit : it’s always been here )
- if Metal gets added  
Ect ect ect

It’s funny though that even early like that I prefer Julia + Reactant way of handling MLIR than Jax or Mojo

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [May 17, 2026, 4:51pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/9 "2026-05-17T16:51:36Z")

</div>

Docs are always welcome! If there’s somethign missing do let us know.

I have no idea what “DDP style training” is so that’s at least one reason why that’s not in a doc xD

As for in progress features, yeah automatic removal of the need of @ trace is something we’re looking into, and there is an open PR adding Metal (and we just added initial trainium support too).

I’m curious what you mean by switching/backend issues?

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [May 17, 2026, 5:39pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/10 "2026-05-17T17:39:53Z")

</div>

I think I saw it somewhere (saw someone ask or an issue) to be able at compile option and/or array creation explicitly ask for it to be on CPU/GPU/TPU instead of / in addition to the global set\_backend way. In a user perceptive not adding that much but for another package wanting to dépend on Reactant that may be nice to have (especially if the package already dispatch on the KernelAbstraction backend selector).  
Also would be very cool to have a talk like [https://m.youtube.com/watch?v=XuMDzRmRPPQ&pp=ygUPSnVsaWFodWIgQ3VUaWxl](https://m.youtube.com/watch?v=XuMDzRmRPPQ&pp=ygUPSnVsaWFodWIgQ3VUaWxl) for Reactant I’m sure JuliaHub would be ok to do that ? The JuliaCon are great but limited in term of complexity and length

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [May 17, 2026, 5:54pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/11 "2026-05-17T17:54:04Z")

</div>

to\_rarray [and similar] take an optional either device, or backend:  
[Core Reactant API | Reactant.jl](https://enzymead.github.io/Reactant.jl/dev/api/api#Converting-Data) , defaulting to the global if not specified.

I think they’ve had those args since the API itself was created, so maybe it’s just not as well documented (help welcome!)? Or are you thinking of something else?

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [May 17, 2026, 5:55pm UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/12 "2026-05-17T17:55:18Z")

</div>

Oh sorry yes I didn’t know it and people may not know indeed. I will make a doc pr about it when I can Thank you

---

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 18, 2026, 4:39am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/13 "2026-05-18T04:39:46Z")

</div>

I re-implemented my training script to use Reactant and Lux:

```julia
function train_lux()
    lux_model = Lux.Chain(
        #Boltz.Vision.ResNet(34), 
        Boltz.Vision.ViT(:base), 
        Lux.Dense(1000, 30), 
    )

    rng = Random.default_rng()
    ps, st = Lux.setup(rng, lux_model)

    # Move to Reactant device
    dev = Lux.reactant_device()
    cdev = Lux.cpu_device()
    ps, st = ps |> dev, st |> dev

    imsize = 256
    dataset = ImageDataset("../../Data/AID", imsize=(imsize, imsize))
    train_data, test_data = Flux.MLUtils.splitobs(dataset, at=0.8, shuffle=true)
    train_loader = Flux.DataLoader(train_data, batchsize=128, shuffle=true, collate=true, parallel=true, partial=false) |> dev
    test_loader = Flux.DataLoader(test_data, batchsize=128, shuffle=true, collate=true, parallel=true, partial=false) |> dev

    model_compiled = Reactant.@compile lux_model(first(train_loader)[1], ps, Lux.testmode(st))

    opt = Optimisers.Adam(1f-4)
    tstate = Lux.Training.TrainState(lux_model, ps, st, opt)

    # 2. Training Loop
    loss_fn = Lux.CrossEntropyLoss(;logits=Val(true)) # or your custom loss function

    for epoch in 1:10
        @info "Epoch $epoch"
        total_loss = 0.0
        ProgressMeter.@showprogress for (xdata, ydata) in train_loader
            _, loss, _, tstate = Lux.Training.single_train_step!(
                Lux.AutoEnzyme(), loss_fn, (xdata, ydata), tstate
            )
            total_loss += loss
        end
        @info "loss:" total_loss / length(train_loader)

        total_acc = 0.0
        st_ = Lux.testmode(tstate.states)
        ProgressMeter.@showprogress for (x, y) in test_loader
            ŷ, st_ = model_compiled(x, tstate.parameters, st_)
            ŷ, y = cdev(ŷ), cdev(y)
            acc = accuracy(ŷ, y)
            total_acc += acc
        end
        @info "Epoch $epoch - Accuracy: $(total_acc / length(test_loader))"
    end
end

function accuracy(y_pred, y_true)
    y_pred = Lux.softmax(y_pred, dims=1)
    y_pred = map(x -> x[1], argmax(y_pred, dims=1))
    y_true = map(x -> x[1], argmax(y_true, dims=1))
    correct = sum(y_pred .== y_true)
    total = size(y_true, 2)
    return sum(correct) / total
end

```

These are the results for ViT Base:

**PyTorch:** 5.1 GB VRAM - 23 seconds / epoch

**Flux:** 16.2 GB VRAM - 38 seconds / epoch

**Lux:** Unknown GB VRAM - 25 seconds / epoch

I couldn’t determine the VRAM usage for LUX/Reactant due to pre-allocation. However, I was able to increase the batch size to a maximum of 128, which is significantly higher than under Flux (24) and equivalent to PyTorch at FP32 precision.

From these, it seems that the main issue is how Flux/Zygote allocates intermediate arrays, which results in much higher memory requirements than Lux/Reactant. If this is the direction that the Julia community is moving, I’ll probably modify my project to use Lux instead of Flux.

---

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 18, 2026, 4:51am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/14 "2026-05-18T04:51:31Z")

</div>

> That being said it’s a very promising direction for ML in Julia. Some very cool stuff is possible like being able to directly load PyTorch models via StableHLO: [PyTorch StableHLO Support · Issue #2065 · EnzymeAD/Reactant.jl · GitHub](https://github.com/EnzymeAD/Reactant.jl/issues/2065)

Does `StableHLO` produce native Julia models, where you can extract and modify layers, or is it essentially a black box like ONNX? The reason I ask is that my current project involves implementing `Timm` models as pure Julia equivalents in `Flux`, then I define a `load_params!` method that takes a PyTorch `state_dict` from the matching `Timm` model/layer and loads the parameters into the corresponding `Flux` layer. This has the advantage of producing a model that can be used like any other `Flux` layer, but obviously requires a fair amount of work to duplicate the original PyTorch code.

---

<div class="post-metadata">

**Author:** ![csvance](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/csvance/32/218927_2.png) [@csvance](https://discourse.julialang.org/u/csvance)\
**Post date:** [May 18, 2026, 4:54am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/15 "2026-05-18T04:54:45Z")

</div>

> [@JoshuaBillson](#):
>
> Does `StableHLO` produce native Julia models, where you can extract and modify layers, or is it essentially a black box like ONNX? The reason I ask is that my current project involves implementing `Timm` models as pure Julia equivalents in `Flux`, then I define a `load_params!` method that takes a PyTorch `state_dict` from the matching `Timm` model/layer and loads the parameters into the corresponding `Flux` layer. This has the advantage of producing a model that can be used like any other `Flux` layer, but obviously requires a fair amount of work to duplicate the original PyTorch code.

It doesn’t produce native Julia models, but you can use it inside the rest of your native Julia model. Well, I suppose as long as you compile it to Reactant.jl.

Also see here, I’m working on a timm port for Lux.jl: [[ANN] Jimm.jl: Lux ports of timm image backbones, with HuggingFace pretrained weights](https://discourse.julialang.org/t/ann-jimm-jl-lux-ports-of-timm-image-backbones-with-huggingface-pretrained-weights/137153/1)

It would likely be pretty easy to do the same sort of workflow I did here to port things over to Flux.jl, but I’m not personally familiar with it.

---

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 18, 2026, 5:01am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/16 "2026-05-18T05:01:00Z")

</div>

> Also see here, I’m working on a timm port for [Lux.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/Lux): [[ANN] Jimm.jl: Lux ports of timm image backbones, with HuggingFace pretrained weights](https://discourse.julialang.org/t/ann-jimm-jl-lux-ports-of-timm-image-backbones-with-huggingface-pretrained-weights/137153/1)

That’s almost exactly what I’m working on (my working name was even Jimm). I’ll take a look and see if I can contribute. So far, I’ve implemented all variants of Timm’s `VisionTransformer`, `ConvNeXt` (both v1 and v2), and `Eva` (basically ViT with rotary positional embeddings used by SAM3). I also have implementations for `Swin`, `PVT`, and `Twins`, but I didn’t get around to adding pre-trained weights yet. It should be relatively straightforward to convert from `Flux` to `Lux`.

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [May 18, 2026, 5:01am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/17 "2026-05-18T05:01:26Z")

</div>

We should be able to write a stablehlo-\>native julia arrays in Reactant/MLIR [and have been meaning to, but it’s not currently high priority – if anyone wants to give it a go, please reach out and I’d be happy to help you get started!]

---

<div class="post-metadata">

**Author:** ![csvance](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/csvance/32/218927_2.png) [@csvance](https://discourse.julialang.org/u/csvance)\
**Post date:** [May 18, 2026, 5:06am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/18 "2026-05-18T05:06:18Z")

</div>

That would be great! I don’t have any ViT added yet. ConvNext V1 and V2 are done along with BiT ResNetV2. I’m not sure all the differences between Flux.jl and Lux.jl, but I went with Lux.jl because I wanted everything to work nicely with the SciML ecosystem.

Having any modern pretrained backbones is already a huge roadblock removed for people working with computer vision problems in Julia.

---

<div class="post-metadata">

**Author:** ![JoshuaBillson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/joshuabillson/32/52706_2.png) [@JoshuaBillson](https://discourse.julialang.org/u/JoshuaBillson)\
**Post date:** [May 18, 2026, 5:16am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/19 "2026-05-18T05:16:08Z")

</div>

> Having any modern pretrained backbones is already a huge roadblock removed for people working with computer vision problems in Julia.

This was exactly my thinking. My research involves land cover classification, and not being able to access state-of-the-art vision models in Julia has been a huge issue. As a result of my project, I successfully fine-tuned `vit_pe_spatial_base_patch16_512.fb`, the same vision encoder used in Meta’s [SAM 3](https://ai.meta.com/research/sam3/), achieving SOTA metrics across several benchmarks. However, high memory usage severely limited my ability to train larger models and use larger batch sizes, which prompted this topic. Since `Lux` seems to resolve this issue, I’m happy to shift my focus to `Jimm`.

---

<div class="post-metadata">

**Author:** ![csvance](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/csvance/32/218927_2.png) [@csvance](https://discourse.julialang.org/u/csvance)\
**Post date:** [May 18, 2026, 5:31am UTC](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124/20 "2026-05-18T05:31:27Z")

</div>

It seems we were literally thinking the exact same thing at the exact same time. I mostly made the announcement for the package to find collaborators; whenever you are ready I can get you full access to the repo.

[Next page](https://discourse.julialang.org/t/significantly-higher-vram-usage-and-slower-training-on-flux-compared-to-pytorch/137124.md?page=2)
