🔧 AI Nachrichten Major AI platforms go down in unprecedented simultaneous outage(03.09.2026 um 17:34 Uhr)
🔧 AI Nachrichten ChatGPT, Claude, and Grok Down? Users Report Widespread Outages(03.09.2026 um 19:14 Uhr)
🔧 AI Nachrichten OpenAI Launches GPT-6 Astra, Says We May Have Entered the AGI Era(03.09.2026 um 22:08 Uhr)
🔧 AI Nachrichten Claude Comes to CarPlay as Fifth Major AI Chatbot App(05.09.2026 um 05:31 Uhr)
🔧 AI Nachrichten OpenAI’s GPT-6 Astra Is AGI, Says NVIDIA CEO Jensen Huang(07.09.2026 um 06:31 Uhr)
🔧 AI Nachrichten Blame AI companies for Mac mini and Mac Studio shortage(31.08.2026 um 10:32 Uhr)
🔧 AI Nachrichten Major AI platforms go down in unprecedented simultaneous outage(03.09.2026 um 17:34 Uhr)
🔧 AI Nachrichten ChatGPT, Claude, and Grok Down? Users Report Widespread Outages(03.09.2026 um 19:14 Uhr)
🔧 AI Nachrichten OpenAI Launches GPT-6 Astra, Says We May Have Entered the AGI Era(03.09.2026 um 22:08 Uhr)
🔧 AI Nachrichten Claude Comes to CarPlay as Fifth Major AI Chatbot App(05.09.2026 um 05:31 Uhr)
🔧 AI Nachrichten OpenAI’s GPT-6 Astra Is AGI, Says NVIDIA CEO Jensen Huang(07.09.2026 um 06:31 Uhr)
🔧 AI Nachrichten Blame AI companies for Mac mini and Mac Studio shortage(31.08.2026 um 10:32 Uhr)

🔧 AI Nachrichten 🕛 kürzlich 22 Min Lesezeit
0

Optimizing Transformer Models for Variable-Length Input Sequences

↗ Quelle (towardsdatascience.com)
🗣️ Stimme:
📑 Inhaltsübersicht
📺
towardsdatascience.com

How PyTorch NestedTensors, FlashAttention2, and xFormers can Boost Performance and Reduce AI Costs

Photo by

As generative AI (genAI) models grow in both popularity and scale, so do the computational demands and costs associated with their training and deployment. Optimizing these models is crucial for enhancing their runtime performance and reducing their operational expenses. At the heart of modern genAI systems is the Transformer architecture and its attention mechanism, which is notably compute-intensive.

In a all of the sample tensors along a new dimension — the batch dimension. However, . This solution requires appropriate masking within the model so that the output is not affected by the irrelevant tensor elements. In the case of attention layers, a padding mask indicates which tokens are padding and should not be attended to (e.g., see sequences along an existing dimension instead of , . Denoting the sum of the lengths of all of the individual by N and adopting ,

  • .
  • Integration into Existing HuggingFace Models

    For teams working with pre-trained models, transitioning to these optimizations might seem challenging. We will demonstrate how and model defined ).

    Transformer Block

    We begin by constructing a basic Transformer block, specifically designed to facilitate experimentation with different attention mechanisms and optimizations. While our block performs the same computation as standard Transformer blocks, we make slight modifications to the usual choice of operators in order to support the possibility of PyTorch ).

    # general imports
    import time, functools

    # torch imports
    import torch
    from torch.utils.data import Dataset, DataLoader
    import torch.nn as nn

    # Define Transformer settings
    BATCH_SIZE = 32
    NUM_HEADS = 16
    HEAD_DIM = 64
    DIM = NUM_HEADS * HEAD_DIM
    DEPTH = 24
    NUM_TOKENS = 1024
    MAX_SEQ_LEN = 1024
    PAD_ID = 0
    DEVICE = 'cuda'

    class MyAttentionBlock(nn.Module):
    def __init__(
    self,
    attn_fn,
    dim,
    num_heads,
    format=None,
    **kwargs
    ):
    super().__init__()
    self.attn_fn = attn_fn
    self.num_heads = num_heads
    self.dim = dim
    self.head_dim = dim // num_heads
    self.norm1 = nn.LayerNorm(dim, bias=False)
    self.norm2 = nn.LayerNorm(dim, bias=False)
    self.qkv = nn.Linear(dim, dim * 3)
    self.proj = nn.Linear(dim, dim)

    # mlp layers
    self.fc1 = nn.Linear(dim, dim * 4)
    self.act = nn.GELU()
    self.fc2 = nn.Linear(dim * 4, dim)

    self.permute = functools.partial(torch.transpose, dim0=1, dim1=2)
    if format == 'bshd':
    self.permute = nn.Identity()

    def mlp(self, x):
    x = self.fc1(x)
    x = self.act(x)
    x = self.fc2(x)
    return x

    def reshape_and_permute(self,x, batch_size):
    x = x.view(batch_size, -1, self.num_heads, self.head_dim)
    return self.permute(x)

    def forward(self, x_in, attn_mask=None):
    batch_size = x_in.size(0)
    x = self.norm1(x_in)
    qkv = self.qkv(x)

    # rather than first reformatting and then splitting the input
    # state, we first split and then reformat q, k, v in order to
    # support PyTorch Nested Tensors
    q, k, v = qkv.chunk(3, -1)
    q = self.reshape_and_permute(q, batch_size)
    k = self.reshape_and_permute(k, batch_size)
    v = self.reshape_and_permute(v, batch_size)

    # call the attn_fn with the input attn_mask
    x = self.attn_fn(q, k, v, attn_mask=attn_mask)

    # reformat output
    x = self.permute(x).reshape(batch_size, -1, self.dim)
    x = self.proj(x)
    x = x + x_in
    x = x + self.mlp(self.norm2(x))
    return x

    Transformer Decoder Model

    Building on our programmable Transformer block, we construct a typical Transformer decoder model.

    class MyDecoder(nn.Module):
    def __init__(
    self,
    block_fn,
    num_tokens,
    dim,
    num_heads,
    num_layers,
    max_seq_len,
    pad_idx=None
    ):
    super().__init__()
    self.num_heads = num_heads
    self.pad_idx = pad_idx
    self.embedding = nn.Embedding(num_tokens, dim, padding_idx=pad_idx)
    self.positional_embedding = nn.Embedding(max_seq_len, dim)
    self.blocks = nn.ModuleList([
    block_fn(
    dim=dim,
    num_heads=num_heads
    )
    for _ in range(num_layers)])
    self.output = nn.Linear(dim, num_tokens)

    def embed_tokens(self, input_ids, position_ids=None):
    x = self.embedding(input_ids)
    if position_ids is None:
    position_ids = torch.arange(input_ids.shape[1],
    device=x.device)
    x = x + self.positional_embedding(position_ids)
    return x

    def forward(self, input_ids, position_ids=None, attn_mask=None):
    # Embed tokens and add positional encoding
    x = self.embed_tokens(input_ids, position_ids)
    if self.pad_idx is not None:
    assert attn_mask is None
    # create a padding mask - we assume boolean masking
    attn_mask = (input_ids != self.pad_idx)
    attn_mask = attn_mask.view(BATCH_SIZE, 1, 1, -1) \
    .expand(-1, self.num_heads, -1, -1)

    for b in self.blocks:
    x = b(x, attn_mask)

    logits = self.output(x)
    return logits

    Variable Length Sequence Input

    Next, we create a dataset containing sequences of variable lengths, where each sequence is made up of randomly generated tokens. For simplicity, we (arbitrarily) select a fixed distribution for the sequence lengths. In real-world scenarios, the distribution of sequence lengths typically reflects the nature of the data, such as the length of documents or audio segments. Note, that the distribution of lengths directly affects the computational inefficiencies caused by padding.

    # Use random data
    class FakeDataset(Dataset):
    def __len__(self):
    return 1000000

    def __getitem__(self, index):
    length = torch.randint(1, MAX_SEQ_LEN, (1,))
    sequence = torch.randint(1, NUM_TOKENS, (length + 1,))
    input = sequence[:-1]
    target = sequence[1:]
    return input, target

    def pad_sequence(sequence, length, pad_val):
    return torch.nn.functional.pad(
    sequence,
    (0, length - sequence.shape[0]),
    value=pad_val
    )

    def collate_with_padding(batch):
    padded_inputs = []
    padded_targets = []
    for b in batch:
    padded_inputs.append(pad_sequence(b[0], MAX_SEQ_LEN, PAD_ID))
    padded_targets.append(pad_sequence(b[1], MAX_SEQ_LEN, PAD_ID))
    padded_inputs = torch.stack(padded_inputs, dim=0)
    padded_targets = torch.stack(padded_targets, dim=0)
    return {
    'inputs': padded_inputs,
    'targets': padded_targets
    }

    def data_to_device(data, device):
    if isinstance(data, dict):
    return {
    key: data_to_device(val,device)
    for key, val in data.items()
    }
    elif isinstance(data, (list, tuple)):
    return type(data)(
    data_to_device(val, device) for val in data
    )
    elif isinstance(data, torch.Tensor):
    return data.to(device=device, non_blocking=True)
    else:
    return data.to(device=device)

    Training/Evaluation Loop

    Lastly, we implement a main function that performs training/evaluation on input sequences of varying length.

    def main(
    block_fn,
    data_collate_fn=collate_with_padding,
    pad_idx=None,
    train=True,
    compile=False
    ):
    torch.random.manual_seed(0)
    device = torch.device(DEVICE)
    torch.set_float32_matmul_precision("high")

    # Create dataset and dataloader
    data_set = FakeDataset()
    data_loader = DataLoader(
    data_set,
    batch_size=BATCH_SIZE,
    collate_fn=data_collate_fn,
    num_workers=12,
    pin_memory=True,
    drop_last=True
    )

    model = MyDecoder(
    block_fn=block_fn,
    num_tokens=NUM_TOKENS,
    dim=DIM,
    num_heads=NUM_HEADS,
    num_layers=DEPTH,
    max_seq_len=MAX_SEQ_LEN,
    pad_idx=pad_idx
    ).to(device)

    if compile:
    model = torch.compile(model)

    # Define loss and optimizer
    criterion = torch.nn.CrossEntropyLoss(ignore_index=PAD_ID)
    optimizer = torch.optim.SGD(model.parameters())

    def train_step(model, inputs, targets,
    position_ids=None, attn_mask=None):
    with torch.amp.autocast(DEVICE, dtype=torch.bfloat16):
    outputs = model(inputs, position_ids, attn_mask)
    outputs = outputs.view(-1, NUM_TOKENS)
    targets = targets.flatten()
    loss = criterion(outputs, targets)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()

    @torch.no_grad()
    def eval_step(model, inputs, targets,
    position_ids=None, attn_mask=None):
    with torch.amp.autocast(DEVICE, dtype=torch.bfloat16):
    outputs = model(inputs, position_ids, attn_mask)
    if outputs.is_nested:
    outputs = outputs.data._values
    targets = targets.data._values
    else:
    outputs = outputs.view(-1, NUM_TOKENS)
    targets = targets.flatten()
    loss = criterion(outputs, targets)
    return loss

    if train:
    model.train()
    step_fn = train_step
    else:
    model.eval()
    step_fn = eval_step

    t0 = time.perf_counter()
    summ = 0
    count = 0

    for step, data in enumerate(data_loader):
    # Copy data to GPU
    data = data_to_device(data, device=device)
    step_fn(model, data['inputs'], data['targets'],
    position_ids=data.get('indices'),
    attn_mask=data.get('attn_mask'))

    # Capture step time
    batch_time = time.perf_counter() - t0
    if step > 20: # Skip first steps
    summ += batch_time
    count += 1
    t0 = time.perf_counter()
    if step >= 100:
    break
    print(f'average step time: {summ / count}')

    PyTorch SDPA with Padding

    For our baseline experiments, we configure our Transformer block to utilize PyTorch’s . These were run on an and in SDPA in evaluation mode. Currently a prototype feature, .

    PyTorch NestedTensors are supported by a we demonstrated the use of from does not support torch.compile (at the time of this writing).

    def collate_concat(batch):
    inputs = torch.concat([b[0] for b in batch]).unsqueeze(0)
    targets = torch.concat([b[1] for b in batch]).unsqueeze(0)
    indices = torch.concat([torch.arange(b[0].shape[0]) for b in batch])
    seqlens = torch.tensor([b[0].shape[0] for b in batch])
    seqlens = torch.cumsum(seqlens, dim=0, dtype=torch.int32)
    cu_seqlens = torch.nn.functional.pad(seqlens, (1, 0))

    return {
    'inputs': inputs,
    'targets': targets,
    'indices': indices,
    'attn_mask': cu_seqlens
    }

    from flash_attn import flash_attn_varlen_func
    fa_varlen = lambda q, k, v, attn_mask: flash_attn_varlen_func(
    q.squeeze(0),
    k.squeeze(0),
    v.squeeze(0),
    cu_seqlens_q=attn_mask,
    cu_seqlens_k=attn_mask,
    max_seqlen_q=MAX_SEQ_LEN,
    max_seqlen_k=MAX_SEQ_LEN
    ).unsqueeze(0)

    fa_varlen_causal = lambda q, k, v, attn_mask: flash_attn_varlen_func(
    q.squeeze(0),
    k.squeeze(0),
    v.squeeze(0),
    cu_seqlens_q=attn_mask,
    cu_seqlens_k=attn_mask,
    max_seqlen_q=MAX_SEQ_LEN,
    max_seqlen_k=MAX_SEQ_LEN,
    causal=True
    ).unsqueeze(0)

    block_fn = functools.partial(MyAttentionBlock,
    attn_fn=fa_varlen,
    format='bshd')

    causal_block_fn = functools.partial(MyAttentionBlock,
    attn_fn=fa_varlen_causal,
    format='bshd')

    print('flash-attn eval')
    main(
    block_fn=block_fn,
    data_collate_fn=collate_concat,
    train=False
    )

    print('flash-attn train')
    main(
    block_fn=causal_block_fn,
    data_collate_fn=collate_concat,
    train=True,
    )

    The impact of this optimization is dramatic, 51 ms for evaluation and 160 ms for training, amounting to 2.6x and 2.1x performance boosts compared to our baseline experiment.

    XFormers Memory Efficient Attention

    In our previous post we demonstrated the use of the . Here we demonstrate the use of which delivered a ~3x performance for evaluation and ~2x performance for training. We caution against deriving any conclusions from these results as the performance impact of different attention functions can vary significantly depending on the specific model and use case.

    Optimizing a HuggingFace Model for Variable-Length Input

    The tools and techniques described above are easy to implement when creating a model from scratch. However, these days it is not uncommon for ML developers to adopt existing (pretrained) models and finetune them for their use case. While the optimizations we have described can be integrated without changing the set of model weights and without altering the model behavior, it is not entirely clear what the best way to do this is. In an ideal world, our ML framework would allow us to program the use of an attention mechanism that is optimized for variable-length inputs. In this section we demonstrate how to optimize HuggingFace models for variable-length inputs.

    A Toy HuggingFace Model - GPT2LMHeadModel

    To facilitate the discussion, we create a toy example in which we train a HuggingFace based on the requested , by setting the attn_implementation parameter to “flash_attention_2”. Behind the scenes, HuggingFace will function we saw above:

    flash_config = GPT2Config(
    n_layer=DEPTH,
    n_embd=DIM,
    n_head=NUM_HEADS,
    vocab_size=NUM_TOKENS,
    attn_implementation='flash_attention_2'
    )

    print(f"HF GPT2 train with flash")
    hf_main(config=flash_config)

    The resultant time step is 620 ms, amounting to a 30% boost (in uncompiled mode) with just a simple flick of a switch.

    FlashAttention2 with Unpadded Input

    Of course, padding the sequences in the collation function only to have them unpadded, hardly seems sensible. In a recent in order to propagate the sequence . The full patch appears in the block below:

    @@ -370,0 +371 @@
    + position_ids = None
    @@ -444,0 +446 @@
    + position_ids=position_ids
    @@ -611,0 +614 @@
    + position_ids=None
    @@ -621,0 +625 @@
    + position_ids=position_ids
    @@ -1140,0 +1145 @@
    + position_ids=position_ids

    We define a collate function that concatenates our sequences and train our hugging face model on unpadded sequences. (Also see the built-in in this series as well as our was originally published in Towards Data Science on Medium, where people are continuing the conversation by highlighting and responding to this story.

    Vollständiger Original-Bericht
    Ausführliche Details, Code-Beispiele & Hersteller-Stellungnahme auf towardsdatascience.com.
    ↗ Original-Artikel auf towardsdatascience.com lesen
    Wie bewertest du diesen Beitrag?
    1 Klick Feedback
    Teilen mit Netzwerk & Team:

    Community-Analysen & Experten-Meinungen 0

    Verfasse deine eigene Analyse, teile Workarounds oder diskutiere diesen Vorfall im Blog.
    Noch keine Community-Analyse verfasst. Markiere einen Textabschnitt oder klicke oben auf Eigene Analyse verfassen“!
    Community Pulse: Relevanz-Einschätzung
    1 Klick Experten-Votum
    🔴 Akute Relevanz 0%
    🟡 In Evaluierung 0%
    🟢 Keine Auswirkung 0%
    Spannende Innovation 0%
    Verwandte Story-Cluster & Quellen (Vektor-KI)
    Port 8095 Engine
    3 Quellen
    GPT-6 Astra Release Today? OpenAI’s Next Major AI Model Is Almost Here
    1 Quelle
    Apple accuses OpenAI of destroying evidence as trade-secrets fight intensifies
    1 Quelle
    Major AI platforms go down in unprecedented simultaneous outage
    Ähnliche Beiträge
    🔍 Verwandte News

    Auch interessante Nachrichten Optimizing Transformer Models for Variable-Length Input Sequences

    Thematisch verwandte Begriffe: Optimizing, Transformer, Models, VariableLength · 6 Treffer

    Laden...

    Videos werden geladen ...

    Laden...

    Beiträge werden geladen ...

    Laden...

    Videos werden geladen ...

    Laden...

    Beiträge werden geladen ...

    Laden...

    Videos werden geladen ...

    Laden...

    Beiträge werden geladen ...

    Laden...

    Videos werden geladen ...

    Laden...

    Beiträge werden geladen ...

    Laden...

    Videos werden geladen ...