# Memory Optimization Techniques in Eagle for Training Large Multimodal Models

> Discover Eagle's memory optimization techniques like FlashAttention and gradient checkpointing for efficient large model training. Train 8B models on a single GPU.

- Repository: [NVIDIA Research Projects/Eagle](https://github.com/NVlabs/Eagle)
- Tags: deep-dive
- Published: 2026-06-28

---

**Eagle combines Flash-Attention kernels, gradient checkpointing with optional CPU offloading, fused linear-cross-entropy Triton kernels, and sharded loss computation to train 8B-parameter vision-language models with 250k-token vocabularies and 8k context lengths on a single A100-40GB GPU.**

The NVlabs/Eagle repository implements a comprehensive memory optimization stack that enables fine-tuning of large-scale multimodal models on limited GPU resources. By avoiding materialization of full attention matrices and eliminating large logits tensors during loss computation, these techniques reduce peak memory usage by orders of magnitude compared to standard PyTorch training loops.

## Flash-Attention and Memory-Efficient Kernels

Eagle replaces standard PyTorch attention implementations with kernel fusion strategies that avoid the **quadratic memory bottleneck** of the full `Q·Kᵀ` matrix.

### Flash-Attention 2 and 3 Integration

In [`modeling_qwen2.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_qwen2.py) and [`modeling_locateanything.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_locateanything.py), Eagle swaps the default `torch.nn.MultiheadAttention` with Flash-Attention kernels that compute attention in-place without materializing the full `N²` attention matrix. This is implemented at line 480 in [`modeling_qwen2.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_qwen2.py) and line 104 in [`modeling_locateanything.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_locateanything.py), where the attention implementation is selected via the `attn_implementation` configuration flag.

### XFormers Memory-Efficient Attention Fallback

When Flash-Attention is unavailable, Eagle falls back to the XFormers `memory_efficient_attention` kernel. The monkey patch in [`llama_xformers_attn_monkey_patch.py`](https://github.com/NVlabs/Eagle/blob/main/llama_xformers_attn_monkey_patch.py) (line 90) intercepts the standard attention call and routes it through the XFormers kernel, which also avoids allocating the full attention score matrix.

### Packed Attention for Multi-Token Positioning

For custom multi-token-position masks (MTP), Eagle patches the Flash-Attention forward pass to accept 2-D masks directly. The implementation in [`packing_attention.py`](https://github.com/NVlabs/Eagle/blob/main/packing_attention.py) (line 322) keeps mask memory usage minimal by avoiding the expansion to full attention bias tensors.

## Gradient Checkpointing and Activation Management

Eagle employs multiple strategies to reduce activation memory during the backward pass, from standard recomputation to aggressive CPU offloading.

### Standard Activation Checkpointing

The repository enables standard PyTorch gradient checkpointing via `torch.utils.checkpoint`, configured in [`train.py`](https://github.com/NVlabs/Eagle/blob/main/train.py) (line 1051) through the `gradient_checkpointing=True` training argument. In [`modeling_qwen2.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_qwen2.py) (line 1231), the model wraps transformer layers with checkpointing functions that store only selected intermediate activations and recompute the rest during backpropagation.

### UnsLoTH Non-Reentrant Optimization

Eagle integrates the UnsLoTH patch from [`unsloth_checkpoint.py`](https://github.com/NVlabs/Eagle/blob/main/unsloth_checkpoint.py) (line 141) to provide a **non-reentrant** checkpointing mode. This implementation avoids the overhead of nested autograd graphs, making it significantly faster for long sequence contexts while maintaining the same memory savings as standard checkpointing.

### CPU Offloading for Checkpoint Buffers

When GPU memory is exhausted, Eagle automatically moves checkpoint buffers to host RAM. The patch in [`unsloth_checkpoint.py`](https://github.com/NVlabs/Eagle/blob/main/unsloth_checkpoint.py) (lines 149 and 221) swaps the standard checkpoint function with an offloading version that maintains the tensor storage in page-locked CPU memory, allowing the model to continue training without reducing batch size or sequence length.

## Low-Memory Loss Computation

For models with large vocabularies (100k+ tokens), the final classification layer typically dominates memory usage. Eagle eliminates the `batch × sequence × vocabulary` logits tensor through custom loss functions.

### Sharded Cross-Entropy Without Logits

The `low_mem_cross_ent` function in [`sp_utils/loss.py`](https://github.com/NVlabs/Eagle/blob/main/sp_utils/loss.py) (line 19) projects hidden states into the vocabulary space in shards, computes the cross-entropy loss for each shard, and discards the logits immediately. This avoids allocating the full `B·T·V` tensor that normally dominates memory for large vocabularies.

```python
import torch
from Embodied.eaglevl.sp_utils.loss import low_mem_cross_ent

hidden = torch.randn(2, 4096, 4096, device="cuda")    # B, T, H

lm_head = torch.randn(250_000, 4096, device="cuda")   # V, H

labels = torch.randint(0, 250_000, (2, 4096), device="cuda")

# Compute loss with 4 shards, avoiding (2, 4096, 250000) logits allocation

loss = low_mem_cross_ent.apply(hidden, lm_head, labels, 4)
loss.backward()

```

### Liger-Fused Linear and Cross-Entropy

Eagle implements a Triton-based kernel in [`liger_loss_weight_ops.py`](https://github.com/NVlabs/Eagle/blob/main/liger_loss_weight_ops.py) that fuses the final linear projection with cross-entropy computation. This kernel uses **in-place gradient accumulation** and never materializes the full logits matrix, supporting vocabularies up to 250k tokens. The implementation uses `torch.amp.custom_fwd/bwd` hooks (line 33) to maintain numerical stability while keeping forward passes in FP32 only where required.

```python
from Embodied.eaglevl.train.liger_loss_weight_ops import LigerFusedLinearCrossEntropyLoss

criterion = LigerFusedLinearCrossEntropyLoss()
logits, grad_input, grad_weight, grad_bias = criterion.forward(
    linear_weight, hidden, target, bias=None
)

```

## Data Loading and Precision Optimization

Beyond model architecture, Eagle optimizes the data pipeline and numerical precision to minimize memory overhead.

### Pin-Memory Streaming

The dataloader configuration in [`locany_finetune_magi_stream.py`](https://github.com/NVlabs/Eagle/blob/main/locany_finetune_magi_stream.py) (line 1186) enables `pin_memory=True`, which copies tensors directly into page-locked host memory. This eliminates extra staging buffers during GPU transfer and reduces host-device synchronization overhead.

### Enabling Flash-Attention and Checkpointing

To activate the full memory optimization stack, configure the training arguments and model configuration as follows:

```python
from transformers import Trainer, TrainingArguments

training_args = TrainingArguments(
    output_dir="out",
    per_device_train_batch_size=1,
    gradient_checkpointing=True,          # Enable activation checkpointing

    fp16=True,
    dataloader_pin_memory=True,           # Enable pin-memory streaming

)

# Activate Flash-Attention kernels

model.config.attn_implementation = "flash_attention_2"

```

## Summary

Eagle implements a layered approach to GPU memory optimization that enables training otherwise impossible model configurations:

- **Flash-Attention 2/3** kernels in [`modeling_qwen2.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_qwen2.py) and [`modeling_locateanything.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_locateanything.py) eliminate the quadratic attention memory bottleneck
- **Gradient checkpointing** with UnsLoTH patches reduces activation memory by recomputing forward passes during backpropagation
- **Offloaded checkpointing** moves activation buffers to CPU RAM when GPU memory is exhausted
- **Sharded and fused loss functions** avoid allocating large logits tensors for massive vocabularies
- **Pin-memory data loading** and **custom AMP hooks** minimize overhead in the data pipeline and precision management

## Frequently Asked Questions

### How does Eagle handle the memory bottleneck of large vocabularies (100k+ tokens)?

Eagle avoids materializing the full `batch × sequence × vocabulary` logits tensor by using **sharded cross-entropy** in [`sp_utils/loss.py`](https://github.com/NVlabs/Eagle/blob/main/sp_utils/loss.py) and **Liger-fused kernels** in [`liger_loss_weight_ops.py`](https://github.com/NVlabs/Eagle/blob/main/liger_loss_weight_ops.py). These compute loss projections in shards or fuse the linear layer with the loss computation, respectively, never storing the full logits matrix in GPU memory.

### What is the difference between standard gradient checkpointing and the UnsLoTH patch in Eagle?

Standard gradient checkpointing in [`modeling_qwen2.py`](https://github.com/NVlabs/Eagle/blob/main/modeling_qwen2.py) uses PyTorch's `torch.utils.checkpoint` with reentrant autograd, which can be slow for long sequences. The UnsLoTH patch in [`unsloth_checkpoint.py`](https://github.com/NVlabs/Eagle/blob/main/unsloth_checkpoint.py) provides a **non-reentrant** implementation that is faster for long contexts while maintaining identical memory savings. The patch also supports offloading checkpoint buffers to CPU RAM when GPU memory is exhausted.

### Can Eagle train models without Flash-Attention installed?

Yes. Eagle includes a fallback to **XFormers memory-efficient attention** via [`llama_xformers_attn_monkey_patch.py`](https://github.com/NVlabs/Eagle/blob/main/llama_xformers_attn_monkey_patch.py). This kernel also avoids materializing the full attention score matrix, though Flash-Attention 2/3 provides superior performance when available.

### Where does Eagle implement the offloading of checkpoint buffers to CPU?

The CPU offloading logic is implemented in [`unsloth_checkpoint.py`](https://github.com/NVlabs/Eagle/blob/main/unsloth_checkpoint.py) (lines 149 and 221). This patch intercepts the standard checkpoint function and moves the hidden states to page-locked host memory before storing them, allowing training to continue on GPUs with limited VRAM by trading compute for host memory bandwidth.