Add kv cache caching wrappers for sync pipeline transformer blocks

This commit is contained in:
ciaranbor
2026-01-06 10:51:20 +00:00
parent 546efe4dd2
commit 085358e5e0
2 changed files with 278 additions and 32 deletions
@@ -8,6 +8,8 @@ from mflux.models.flux.model.flux_transformer.transformer import Transformer
from exo.shared.types.worker.shards import PipelineShardMetadata
from exo.worker.engines.mflux.pipefusion.kv_cache import JointPatchKVCache, PatchKVCache
from exo.worker.engines.mflux.pipefusion.patched_blocks import (
CachingJointTransformerBlock,
CachingSingleTransformerBlock,
PatchedJointTransformerBlock,
PatchedSingleTransformerBlock,
)
@@ -125,6 +127,43 @@ class DistributedDenoising:
def is_last_stage(self) -> bool:
return self.rank == self.world_size - 1
def _initialize_kv_caches(
self,
batch_size: int,
text_seq_len: int,
num_img_tokens: int,
dtype: mx.Dtype,
) -> None:
"""Initialize KV caches for both sync and async pipelines.
Args:
batch_size: Batch size
text_seq_len: Length of text sequence
num_img_tokens: Number of image tokens
dtype: Data type for cache tensors
"""
self._joint_kv_caches = [
JointPatchKVCache(
batch_size=batch_size,
num_heads=24,
text_seq_len=text_seq_len,
image_seq_len=num_img_tokens,
head_dim=128,
dtype=dtype,
)
for _ in range(len(self.transformer_blocks))
]
self._single_kv_caches = [
PatchKVCache(
batch_size=batch_size,
num_heads=24,
total_seq_len=text_seq_len + num_img_tokens,
head_dim=128,
dtype=dtype,
)
for _ in range(len(self.single_transformer_blocks))
]
def _sync_pipeline(
self,
t: int,
@@ -148,7 +187,20 @@ class DistributedDenoising:
prompt_embeds, self.transformer.pos_embed, config, kontext_image_ids
)
# === PHASE 2: Joint Blocks with Communication ===
# === Initialize KV caches to populate during sync for async warmstart ===
batch_size = hidden_states.shape[0]
num_img_tokens = hidden_states.shape[1]
text_seq_len = encoder_hidden_states.shape[1]
if self._joint_kv_caches is None:
self._initialize_kv_caches(
batch_size=batch_size,
text_seq_len=text_seq_len,
num_img_tokens=num_img_tokens,
dtype=hidden_states.dtype,
)
# === PHASE 2: Joint Blocks with Communication and Caching ===
if self.has_joint_blocks:
# Receive from previous stage (if not first stage)
if not self.is_first_stage:
@@ -159,9 +211,12 @@ class DistributedDenoising:
encoder_hidden_states, self.rank - 1, group=self.group
)
# Run assigned joint blocks
for block in self.transformer_blocks:
encoder_hidden_states, hidden_states = block(
# Run assigned joint blocks with caching wrappers
for block_idx, block in enumerate(self.transformer_blocks):
caching_block = CachingJointTransformerBlock(
block, self._joint_kv_caches[block_idx]
)
encoder_hidden_states, hidden_states = caching_block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
text_embeddings=text_embeddings,
@@ -187,7 +242,7 @@ class DistributedDenoising:
mx.distributed.send(hidden_states, self.rank + 1, group=self.group)
mx.distributed.send(encoder_hidden_states, self.rank + 1, group=self.group)
# === PHASE 4: Single Blocks with Communication ===
# === PHASE 4: Single Blocks with Communication and Caching ===
if self.has_single_blocks:
# Receive from previous stage if we didn't do concatenation
if not self.owns_concat_stage and not self.is_first_stage:
@@ -199,9 +254,12 @@ class DistributedDenoising:
)
mx.eval(hidden_states)
# Run assigned single blocks
for block in self.single_transformer_blocks:
hidden_states = block(
# Run assigned single blocks with caching wrappers
for block_idx, block in enumerate(self.single_transformer_blocks):
caching_block = CachingSingleTransformerBlock(
block, self._single_kv_caches[block_idx]
)
hidden_states = caching_block(
hidden_states=hidden_states,
text_embeddings=text_embeddings,
rotary_embeddings=image_rotary_embeddings,
@@ -274,32 +332,17 @@ class DistributedDenoising:
)
token_indices = calculate_token_indices(patch_heights, latent_width, patch_size)
# === Initialize KV caches (only on first async timestep, reused across timesteps) ===
# === Initialize KV caches if not already done (reused across timesteps) ===
# This enables true PipeFusion behavior: at timestep T, patches not yet processed
# have stale K/V from timestep T-1 (not zeros)
# have stale K/V from timestep T-1. If sync pipeline ran first, caches already
# contain valid K/V from the last sync timestep.
if self._joint_kv_caches is None:
self._joint_kv_caches = [
JointPatchKVCache(
batch_size=batch_size,
num_heads=24,
text_seq_len=text_seq_len,
image_seq_len=num_img_tokens,
head_dim=128,
dtype=full_hidden.dtype,
)
for _ in range(len(self.transformer_blocks))
]
if self._single_kv_caches is None:
self._single_kv_caches = [
PatchKVCache(
batch_size=batch_size,
num_heads=24,
total_seq_len=total_seq_len,
head_dim=128,
dtype=full_hidden.dtype,
)
for _ in range(len(self.single_transformer_blocks))
]
self._initialize_kv_caches(
batch_size=batch_size,
text_seq_len=text_seq_len,
num_img_tokens=num_img_tokens,
dtype=full_hidden.dtype,
)
# Use persistent caches (stale K/V from previous timestep for unprocessed patches)
joint_kv_caches = self._joint_kv_caches
@@ -404,3 +404,206 @@ class PatchedSingleTransformerBlock:
)
return residual + hidden_states
class CachingJointTransformerBlock:
"""Joint transformer block that captures K/V for cache during sync mode.
Runs full (non-patched) attention but stores K/V in the cache for
subsequent async timesteps to use as stale values.
"""
def __init__(self, block: JointTransformerBlock, kv_cache: JointPatchKVCache):
"""Wrap an existing JointTransformerBlock with a KV cache.
Args:
block: The original JointTransformerBlock
kv_cache: KV cache to populate during forward pass
"""
self.block = block
self.kv_cache = kv_cache
def __call__(
self,
hidden_states: mx.array,
encoder_hidden_states: mx.array,
text_embeddings: mx.array,
rotary_embeddings: mx.array,
) -> tuple[mx.array, mx.array]:
"""Forward pass that also populates the KV cache.
Args:
hidden_states: Full image hidden states [B, img_len, D]
encoder_hidden_states: Full text hidden states [B, text_len, D]
text_embeddings: Time + pooled text conditioning
rotary_embeddings: Full rotary embeddings for [text + full_image]
Returns:
Tuple of (encoder_hidden_states, hidden_states) after block processing
"""
# Run standard block (computes full attention)
encoder_out, hidden_out = self.block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
text_embeddings=text_embeddings,
rotary_embeddings=rotary_embeddings,
)
# Populate KV cache for async pipeline warmstart
self._populate_cache(
hidden_states, encoder_hidden_states, text_embeddings, rotary_embeddings
)
return encoder_out, hidden_out
def _populate_cache(
self,
hidden_states: mx.array,
encoder_hidden_states: mx.array,
text_embeddings: mx.array,
rotary_embeddings: mx.array,
) -> None:
"""Compute and store K/V in cache for async pipeline warmstart."""
attn = self.block.attn
text_seq_len = encoder_hidden_states.shape[1]
num_img_tokens = hidden_states.shape[1]
# Get normalized inputs (same as what attention would see)
norm_hidden, *_ = self.block.norm1(
hidden_states=hidden_states,
text_embeddings=text_embeddings,
)
norm_encoder, *_ = self.block.norm1_context(
hidden_states=encoder_hidden_states,
text_embeddings=text_embeddings,
)
# Compute K, V for image (full, not patched)
_, img_key, img_value = AttentionUtils.process_qkv(
hidden_states=norm_hidden,
to_q=attn.to_q,
to_k=attn.to_k,
to_v=attn.to_v,
norm_q=attn.norm_q,
norm_k=attn.norm_k,
num_heads=attn.num_heads,
head_dim=attn.head_dimension,
)
# Compute K, V for text
_, txt_key, txt_value = AttentionUtils.process_qkv(
hidden_states=norm_encoder,
to_q=attn.add_q_proj,
to_k=attn.add_k_proj,
to_v=attn.add_v_proj,
norm_q=attn.norm_added_q,
norm_k=attn.norm_added_k,
num_heads=attn.num_heads,
head_dim=attn.head_dimension,
)
# Concatenate and apply RoPE
full_key = mx.concatenate([txt_key, img_key], axis=2)
full_value = mx.concatenate([txt_value, img_value], axis=2)
_, full_key = AttentionUtils.apply_rope(
xq=full_key, xk=full_key, freqs_cis=rotary_embeddings
)
# Store full sequence in cache
self.kv_cache.update_text(
key=full_key[:, :, :text_seq_len, :],
value=full_value[:, :, :text_seq_len, :],
)
self.kv_cache.update_image_patch(
patch_start=0,
patch_end=num_img_tokens,
key=full_key[:, :, text_seq_len:, :],
value=full_value[:, :, text_seq_len:, :],
)
class CachingSingleTransformerBlock:
"""Single transformer block that captures K/V for cache during sync mode.
Runs full (non-patched) attention but stores K/V in the cache for
subsequent async timesteps to use as stale values.
"""
def __init__(self, block: SingleTransformerBlock, kv_cache: PatchKVCache):
"""Wrap an existing SingleTransformerBlock with a KV cache.
Args:
block: The original SingleTransformerBlock
kv_cache: KV cache to populate during forward pass
"""
self.block = block
self.kv_cache = kv_cache
def __call__(
self,
hidden_states: mx.array,
text_embeddings: mx.array,
rotary_embeddings: mx.array,
) -> mx.array:
"""Forward pass that also populates the KV cache.
Args:
hidden_states: Full [text + image] hidden states [B, text_len + img_len, D]
text_embeddings: Time + pooled text conditioning
rotary_embeddings: Full rotary embeddings for [text + full_image]
Returns:
Output hidden states after block processing
"""
# Run standard block
hidden_out = self.block(
hidden_states=hidden_states,
text_embeddings=text_embeddings,
rotary_embeddings=rotary_embeddings,
)
# Populate KV cache for async pipeline warmstart
self._populate_cache(hidden_states, text_embeddings, rotary_embeddings)
return hidden_out
def _populate_cache(
self,
hidden_states: mx.array,
text_embeddings: mx.array,
rotary_embeddings: mx.array,
) -> None:
"""Compute and store K/V in cache for async pipeline warmstart."""
attn = self.block.attn
total_seq_len = hidden_states.shape[1]
# Get normalized inputs
norm_hidden, _ = self.block.norm(
hidden_states=hidden_states,
text_embeddings=text_embeddings,
)
# Compute K, V
_, key, value = AttentionUtils.process_qkv(
hidden_states=norm_hidden,
to_q=attn.to_q,
to_k=attn.to_k,
to_v=attn.to_v,
norm_q=attn.norm_q,
norm_k=attn.norm_k,
num_heads=attn.num_heads,
head_dim=attn.head_dimension,
)
# Apply RoPE
_, key = AttentionUtils.apply_rope(
xq=key, xk=key, freqs_cis=rotary_embeddings
)
# Store full sequence in cache
self.kv_cache.update(
patch_start=0,
patch_end=total_seq_len,
key=key,
value=value,
)