[diffusion] Support SP for Krea-2 (#29777)
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
9bf122a455
commit
68901ba387
@@ -11,6 +11,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
ImagePipelineConfig,
|
||||
ModelTaskType,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_sp_parallel_rank,
|
||||
get_sp_world_size,
|
||||
)
|
||||
|
||||
# Resolution-interpolation endpoints for the time-shift `mu` (reference sampler):
|
||||
# mu is linear in image-token count between (min_res, 0.5) and (max_res, 1.15).
|
||||
@@ -74,8 +78,22 @@ class Krea2PipelineConfig(ImagePipelineConfig):
|
||||
img_ids[..., 2] = torch.arange(w_tok, device=device)[None, :]
|
||||
img_pos = img_ids.reshape(h_tok * w_tok, 3).unsqueeze(0).expand(b, -1, -1)
|
||||
txt_pos = torch.zeros(b, txt_len, 3, device=device)
|
||||
pos = torch.cat([txt_pos, img_pos], dim=1)
|
||||
|
||||
sp_world_size = get_sp_world_size()
|
||||
if sp_world_size > 1:
|
||||
# Shard the image RoPE positions to match the denoise stage's latent
|
||||
# sharding; the text prefix stays replicated (kept out of the all-to-all
|
||||
# via num_replicated_prefix). The masked path is incompatible with
|
||||
# replicated-prefix, so ragged multi-prompt batches aren't supported.
|
||||
if text_mask is not None and not bool(text_mask.all()):
|
||||
raise ValueError(
|
||||
"Krea-2 sequence parallelism does not support ragged/padded "
|
||||
"multi-prompt batches; use a single prompt or --tp-size."
|
||||
)
|
||||
img_pos = self._shard_img_pos_for_sp(img_pos, sp_world_size)
|
||||
return {"pos": torch.cat([txt_pos, img_pos], dim=1), "mask": None}
|
||||
|
||||
pos = torch.cat([txt_pos, img_pos], dim=1)
|
||||
img_mask = torch.ones(b, h_tok * w_tok, dtype=torch.bool, device=device)
|
||||
if text_mask is None:
|
||||
txt_mask = torch.ones(b, txt_len, dtype=torch.bool, device=device)
|
||||
@@ -84,6 +102,18 @@ class Krea2PipelineConfig(ImagePipelineConfig):
|
||||
mask = torch.cat([txt_mask, img_mask], dim=1)
|
||||
return {"pos": pos, "mask": mask}
|
||||
|
||||
@staticmethod
|
||||
def _shard_img_pos_for_sp(img_pos, sp_world_size):
|
||||
# This rank's contiguous image-position slice, padded to a multiple of
|
||||
# sp_world_size (mirrors the latent sharder).
|
||||
s = img_pos.shape[1]
|
||||
if s % sp_world_size != 0:
|
||||
pad = img_pos[:, -1:].repeat(1, sp_world_size - (s % sp_world_size), 1)
|
||||
img_pos = torch.cat([img_pos, pad], dim=1)
|
||||
local = img_pos.shape[1] // sp_world_size
|
||||
rank = get_sp_parallel_rank()
|
||||
return img_pos[:, rank * local : (rank + 1) * local]
|
||||
|
||||
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||
text_mask = batch.prompt_embeds_mask[0] if batch.prompt_embeds_mask else None
|
||||
return self._build_pos_and_mask(
|
||||
|
||||
@@ -978,12 +978,16 @@ class USPAttention(nn.Module):
|
||||
k_shard = _usp_input_all_to_all(k_shard, head_dim=2)
|
||||
v_shard = _usp_input_all_to_all(v_shard, head_dim=2)
|
||||
|
||||
# Q and KV can have different head counts (GQA), so slice each replicated
|
||||
# prefix by its own per-rank head shard to match the all-to-all'd suffix.
|
||||
# For MHA (kv heads == q heads) this is identical to the q shard.
|
||||
h_local = q_shard.shape[2]
|
||||
kv_h_local = k_shard.shape[2]
|
||||
h_start = sp_rank * h_local
|
||||
h_end = h_start + h_local
|
||||
q_rep = q_rep[:, :, h_start:h_end, :].contiguous()
|
||||
k_rep = k_rep[:, :, h_start:h_end, :].contiguous()
|
||||
v_rep = v_rep[:, :, h_start:h_end, :].contiguous()
|
||||
kv_h_start = sp_rank * kv_h_local
|
||||
q_rep = q_rep[:, :, h_start : h_start + h_local, :].contiguous()
|
||||
k_rep = k_rep[:, :, kv_h_start : kv_h_start + kv_h_local, :].contiguous()
|
||||
v_rep = v_rep[:, :, kv_h_start : kv_h_start + kv_h_local, :].contiguous()
|
||||
|
||||
q = torch.cat([q_rep, q_shard], dim=1)
|
||||
k = torch.cat([k_rep, k_shard], dim=1)
|
||||
|
||||
@@ -18,7 +18,10 @@ from einops import rearrange
|
||||
from torch import Tensor
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.krea2 import Krea2DitConfig
|
||||
from sglang.multimodal_gen.runtime.distributed import get_tp_world_size
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_sp_world_size,
|
||||
get_tp_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||
from sglang.multimodal_gen.runtime.layers.attention.layer import build_varlen_mask_meta
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
@@ -269,6 +272,8 @@ class Attention(nn.Module):
|
||||
freqs: Tensor | None = None,
|
||||
key_mask: Tensor | None = None,
|
||||
mask_meta: dict | None = None,
|
||||
num_replicated_prefix: int = 0,
|
||||
skip_sequence_parallel: bool = False,
|
||||
) -> Tensor:
|
||||
q, _ = self.to_q(qkv)
|
||||
k, _ = self.to_k(qkv)
|
||||
@@ -311,7 +316,13 @@ class Attention(nn.Module):
|
||||
rope_dim=hd,
|
||||
)
|
||||
out = self.attn(
|
||||
q, k, v, attn_mask=key_mask, attn_mask_meta=mask_meta
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=key_mask,
|
||||
attn_mask_meta=mask_meta,
|
||||
num_replicated_prefix=num_replicated_prefix,
|
||||
skip_sequence_parallel_override=skip_sequence_parallel,
|
||||
).flatten(2)
|
||||
else:
|
||||
q, k, v = (
|
||||
@@ -330,6 +341,8 @@ class Attention(nn.Module):
|
||||
v.transpose(1, 2).contiguous(),
|
||||
attn_mask=key_mask,
|
||||
attn_mask_meta=mask_meta,
|
||||
num_replicated_prefix=num_replicated_prefix,
|
||||
skip_sequence_parallel_override=skip_sequence_parallel,
|
||||
).flatten(2)
|
||||
out, _ = self.to_out[0](out * F.sigmoid(gate))
|
||||
return out
|
||||
@@ -371,7 +384,13 @@ class TextFusionBlock(nn.Module):
|
||||
key_mask: Tensor | None = None,
|
||||
mask_meta: dict | None = None,
|
||||
) -> Tensor:
|
||||
x = x + self.attn(self.norm1(x), key_mask=key_mask, mask_meta=mask_meta)
|
||||
# Text-fusion runs on the full replicated text, so skip the SP all-to-all.
|
||||
x = x + self.attn(
|
||||
self.norm1(x),
|
||||
key_mask=key_mask,
|
||||
mask_meta=mask_meta,
|
||||
skip_sequence_parallel=True,
|
||||
)
|
||||
x = x + self.ff(self.norm2(x))
|
||||
return x
|
||||
|
||||
@@ -450,6 +469,7 @@ class SingleStreamBlock(nn.Module):
|
||||
freqs: Tensor,
|
||||
key_mask: Tensor | None = None,
|
||||
mask_meta: dict | None = None,
|
||||
num_replicated_prefix: int = 0,
|
||||
) -> Tensor:
|
||||
mod = vec + self.scale_shift_table.reshape(-1)
|
||||
prescale, preshift, pregate, postscale, postshift, postgate = mod.chunk(
|
||||
@@ -466,6 +486,7 @@ class SingleStreamBlock(nn.Module):
|
||||
freqs,
|
||||
key_mask,
|
||||
mask_meta,
|
||||
num_replicated_prefix=num_replicated_prefix,
|
||||
)
|
||||
hidden_states = hidden_states + postgate * self.ff(
|
||||
norm_scale_shift(
|
||||
@@ -573,8 +594,18 @@ class Krea2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
combined = torch.cat((context, img), dim=1)
|
||||
freqs = self.posemb(pos)
|
||||
|
||||
# Under SP the image tokens are sharded across ranks while the text prefix
|
||||
# stays replicated; keep the leading txtlen tokens out of the all-to-all.
|
||||
num_replicated_prefix = txtlen if get_sp_world_size() > 1 else 0
|
||||
for block in self.transformer_blocks:
|
||||
combined = block(combined, tvec, freqs, joint_key, joint_meta)
|
||||
combined = block(
|
||||
combined,
|
||||
tvec,
|
||||
freqs,
|
||||
joint_key,
|
||||
joint_meta,
|
||||
num_replicated_prefix=num_replicated_prefix,
|
||||
)
|
||||
|
||||
final = self.final_layer(combined, t)
|
||||
output = final[:, txtlen : txtlen + imglen, :]
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
"""Regression test for USPAttention GQA replicated-prefix head sharding.
|
||||
|
||||
``USPAttention._forward_with_replicated_prefix`` keeps a replicated token prefix
|
||||
(e.g. text) out of the Ulysses all-to-all and slices that prefix down to the local
|
||||
head shard. For a GQA model (kv heads < q heads) the K/V prefix must be sliced by
|
||||
the *KV* head shard, not the query head shard -- otherwise the per-rank query
|
||||
offset overshoots the KV head dim, the prefix slice is empty/mismatched, and the
|
||||
``cat`` with the all-to-all'd suffix raises. MHA (kv heads == q heads) is unaffected.
|
||||
|
||||
Single-process test: the Ulysses world size, rank, all-to-all helpers, and
|
||||
all_gather are mocked so the per-rank slicing logic runs on CPU.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.layer import USPAttention
|
||||
|
||||
_LAYER = "sglang.multimodal_gen.runtime.layers.attention.layer"
|
||||
_SP = 2
|
||||
|
||||
|
||||
def _fake_input_all_to_all(x, **_):
|
||||
# Ulysses input all-to-all: gather sequence (xSP), shard heads (/SP). Only the
|
||||
# resulting shape matters for this slicing test.
|
||||
h = x.shape[2]
|
||||
return x[:, :, : h // _SP, :].repeat_interleave(_SP, dim=1).contiguous()
|
||||
|
||||
|
||||
def _fake_output_all_to_all(x, **_):
|
||||
# Inverse of the input all-to-all: shard sequence (/SP), gather heads (xSP).
|
||||
s = x.shape[1]
|
||||
return x[:, : s // _SP, :, :].repeat_interleave(_SP, dim=2).contiguous()
|
||||
|
||||
|
||||
class _CaptureAttn:
|
||||
"""Stand-in attn backend that records the q/k/v it receives."""
|
||||
|
||||
def __init__(self):
|
||||
self.q = self.k = self.v = None
|
||||
|
||||
def forward(self, q, k, v, _ctx):
|
||||
self.q, self.k, self.v = q, k, v
|
||||
return q.clone()
|
||||
|
||||
|
||||
class TestUSPAttentionReplicatedPrefix(unittest.TestCase):
|
||||
def _run(self, q_heads, kv_heads, sp_rank, num_rep=3, suffix=4, head_dim=4):
|
||||
attn = _CaptureAttn()
|
||||
obj = USPAttention.__new__(USPAttention) # bypass __init__/backend setup
|
||||
obj.attn_impl = attn
|
||||
|
||||
seq = num_rep + suffix
|
||||
q = torch.randn(1, seq, q_heads, head_dim)
|
||||
k = torch.randn(1, seq, kv_heads, head_dim)
|
||||
v = torch.randn(1, seq, kv_heads, head_dim)
|
||||
|
||||
sp_group = MagicMock()
|
||||
sp_group.ulysses_group = None
|
||||
|
||||
def fake_all_gather(out_list, tensor, **_):
|
||||
for t in out_list:
|
||||
t.copy_(tensor)
|
||||
|
||||
with (
|
||||
patch(f"{_LAYER}.get_ulysses_parallel_world_size", return_value=_SP),
|
||||
patch(f"{_LAYER}.get_sp_parallel_rank", return_value=sp_rank),
|
||||
patch(
|
||||
f"{_LAYER}._usp_input_all_to_all", side_effect=_fake_input_all_to_all
|
||||
),
|
||||
patch(
|
||||
f"{_LAYER}._usp_output_all_to_all",
|
||||
side_effect=_fake_output_all_to_all,
|
||||
),
|
||||
patch(f"{_LAYER}.get_sp_group", return_value=sp_group),
|
||||
patch("torch.distributed.all_gather", side_effect=fake_all_gather),
|
||||
):
|
||||
out = USPAttention._forward_with_replicated_prefix(
|
||||
obj, q, k, v, None, num_rep
|
||||
)
|
||||
return attn, out, q.shape
|
||||
|
||||
def test_gqa_slices_kv_prefix_by_kv_heads(self):
|
||||
# GQA: 8 query heads, 2 kv heads. The old code sliced the K/V prefix by the
|
||||
# query head shard, producing an empty/mismatched prefix and a cat error.
|
||||
for sp_rank in range(_SP):
|
||||
with self.subTest(sp_rank=sp_rank):
|
||||
attn, out, q_shape = self._run(q_heads=8, kv_heads=2, sp_rank=sp_rank)
|
||||
# q keeps q_heads/SP, k/v keep kv_heads/SP -> GQA grouping preserved.
|
||||
self.assertEqual(attn.q.shape[2], 8 // _SP)
|
||||
self.assertEqual(attn.k.shape[2], 2 // _SP)
|
||||
self.assertEqual(attn.v.shape[2], 2 // _SP)
|
||||
# prefix + all-to-all'd suffix line up on the sequence axis.
|
||||
self.assertEqual(attn.k.shape[1], attn.q.shape[1])
|
||||
# output is restored to the input layout.
|
||||
self.assertEqual(tuple(out.shape), tuple(q_shape))
|
||||
|
||||
def test_mha_prefix_unchanged(self):
|
||||
# MHA: q heads == kv heads, so the KV-shard slicing is identical to before.
|
||||
attn, out, q_shape = self._run(q_heads=8, kv_heads=8, sp_rank=1)
|
||||
self.assertEqual(attn.q.shape[2], 8 // _SP)
|
||||
self.assertEqual(attn.k.shape[2], 8 // _SP)
|
||||
self.assertEqual(tuple(out.shape), tuple(q_shape))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user