Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
79 changes: 67 additions & 12 deletions benchmarks/benchmark_training_throughput.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from transformers.optimization import get_cosine_schedule_with_warmup

import fla
from benchmarks.distributions import sample_lognormal_packed_lengths

classes = [getattr(fla.models, i) for i in fla.models.__all__]
configs = {i.model_type: i() for i in classes if issubclass(i, PretrainedConfig)}
Expand Down Expand Up @@ -66,19 +67,46 @@ def prepare_inputs(
varlen: bool,
vocab_size: int,
device: torch.device,
length_distribution: str = 'random',
num_sequences: int | None = None,
length_sigma: float = 1.0,
generator: torch.Generator | None = None,
):
if varlen:
tokens = torch.randint(high=vocab_size, size=(1, batch_size * seq_len), device=device)
cu_seqlens = torch.cat([
torch.tensor([0]),
torch.randperm(batch_size * seq_len - 16)[:torch.randint(8, 64, size=(1,))] + 16,
torch.tensor([batch_size * seq_len]),
], 0).sort()[0].to(dtype=torch.int32, device=device)
if context_len is not None:
cu_seqlens = torch.cat(
[torch.arange(i, j, context_len) for i, j in zip(cu_seqlens[:-1].tolist(), cu_seqlens[1:].tolist())] +
[torch.tensor([len(tokens[0])])],
).to(dtype=torch.int32, device=device)
total_tokens = batch_size * seq_len
tokens = torch.randint(high=vocab_size, size=(1, total_tokens), device=device)
if length_distribution == 'random':
num_cuts = int(torch.randint(8, 64, size=(1,), generator=generator))
cut_points = torch.randperm(total_tokens - 16, generator=generator)[:num_cuts] + 16
cu_seqlens = torch.cat([
torch.tensor([0]),
cut_points,
torch.tensor([total_tokens]),
], 0).sort()[0]
if context_len is not None:
cu_seqlens = torch.cat(
[torch.arange(i, j, context_len) for i, j in zip(cu_seqlens[:-1].tolist(), cu_seqlens[1:].tolist())] +
[torch.tensor([total_tokens])],
)
elif length_distribution == 'lognormal':
max_length = context_len or seq_len
min_sequences = (total_tokens + max_length - 1) // max_length
max_sequences = total_tokens // 16
if num_sequences is None:
default_sequences = min(64, max(8, batch_size * 4))
num_sequences = min(max(default_sequences, min_sequences), max_sequences)
lengths = sample_lognormal_packed_lengths(
total_tokens=total_tokens,
num_sequences=num_sequences,
max_length=max_length,
min_length=16,
sigma=length_sigma,
generator=generator,
)
cu_seqlens = torch.cat([torch.zeros(1, dtype=torch.long), lengths.cumsum(0)])
else:
raise ValueError(f"unsupported length_distribution: {length_distribution!r}")
cu_seqlens = cu_seqlens.to(dtype=torch.int32, device=device)
else:
tokens = torch.randint(high=vocab_size, size=(batch_size, seq_len), device=device)
cu_seqlens = None
Expand Down Expand Up @@ -106,8 +134,13 @@ def profile(
enable_profile: bool = False,
profile_steps: int = 64,
profile_trace: str | None = None,
length_distribution: str = 'random',
num_sequences: int | None = None,
length_sigma: float = 1.0,
seed: int = 42,
):
device = torch.device('cuda')
torch.manual_seed(seed)
config = configs[name] if name in configs else AutoConfig.from_pretrained(name)
if num_heads is not None:
if not hasattr(config, 'num_heads'):
Expand Down Expand Up @@ -149,7 +182,8 @@ def _mark(override):
_print_run_header({
'model': name,
'arch': ' '.join(arch_parts),
'data': f"B={batch_size} T={seq_len} ctx={context_len} varlen={varlen}",
'data': f"B={batch_size} T={seq_len} ctx={context_len} varlen={varlen} "
f"lengths={length_distribution} n={num_sequences} sigma={length_sigma} seed={seed}",
'training': f"{dtype} (mixed={mixed_precision}) compile={compile} "
f"warmup={warmup_steps} steps={steps}",
'profile': profile_str,
Expand Down Expand Up @@ -178,6 +212,7 @@ def _mark(override):
bar = trange(warmup_steps)

model, optimizer, scheduler = accelerator.prepare(model, optimizer, scheduler)
length_generator = torch.Generator().manual_seed(seed)
torch.cuda.synchronize(device)
for _ in bar:
# forward pass
Expand All @@ -188,6 +223,10 @@ def _mark(override):
varlen=varlen,
vocab_size=config.vocab_size,
device=device,
length_distribution=length_distribution,
num_sequences=num_sequences,
length_sigma=length_sigma,
generator=length_generator,
)
outputs = model(tokens, labels=tokens, cu_seqlens=cu_seqlens)
# backward pass
Expand All @@ -209,6 +248,10 @@ def _mark(override):
varlen=varlen,
vocab_size=config.vocab_size,
device=device,
length_distribution=length_distribution,
num_sequences=num_sequences,
length_sigma=length_sigma,
generator=length_generator,
)
outputs = model(tokens, labels=tokens, cu_seqlens=cu_seqlens)
# backward pass
Expand Down Expand Up @@ -246,6 +289,10 @@ def _mark(override):
varlen=varlen,
vocab_size=config.vocab_size,
device=device,
length_distribution=length_distribution,
num_sequences=num_sequences,
length_sigma=length_sigma,
generator=length_generator,
)
outputs = model(tokens, labels=tokens, cu_seqlens=cu_seqlens)
accelerator.backward(outputs.loss)
Expand Down Expand Up @@ -276,6 +323,10 @@ def _mark(override):
parser.add_argument("--seq_len", default=4096, type=int)
parser.add_argument("--context_len", default=None, type=int)
parser.add_argument("--varlen", action='store_true')
parser.add_argument("--length_distribution", choices=['random', 'lognormal'], default='random')
parser.add_argument("--num_sequences", default=None, type=int)
parser.add_argument("--length_sigma", default=1.0, type=float)
parser.add_argument("--seed", default=42, type=int)
parser.add_argument("--num_heads", default=None, type=int)
parser.add_argument("--head_dim", default=None, type=int)
parser.add_argument("--num_hidden_layers", default=None, type=int)
Expand All @@ -291,6 +342,10 @@ def _mark(override):
seq_len=args.seq_len,
context_len=args.context_len,
varlen=args.varlen,
length_distribution=args.length_distribution,
num_sequences=args.num_sequences,
length_sigma=args.length_sigma,
seed=args.seed,
num_heads=args.num_heads,
head_dim=args.head_dim,
num_hidden_layers=args.num_hidden_layers,
Expand Down
61 changes: 61 additions & 0 deletions benchmarks/distributions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

from __future__ import annotations

import torch


def sample_lognormal_packed_lengths(
total_tokens: int,
num_sequences: int,
max_length: int,
min_length: int = 16,
sigma: float = 1.0,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""Sample right-skewed sequence lengths with an exact packed-token budget.

The log-normal weights are a configurable workload proxy, not a claim about any particular dataset.
Pass an observed sequence count, length bounds, and sigma from the workload being measured.
"""
if total_tokens <= 0:
raise ValueError(f"total_tokens must be positive, got {total_tokens}")
if num_sequences <= 0:
raise ValueError(f"num_sequences must be positive, got {num_sequences}")
if not 0 < min_length <= max_length:
raise ValueError(f"expected 0 < min_length <= max_length, got {min_length=} and {max_length=}")
if sigma <= 0:
raise ValueError(f"sigma must be positive, got {sigma}")

min_tokens = num_sequences * min_length
max_tokens = num_sequences * max_length
if not min_tokens <= total_tokens <= max_tokens:
raise ValueError(
f"cannot pack {total_tokens} tokens into {num_sequences} sequences with "
f"lengths in [{min_length}, {max_length}]"
)

weights = torch.empty(num_sequences, dtype=torch.float64).log_normal_(
mean=0.0,
std=sigma,
generator=generator,
)
lengths = torch.full((num_sequences,), min_length, dtype=torch.long)
capacity = torch.full_like(lengths, max_length - min_length)
remaining = total_tokens - min_tokens

while remaining:
active_weights = weights.masked_fill(capacity == 0, 0)
picks = torch.multinomial(active_weights, remaining, replacement=True, generator=generator)
requested = torch.bincount(picks, minlength=num_sequences)
added = torch.minimum(requested, capacity)
lengths += added
capacity -= added
remaining -= int(added.sum())

return lengths
Loading
Loading