| name | streambp-efficient-backprop |
| title | StreamBP: Memory-Efficient Exact Backpropagation for Long Sequence Training of LLMs |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.03077 |
| keywords | ["training-efficiency","backpropagation","long-sequences","memory-optimization","llms"] |
| description | Enables 2.8-5.5x longer sequences during LLM training via linear decomposition of chain rule along sequence dimension, maintaining exact gradients with lower memory cost. |
StreamBP: Memory-Efficient Backpropagation
Core Concept
Training language models on long sequences requires storing activations for all tokens during backpropagation, consuming prohibitive GPU memory. StreamBP decomposes gradient computation into D sequential chunks along the sequence dimension, exploiting causal masking to compute each chunk's gradients independently. This streaming approach maintains exact backpropagation (not approximation) while reducing activation memory from O(T) to O(T/D), enabling sequences 2.8-5.5× longer than standard training with gradient checkpointing.
Architecture Overview
- Chunk-Based Decomposition: Partitions gradient computation into D sequential chunks, each processed independently
- Causal Structure Exploitation: Leverages left-to-right dependency pattern in language models to enable chunked computation
- Exact Gradients: No approximation—gradient values identical to full backpropagation
- Layer-Wise Application: Applies chunking at each transformer layer for efficient implementation
- Distributed Training: Communication-efficient variant supporting multi-GPU with DeepSpeed ZeRO
- Multiple Objectives: Works with SFT, GRPO, and DPO training objectives
Implementation
The following code demonstrates the StreamBP algorithm:
import torch
import torch.nn as nn
from typing import List, Tuple, Optional, Callable
class StreamBPGradientComputation:
"""
Streaming backpropagation with chunked gradient computation.
"""
def __init__(self, num_chunks: int = 4, gradient_accumulation_steps: int = 1):
self.num_chunks = num_chunks
self.gradient_accumulation_steps = gradient_accumulation_steps
def decompose_sequence(self, sequence_length: int) -> List[Tuple[int, int]]:
"""
Decompose sequence into D chunks for streaming backprop.
Returns list of (start, end) indices for each chunk.
"""
chunk_size = sequence_length // self.num_chunks
chunks = []
for i in range(self.num_chunks):
start = i * chunk_size
end = (i + 1) * chunk_size if i < self.num_chunks - 1 else sequence_length
chunks.append((start, end))
return chunks
def compute_chunk_gradient(self, activations: torch.Tensor,
weights: torch.Tensor,
chunk_bounds: [, ],
upstream_grad: torch.Tensor) -> [torch.Tensor, torch.Tensor, torch.Tensor]:
start, end = chunk_bounds
chunk_activations = activations[start:end]
chunk_upstream = upstream_grad[start:end]
grad_weight = torch.matmul(chunk_activations.t(), chunk_upstream)
grad_input = torch.matmul(chunk_upstream, weights.t())
grad_bias = chunk_upstream.(dim=)
grad_input, grad_weight, grad_bias
() -> [torch.Tensor, [torch.Tensor]]:
seq_len, vocab_size = logits.shape
loss_per_token = nn.functional.cross_entropy(
logits.view(-, vocab_size),
targets.view(-),
reduction=
).view(seq_len)
logits_grad = torch.zeros_like(logits)
num_chunks = (seq_len + chunk_size - ) // chunk_size
chunk_gradients = []
chunk_idx (num_chunks):
start = chunk_idx * chunk_size
end = (start + chunk_size, seq_len)
probs = torch.softmax(logits[start:end], dim=)
probs[torch.arange(end - start), targets[start:end]] -=
logits_grad[start:end] = probs / (end - start)
chunk_gradients.append(logits_grad[start:end].clone())
logits_grad, chunk_gradients
() -> [, torch.Tensor]:
seq_len = activations.shape[]
chunks = .decompose_sequence(seq_len)
weight_grads = {name: torch.zeros_like(w) name, w weights.items()}
activation_grads = torch.zeros_like(activations)
chunk_start, chunk_end chunks:
chunk_act = activations[chunk_start:chunk_end]
chunk_upstream = upstream_grad[chunk_start:chunk_end]
chunk_output = layer_fn(chunk_act, weights)
chunk_output.backward(chunk_upstream)
name, w weights.items():
w.grad :
weight_grads[name] += w.grad.clone()
w.grad.zero_()
activation_grads[chunk_start:chunk_end] = chunk_act.grad.clone() chunk_act.grad
weight_grads
:
():
.model = model
.stream_bp = StreamBPGradientComputation(num_chunks=num_chunks)
.max_seq_len = max_seq_len
.optimizer = torch.optim.AdamW(model.parameters(), lr=)
() -> :
seq_len = input_ids.shape[]
torch.enable_grad():
logits = .model(input_ids)
loss = torch.nn.functional.cross_entropy(
logits.view(-, logits.shape[-]),
target_ids.view(-)
)
loss.backward()
torch.nn.utils.clip_grad_norm_(.model.parameters(), max_norm=)
.optimizer.step()
.optimizer.zero_grad()
(loss)
() -> :
model_param_bytes = (p.numel() * p .model.parameters())
optimizer_state_bytes = model_param_bytes *
per_token_bytes = .stream_bp.num_chunks * * *
available_bytes = gpu_memory_gb * ( ** )
fixed_overhead = model_param_bytes + optimizer_state_bytes
tokens_per_batch = (available_bytes - fixed_overhead) / (batch_size * per_token_bytes)
(tokens_per_batch)
Practical Guidance
Number of Chunks: Use D=4-8 chunks for good balance. More chunks reduce memory but increase computation slightly. Fewer chunks waste memory.
Chunk Size Boundaries: Ensure chunk boundaries align with token positions. Attention masks must be applied correctly across chunk boundaries.
Causal Masking: Verify that your model uses causal (left-to-right) masking. StreamBP exploits this; bi-directional attention requires different handling.
Gradient Accumulation: StreamBP is orthogonal to gradient accumulation. Combine them: accumulate gradients over micro-batches, apply StreamBP within each batch.
Distributed Training: Use DeepSpeed ZeRO-2 compatibility for multi-GPU training. Communication happens only for synchronized optimizer steps, not intermediate chunk gradients.
Sequence Length Scheduling: Start with moderate lengths (8K tokens), then increase as training stabilizes. Longer sequences early can cause instability.
Loss Objectives: StreamBP works with SFT (cross-entropy), GRPO (policy gradient), and DPO (preference learning). Ensure loss computation respects chunk decomposition.
Reference
StreamBP enables substantial sequence length scaling:
- 2.8-5.5× longer sequences under same GPU memory as standard training
- 10-12% faster backward pass at 18K+ tokens
- 4.5× larger batch size for 8B model SFT training
- Exact gradients (no approximation error)
Empirically tested on Qwen 3 models with 80GB GPUs. The method is particularly valuable for training reasoning-heavy tasks (MATH, code) where longer sequences provide more signal for credit assignment.