# How a Decoder-Only Text Encoder Enhances Text-Image Alignment in Sana

> Discover how Sana uses a decoder-only text encoder to improve text-image alignment, extracting richer embeddings for optimized generation.

- Repository: [NVIDIA Research Projects/Sana](https://github.com/NVlabs/Sana)
- Tags: deep-dive
- Published: 2026-05-19

---

**Sana improves text-image alignment by replacing the traditional T5 encoder with a decoder-only language model that extracts richer, generation-optimized embeddings from its hidden states.**

NVlabs/Sana employs a decoder-only text encoder architecture that fundamentally changes how textual prompts condition the diffusion process. By leveraging decoder-only language models like Gemma instead of conventional encoder-only architectures such as T5, Sana captures embeddings that reflect the generative nature of language, resulting in tighter alignment between prompts and generated visual content while maintaining computational efficiency on consumer hardware.

## Architecture of the Decoder-Only Text Encoder

Sana's text encoding pipeline centers on repurposing causal language models as feature extractors. Unlike encoder-decoder or encoder-only designs, this approach utilizes the decoder's hidden states to create conditioning signals for the diffusion model.

### Model Loading via `get_tokenizer_and_text_encoder`

In [`diffusion/model/builder.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/builder.py), the function `get_tokenizer_and_text_encoder` handles the instantiation of decoder-only models. When a `gemma-*` model name is specified, the code loads the checkpoint using HuggingFace's `AutoModelForCausalLM` and immediately extracts the decoder component:

```python

# From diffusion/model/builder.py (lines 85-93)

text_encoder = AutoModelForCausalLM.from_pretrained(
    model_name, 
    torch_dtype=dtype, 
    device_map=device
).get_decoder()

```

This extraction strategy isolates the transformer decoder blocks, discarding the language modeling head while preserving the rich contextual representations learned during pretraining on massive text corpora.

### Tokenization Consistency

The implementation maintains strict token-level consistency by using the original model's tokenizer (e.g., Gemma's `AutoTokenizer`). This preserves the embedding matrix alignment between token IDs and their vector representations, ensuring that the decoder receives inputs compatible with its pretrained weights. The tokenizer processes prompts with fixed maximum lengths (typically 77 tokens) and padding to create uniform tensor shapes for batch processing.

### Embedding Extraction Strategy

During training and inference, Sana extracts embeddings from the decoder's last hidden state. In [`train_video_scripts/train_video_ivjoint_chunk.py`](https://github.com/NVlabs/Sana/blob/main/train_video_scripts/train_video_ivjoint_chunk.py) (lines 506-512), the forward pass is executed as:

```python
txt_tokens = tokenizer(
    prompt, 
    return_tensors="pt", 
    padding="max_length", 
    max_length=77, 
    truncation=True
)

# Extract last hidden state from decoder-only encoder

text_embeddings = text_encoder(
    txt_tokens.input_ids, 
    attention_mask=txt_tokens.attention_mask
)[0]  # Shape: (batch, seq_len, hidden_dim)

```

The `[0]` index retrieves the hidden states tensor, which serves as the dense text embedding fed into the diffusion model's conditioning pathway. Because the decoder has been optimized for next-token prediction, its hidden states embed strong semantic and syntactic cues that guide the diffusion process more precisely than traditional encoder outputs.

## Mechanisms for Enhanced Text-Image Alignment

The decoder-only architecture improves cross-modal alignment through three primary mechanisms: generative training objectives, in-context adaptation capabilities, and efficient model scaling.

### Context-Aware Embeddings from Generative Training

Decoder-only models like Gemma are trained with a causal language modeling objective—predicting the next token given previous context. This training paradigm forces the model to develop deep contextual understanding and maintain rich representations of semantic relationships throughout the hidden layers. When these hidden states serve as conditioning signals for image generation, they carry stronger generative semantics compared to encoder-only models that optimize purely for bidirectional encoding.

### In-Context Learning for Instruction Following

According to [`docs/8bit_sana.md`](https://github.com/NVlabs/Sana/blob/main/docs/8bit_sana.md) (lines 21-23), Sana leverages "complex human instruction with in-context learning to enhance the image-text alignment." By constructing prompts that include system instructions, few-shot examples, or style descriptors before the actual caption, the decoder adapts its hidden representations to match the specific aesthetic or structural requirements of the target output. This dynamic adaptation allows the same base model to condition the diffusion process differently based on instructional context, yielding more accurate prompt adherence.

### Efficiency Gains with Smaller Models

Sana utilizes compact decoder-only models such as Gemma-2B and Gemma-2-9B, which are significantly smaller than encoder-decoder alternatives like T5-XXL (11B parameters). This size reduction enables:

- **Consumer GPU compatibility**: The 2B parameter model runs efficiently on standard consumer graphics cards.
- **Faster inference**: Reduced memory bandwidth and computational requirements speed up text encoding.
- **Maintained quality**: Despite the smaller footprint, the decoder's generative pretraining provides embedding quality that meets or exceeds larger encoder-only models for diffusion conditioning.

## Practical Implementation Example

The following code demonstrates how to load and utilize Sana's decoder-only text encoder for extracting prompt embeddings:

```python
from diffusion.model.builder import get_tokenizer_and_text_encoder

# Initialize Gemma-2B as the decoder-only text encoder

tokenizer, text_encoder = get_tokenizer_and_text_encoder(
    name="gemma-2b", 
    device="cuda"
)

def embed_prompt(prompt: str):
    """
    Tokenize input and extract mean-pooled embeddings 
    from the decoder's last hidden state.
    """
    # Tokenize with fixed length padding

    tokens = tokenizer(
        prompt,
        return_tensors="pt",
        padding="max_length",
        max_length=77,
        truncation=True
    ).to(text_encoder.device)

    # Forward pass through decoder-only model

    hidden_states = text_encoder(
        tokens.input_ids,
        attention_mask=tokens.attention_mask
    )[0]  # (batch, seq_len, hidden_dim)

    # Mean pool across sequence dimension for single vector

    return hidden_states.mean(dim=1)  # (batch, hidden_dim)

```

The returned tensor can be directly concatenated with latent representations or fed into cross-attention layers within the diffusion model, exactly as implemented in Sana's training pipeline.

## Summary

- **Decoder-only extraction**: Sana loads decoder-only LMs (Gemma) via `get_tokenizer_and_text_encoder` in [`diffusion/model/builder.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/builder.py), extracting the decoder backbone for feature extraction.
- **Generative embeddings**: Hidden states from next-token prediction training provide richer semantic conditioning than encoder-only architectures.
- **In-context adaptation**: Complex prompting strategies allow the decoder to adjust embeddings based on instructional context, improving alignment quality.
- **Efficiency optimization**: Models like Gemma-2B offer high-quality embeddings with significantly lower memory requirements than T5-XXL.
- **Training integration**: The encoder is invoked in training scripts like [`train_video_scripts/train_video_ivjoint_chunk.py`](https://github.com/NVlabs/Sana/blob/main/train_video_scripts/train_video_ivjoint_chunk.py) to produce text embeddings that guide the diffusion process.

## Frequently Asked Questions

### What is the main advantage of using a decoder-only text encoder over T5 in Sana?

The primary advantage lies in the **generative training objective**. Decoder-only models optimize for next-token prediction, which produces hidden states that capture richer contextual relationships and generation-aware semantics. These embeddings guide the diffusion model more effectively than T5's bidirectional encoding, resulting in tighter text-image alignment while using significantly fewer parameters (2B vs 11B).

### How does Sana extract embeddings from a decoder-only language model?

Sana extracts embeddings by calling the decoder's forward pass and retrieving the last hidden state tensor. Specifically, in the training scripts, the code executes `text_encoder(input_ids, attention_mask=attention_mask)[0]`, where index `[0]` accesses the hidden states. The shape `(batch, seq_len, hidden_dim)` provides token-level embeddings that are typically mean-pooled or used directly for cross-modal conditioning.

### Can I use other decoder-only models besides Gemma with Sana?

Yes. The `get_tokenizer_and_text_encoder` function in [`diffusion/model/builder.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/builder.py) is designed to support any decoder-only architecture available through HuggingFace's `AutoModelForCausalLM`. You can specify alternative model checkpoints in the configuration, provided they implement the `.get_decoder()` method and maintain compatible hidden state dimensions for the diffusion model's conditioning pathway.

### Does using a decoder-only encoder impact inference speed?

No. In fact, Sana's decoder-only approach improves inference efficiency. Models like Gemma-2B contain fewer parameters than T5-XXL and require less memory bandwidth during text encoding. The [`docs/8bit_sana.md`](https://github.com/NVlabs/Sana/blob/main/docs/8bit_sana.md) documentation highlights this efficiency as a core design goal, enabling high-quality text-image alignment on consumer GPUs without the latency penalties associated with larger encoder-decoder models.