[diffusion] feat: support key masks on USPAttention's replicated-prefix path (#36735)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
14444c6a04
commit
f5bed255c0
@@ -147,6 +147,26 @@ def _count_active_replicated_modes(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_sdpa_mask(
|
||||||
|
mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device
|
||||||
|
) -> torch.Tensor:
|
||||||
|
mask = mask.to(device=device)
|
||||||
|
if torch.is_floating_point(mask):
|
||||||
|
mask = mask.to(dtype=dtype)
|
||||||
|
if mask.dim() == 2:
|
||||||
|
mask = mask[:, None, None, :]
|
||||||
|
elif mask.dim() == 3:
|
||||||
|
mask = mask[:, None, :, :]
|
||||||
|
return mask
|
||||||
|
|
||||||
|
mask = mask.to(dtype=dtype)
|
||||||
|
if mask.dim() == 2:
|
||||||
|
mask = mask[:, None, None, :]
|
||||||
|
elif mask.dim() == 3:
|
||||||
|
mask = mask[:, None, :, :]
|
||||||
|
return (mask - 1.0) * torch.finfo(dtype).max
|
||||||
|
|
||||||
|
|
||||||
def build_varlen_mask_meta(
|
def build_varlen_mask_meta(
|
||||||
key_mask: torch.Tensor,
|
key_mask: torch.Tensor,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
@@ -905,6 +925,15 @@ class USPAttention(nn.Module):
|
|||||||
out = self.attn_impl.postprocess_output(out, ctx_attn_metadata)
|
out = self.attn_impl.postprocess_output(out, ctx_attn_metadata)
|
||||||
return _usp_output_all_to_all_varlen(out, seq_lens, head_dim=2)
|
return _usp_output_all_to_all_varlen(out, seq_lens, head_dim=2)
|
||||||
|
|
||||||
|
# Under BCG each attention break point snapshots its own kwargs, so
|
||||||
|
# per-segment mask buffers differ while this meta object is shared by
|
||||||
|
# every segment; the replicated-prefix path memoizes its gathered
|
||||||
|
# joint mask on it (see _forward_with_replicated_prefix).
|
||||||
|
joint_mask_memo_host = (
|
||||||
|
attn_mask_meta
|
||||||
|
if isinstance(attn_mask_meta, DynamicVarlenMaskMeta)
|
||||||
|
else attn_mask
|
||||||
|
)
|
||||||
if isinstance(attn_mask_meta, DynamicVarlenMaskMeta):
|
if isinstance(attn_mask_meta, DynamicVarlenMaskMeta):
|
||||||
attn_mask_meta = attn_mask_meta.resolve(attn_mask)
|
attn_mask_meta = attn_mask_meta.resolve(attn_mask)
|
||||||
|
|
||||||
@@ -975,36 +1004,41 @@ class USPAttention(nn.Module):
|
|||||||
and not effective_skip_sp
|
and not effective_skip_sp
|
||||||
and get_sequence_parallel_world_size() > 1
|
and get_sequence_parallel_world_size() > 1
|
||||||
):
|
):
|
||||||
# Under SP this path shards every row through the all-to-all;
|
# A replicated prefix with a [B, S] key mask is served by the
|
||||||
# a replicated prefix/suffix would be duplicated across ranks
|
# replicated-prefix flow: after its all-to-all the attention is
|
||||||
# and silently corrupt the output, so refuse loudly instead.
|
# rank-local over the full sequence, so the mask (identical on
|
||||||
# On a single rank the mask already describes the full
|
# every rank for the prefix, gathered for the suffix) applies
|
||||||
# sequence and the replicated counts are meaningless, so the
|
# exactly as on a single rank.
|
||||||
# call is legal.
|
masked_prefix_supported = (
|
||||||
raise NotImplementedError(
|
num_replicated_prefix > 0
|
||||||
"USPAttention's masked path does not support replicated "
|
and not num_replicated_suffix
|
||||||
"prefix/suffix tokens under sequence parallelism; drop "
|
and not num_replicated_kv_prefix
|
||||||
"attn_mask/attn_mask_meta or the replicated segment."
|
and attn_mask is not None
|
||||||
|
and attn_mask.dim() == 2
|
||||||
|
and not qkv_pre_all_to_all
|
||||||
|
and get_ring_parallel_world_size() == 1
|
||||||
|
)
|
||||||
|
if not masked_prefix_supported:
|
||||||
|
# Sharding every row through the all-to-all would duplicate
|
||||||
|
# a replicated segment across ranks and silently corrupt
|
||||||
|
# the output, so refuse loudly. On a single rank the mask
|
||||||
|
# already describes the full sequence and the replicated
|
||||||
|
# counts are meaningless, so the call is legal.
|
||||||
|
raise NotImplementedError(
|
||||||
|
"USPAttention's masked path does not support replicated "
|
||||||
|
"prefix/suffix tokens under sequence parallelism; drop "
|
||||||
|
"attn_mask/attn_mask_meta or the replicated segment."
|
||||||
|
)
|
||||||
|
return self._forward_with_replicated_prefix(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
ctx_attn_metadata,
|
||||||
|
num_replicated_prefix,
|
||||||
|
attn_mask=attn_mask,
|
||||||
|
drop_masked_query_rows=attn_mask_meta is not None,
|
||||||
|
joint_mask_memo_host=joint_mask_memo_host,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _prepare_sdpa_mask(
|
|
||||||
mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device
|
|
||||||
) -> torch.Tensor:
|
|
||||||
mask = mask.to(device=device)
|
|
||||||
if torch.is_floating_point(mask):
|
|
||||||
mask = mask.to(dtype=dtype)
|
|
||||||
if mask.dim() == 2:
|
|
||||||
mask = mask[:, None, None, :]
|
|
||||||
elif mask.dim() == 3:
|
|
||||||
mask = mask[:, None, :, :]
|
|
||||||
return mask
|
|
||||||
|
|
||||||
mask = mask.to(dtype=dtype)
|
|
||||||
if mask.dim() == 2:
|
|
||||||
mask = mask[:, None, None, :]
|
|
||||||
elif mask.dim() == 3:
|
|
||||||
mask = mask[:, None, :, :]
|
|
||||||
return (mask - 1.0) * torch.finfo(dtype).max
|
|
||||||
|
|
||||||
sp_world_size = get_sequence_parallel_world_size()
|
sp_world_size = get_sequence_parallel_world_size()
|
||||||
if effective_skip_sp or sp_world_size == 1:
|
if effective_skip_sp or sp_world_size == 1:
|
||||||
@@ -1596,6 +1630,9 @@ class USPAttention(nn.Module):
|
|||||||
v: torch.Tensor,
|
v: torch.Tensor,
|
||||||
ctx_attn_metadata,
|
ctx_attn_metadata,
|
||||||
num_rep: int,
|
num_rep: int,
|
||||||
|
attn_mask: torch.Tensor | None = None,
|
||||||
|
drop_masked_query_rows: bool = False,
|
||||||
|
joint_mask_memo_host=None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Ulysses attention where the first *num_rep* tokens are replicated
|
"""Ulysses attention where the first *num_rep* tokens are replicated
|
||||||
across SP ranks (e.g. text tokens) and should NOT be duplicated by the
|
across SP ranks (e.g. text tokens) and should NOT be duplicated by the
|
||||||
@@ -1607,6 +1644,11 @@ class USPAttention(nn.Module):
|
|||||||
3. Locally slice the replicated prefix to the same head shard.
|
3. Locally slice the replicated prefix to the same head shard.
|
||||||
4. Concatenate [prefix_h_local, gathered_suffix] and run attention.
|
4. Concatenate [prefix_h_local, gathered_suffix] and run attention.
|
||||||
5. Split output, all-to-all back the suffix, all-gather prefix heads.
|
5. Split output, all-to-all back the suffix, all-gather prefix heads.
|
||||||
|
|
||||||
|
``attn_mask`` is a [B, num_rep + S_local] key mask in the caller's
|
||||||
|
local layout: its prefix columns are replicated (identical on every
|
||||||
|
rank) and its suffix columns are gathered to match step 2's sequence
|
||||||
|
gather, so step 4 can apply it exactly as a single rank would.
|
||||||
"""
|
"""
|
||||||
sp_size = get_ulysses_parallel_world_size()
|
sp_size = get_ulysses_parallel_world_size()
|
||||||
u_rank = get_ulysses_parallel_rank()
|
u_rank = get_ulysses_parallel_rank()
|
||||||
@@ -1619,6 +1661,33 @@ class USPAttention(nn.Module):
|
|||||||
k_shard = _usp_input_all_to_all(k_shard, head_dim=2)
|
k_shard = _usp_input_all_to_all(k_shard, head_dim=2)
|
||||||
v_shard = _usp_input_all_to_all(v_shard, head_dim=2)
|
v_shard = _usp_input_all_to_all(v_shard, head_dim=2)
|
||||||
|
|
||||||
|
joint_mask = joint_mask_meta = None
|
||||||
|
if attn_mask is not None:
|
||||||
|
cache_key = (
|
||||||
|
num_rep,
|
||||||
|
sp_size,
|
||||||
|
tuple(attn_mask.shape),
|
||||||
|
get_current_replay_token(),
|
||||||
|
)
|
||||||
|
host = (
|
||||||
|
joint_mask_memo_host if joint_mask_memo_host is not None else attn_mask
|
||||||
|
)
|
||||||
|
cached = getattr(host, "_sglang_usp_joint_mask_cache", None)
|
||||||
|
if cached is not None and cached[0] == cache_key:
|
||||||
|
_, joint_mask, joint_mask_meta = cached
|
||||||
|
else:
|
||||||
|
mask_suffix = sequence_model_parallel_all_gather(
|
||||||
|
attn_mask[:, num_rep:].contiguous(), dim=1
|
||||||
|
)
|
||||||
|
joint_mask = torch.cat([attn_mask[:, :num_rep], mask_suffix], dim=1)
|
||||||
|
if drop_masked_query_rows:
|
||||||
|
joint_mask_meta = build_varlen_mask_meta(joint_mask)
|
||||||
|
host._sglang_usp_joint_mask_cache = (
|
||||||
|
cache_key,
|
||||||
|
joint_mask,
|
||||||
|
joint_mask_meta,
|
||||||
|
)
|
||||||
|
|
||||||
# Q and KV can have different head counts (GQA), so slice each replicated
|
# 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.
|
# 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.
|
# For MHA (kv heads == q heads) this is identical to the q shard.
|
||||||
@@ -1632,7 +1701,14 @@ class USPAttention(nn.Module):
|
|||||||
|
|
||||||
q = torch.cat([q_rep, q_shard], dim=1)
|
q = torch.cat([q_rep, q_shard], dim=1)
|
||||||
out = self._replicated_kv_attention(
|
out = self._replicated_kv_attention(
|
||||||
q, k_shard, v_shard, k_rep, v_rep, ctx_attn_metadata
|
q,
|
||||||
|
k_shard,
|
||||||
|
v_shard,
|
||||||
|
k_rep,
|
||||||
|
v_rep,
|
||||||
|
ctx_attn_metadata,
|
||||||
|
attn_mask=joint_mask,
|
||||||
|
attn_mask_meta=joint_mask_meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
out_rep = out[:, :num_rep]
|
out_rep = out[:, :num_rep]
|
||||||
@@ -1660,6 +1736,8 @@ class USPAttention(nn.Module):
|
|||||||
v_rep: torch.Tensor,
|
v_rep: torch.Tensor,
|
||||||
ctx_attn_metadata,
|
ctx_attn_metadata,
|
||||||
rep_first: bool = True,
|
rep_first: bool = True,
|
||||||
|
attn_mask: torch.Tensor | None = None,
|
||||||
|
attn_mask_meta: dict | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Attention of q against replicated + ring-sharded KV.
|
"""Attention of q against replicated + ring-sharded KV.
|
||||||
|
|
||||||
@@ -1669,8 +1747,12 @@ class USPAttention(nn.Module):
|
|||||||
parallelism the sharded KV rotates around the ring while the
|
parallelism the sharded KV rotates around the ring while the
|
||||||
replicated KV contributes one extra local partial, LSE-merged with
|
replicated KV contributes one extra local partial, LSE-merged with
|
||||||
the ring result (exact up to float reordering).
|
the ring result (exact up to float reordering).
|
||||||
|
|
||||||
|
``attn_mask`` is a [B, S] key mask over the concatenated KV order and
|
||||||
|
is only supported on the non-ring path (the caller gates ring out).
|
||||||
"""
|
"""
|
||||||
if get_ring_parallel_world_size() > 1:
|
if get_ring_parallel_world_size() > 1:
|
||||||
|
assert attn_mask is None, "masked replicated KV does not support ring"
|
||||||
out_ring, lse_ring = ring_attn(
|
out_ring, lse_ring = ring_attn(
|
||||||
q, k_shard, v_shard, self.attn_impl, return_softmax_lse=True
|
q, k_shard, v_shard, self.attn_impl, return_softmax_lse=True
|
||||||
)
|
)
|
||||||
@@ -1689,8 +1771,76 @@ class USPAttention(nn.Module):
|
|||||||
)
|
)
|
||||||
k = torch.cat(kv_parts[0], dim=1)
|
k = torch.cat(kv_parts[0], dim=1)
|
||||||
v = torch.cat(kv_parts[1], dim=1)
|
v = torch.cat(kv_parts[1], dim=1)
|
||||||
|
if attn_mask is not None:
|
||||||
|
assert rep_first, "the joint key mask is built in [rep, shard] order"
|
||||||
|
return self._masked_local_attention(
|
||||||
|
q, k, v, attn_mask, attn_mask_meta=attn_mask_meta
|
||||||
|
)
|
||||||
return self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
return self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||||
|
|
||||||
|
def _masked_local_attention(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
attn_mask: torch.Tensor,
|
||||||
|
attn_mask_meta: dict | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Rank-local attention with a [B, S] key mask over full-sequence
|
||||||
|
q/k/v (heads may already be sharded). Mirrors the single-rank masked
|
||||||
|
path: varlen FA with precomputed meta when the caller opted into
|
||||||
|
dropped-query-row semantics (masked rows zero-filled on output), SDPA
|
||||||
|
with an additive mask otherwise.
|
||||||
|
"""
|
||||||
|
if (
|
||||||
|
_VARLEN_FA_ENABLED
|
||||||
|
and attn_mask_meta is not None
|
||||||
|
and self.backend == AttentionBackendEnum.FA
|
||||||
|
and attn_mask.dtype in (torch.bool, torch.uint8, torch.int32, torch.int64)
|
||||||
|
and q.device.type == "cuda"
|
||||||
|
and attn_mask.device == q.device
|
||||||
|
and q.dtype in (torch.float16, torch.bfloat16)
|
||||||
|
and q.shape[:2] == attn_mask.shape == k.shape[:2] == v.shape[:2]
|
||||||
|
):
|
||||||
|
bs, seq = q.shape[0], q.shape[1]
|
||||||
|
meta = attn_mask_meta
|
||||||
|
indices = meta["indices"]
|
||||||
|
if indices.shape[0] > 0:
|
||||||
|
q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices)
|
||||||
|
out_unpad = flash_attn_varlen_func(
|
||||||
|
q=q_unpad,
|
||||||
|
k=k_unpad,
|
||||||
|
v=v_unpad,
|
||||||
|
cu_seqlens_q=meta["cu_seqlens"],
|
||||||
|
cu_seqlens_k=meta["cu_seqlens"],
|
||||||
|
max_seqlen_q=meta["max_seqlen"],
|
||||||
|
max_seqlen_k=meta["max_seqlen"],
|
||||||
|
softmax_scale=self.softmax_scale,
|
||||||
|
causal=False,
|
||||||
|
ver=_fa_backend.fa_ver,
|
||||||
|
)
|
||||||
|
return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq)
|
||||||
|
|
||||||
|
q_ = q.transpose(1, 2)
|
||||||
|
k_ = k.transpose(1, 2)
|
||||||
|
v_ = v.transpose(1, 2)
|
||||||
|
mask = _prepare_sdpa_mask(attn_mask, dtype=q_.dtype, device=q_.device)
|
||||||
|
sdpa_context = (
|
||||||
|
sdpa_kernel(_PYTORCH_DEFAULT_CUDA_SDP_BACKENDS)
|
||||||
|
if self.allow_cudnn_sdp and q_.device.type == "cuda"
|
||||||
|
else nullcontext()
|
||||||
|
)
|
||||||
|
with sdpa_context:
|
||||||
|
return torch.nn.functional.scaled_dot_product_attention(
|
||||||
|
q_,
|
||||||
|
k_,
|
||||||
|
v_,
|
||||||
|
attn_mask=mask,
|
||||||
|
dropout_p=0.0,
|
||||||
|
is_causal=False,
|
||||||
|
scale=self.softmax_scale,
|
||||||
|
).transpose(1, 2)
|
||||||
|
|
||||||
def forward_with_replicated_kv_prefix(
|
def forward_with_replicated_kv_prefix(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
|
|||||||
@@ -67,7 +67,8 @@ class TestUSPAttentionReplicatedPrefix(unittest.TestCase):
|
|||||||
|
|
||||||
with (
|
with (
|
||||||
patch(f"{_LAYER}.get_ulysses_parallel_world_size", return_value=_SP),
|
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}.get_ulysses_parallel_rank", return_value=sp_rank),
|
||||||
|
patch(f"{_LAYER}.get_ring_parallel_world_size", return_value=1),
|
||||||
patch(
|
patch(
|
||||||
f"{_LAYER}._usp_input_all_to_all", side_effect=_fake_input_all_to_all
|
f"{_LAYER}._usp_input_all_to_all", side_effect=_fake_input_all_to_all
|
||||||
),
|
),
|
||||||
@@ -153,3 +154,155 @@ class TestUSPAttentionMaskedReplicatedGuard(unittest.TestCase):
|
|||||||
):
|
):
|
||||||
out = obj.forward(q, q, q, attn_mask=mask, num_replicated_prefix=2)
|
out = obj.forward(q, q, q, attn_mask=mask, num_replicated_prefix=2)
|
||||||
self.assertEqual(out.shape, q.shape)
|
self.assertEqual(out.shape, q.shape)
|
||||||
|
|
||||||
|
|
||||||
|
class TestUSPAttentionMaskedReplicatedPrefix(unittest.TestCase):
|
||||||
|
"""A [B, S] key mask with a replicated prefix now rides the
|
||||||
|
replicated-prefix flow under ulysses SP instead of raising: after the
|
||||||
|
all-to-all the attention is rank-local over the full sequence, so the
|
||||||
|
output must match a plain full-sequence masked SDPA reference."""
|
||||||
|
|
||||||
|
_H, _D = 4, 8
|
||||||
|
_NUM_REP, _S_LOCAL, _S_OTHER = 3, 4, 4
|
||||||
|
|
||||||
|
def _run_rank0(self, q_full, k_full, v_full, text_mask, other_suffix_mask):
|
||||||
|
|
||||||
|
num_rep, s_local = self._NUM_REP, self._S_LOCAL
|
||||||
|
h_local = self._H // _SP
|
||||||
|
obj = USPAttention.__new__(USPAttention)
|
||||||
|
obj.attn_impl = _CaptureAttn()
|
||||||
|
obj.skip_sequence_parallel = False
|
||||||
|
obj.sp_attention_mode = "ulysses"
|
||||||
|
obj.sp_attention_mode_is_auto = True
|
||||||
|
obj.backend = None # not FA -> the masked branch takes SDPA
|
||||||
|
obj.softmax_scale = self._D**-0.5
|
||||||
|
obj.allow_cudnn_sdp = False
|
||||||
|
obj.enable_packed_qkv_input_a2a = False
|
||||||
|
|
||||||
|
local = slice(0, num_rep + s_local)
|
||||||
|
q, k, v = q_full[:, local], k_full[:, local], v_full[:, local]
|
||||||
|
other = {
|
||||||
|
"suffix": lambda full: full[:, num_rep + s_local :],
|
||||||
|
}
|
||||||
|
gathered = iter(
|
||||||
|
[
|
||||||
|
torch.cat([q[:, num_rep:], other["suffix"](q_full)], dim=1),
|
||||||
|
torch.cat([k[:, num_rep:], other["suffix"](k_full)], dim=1),
|
||||||
|
torch.cat([v[:, num_rep:], other["suffix"](v_full)], dim=1),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
def fake_input_a2a(x, **_):
|
||||||
|
# rank0 view: gather sequence across ranks, keep the first head shard.
|
||||||
|
return next(gathered)[:, :, :h_local, :].contiguous()
|
||||||
|
|
||||||
|
def fake_output_a2a(x, **_):
|
||||||
|
# inverse: keep rank0's sequence shard, duplicate the head shard.
|
||||||
|
return x[:, :s_local].repeat(1, 1, _SP, 1)
|
||||||
|
|
||||||
|
def fake_mask_gather(m, **_):
|
||||||
|
return torch.cat([m, other_suffix_mask], dim=1)
|
||||||
|
|
||||||
|
def fake_all_gather(out_list, tensor, **_):
|
||||||
|
for t in out_list:
|
||||||
|
t.copy_(tensor)
|
||||||
|
|
||||||
|
sp_group = MagicMock()
|
||||||
|
sp_group.ulysses_group = None
|
||||||
|
local_mask = torch.cat(
|
||||||
|
[text_mask, torch.ones(1, s_local, dtype=torch.bool)], dim=1
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
f"{_LAYER}.get_forward_context",
|
||||||
|
return_value=MagicMock(attn_metadata=None),
|
||||||
|
),
|
||||||
|
patch(f"{_LAYER}.get_sequence_parallel_world_size", return_value=_SP),
|
||||||
|
patch(f"{_LAYER}.get_ulysses_parallel_world_size", return_value=_SP),
|
||||||
|
patch(f"{_LAYER}.get_ring_parallel_world_size", return_value=1),
|
||||||
|
patch(f"{_LAYER}.get_ulysses_parallel_rank", return_value=0),
|
||||||
|
patch(f"{_LAYER}._usp_input_all_to_all", side_effect=fake_input_a2a),
|
||||||
|
patch(f"{_LAYER}._usp_output_all_to_all", side_effect=fake_output_a2a),
|
||||||
|
patch(
|
||||||
|
f"{_LAYER}.sequence_model_parallel_all_gather",
|
||||||
|
side_effect=fake_mask_gather,
|
||||||
|
),
|
||||||
|
patch(f"{_LAYER}.get_sp_group", return_value=sp_group),
|
||||||
|
patch("torch.distributed.all_gather", side_effect=fake_all_gather),
|
||||||
|
):
|
||||||
|
return obj.forward(
|
||||||
|
q, k, v, attn_mask=local_mask, num_replicated_prefix=num_rep
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_matches_full_sequence_masked_sdpa(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
seq = self._NUM_REP + self._S_LOCAL + self._S_OTHER
|
||||||
|
q = torch.randn(1, seq, self._H, self._D)
|
||||||
|
k = torch.randn(1, seq, self._H, self._D)
|
||||||
|
v = torch.randn(1, seq, self._H, self._D)
|
||||||
|
text_mask = torch.tensor([[True, True, False]])
|
||||||
|
other_suffix_mask = torch.ones(1, self._S_OTHER, dtype=torch.bool)
|
||||||
|
|
||||||
|
out = self._run_rank0(q, k, v, text_mask, other_suffix_mask)
|
||||||
|
|
||||||
|
full_mask = torch.cat(
|
||||||
|
[text_mask, torch.ones(1, self._S_LOCAL + self._S_OTHER, dtype=torch.bool)],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
|
ref = torch.nn.functional.scaled_dot_product_attention(
|
||||||
|
q.transpose(1, 2),
|
||||||
|
k.transpose(1, 2),
|
||||||
|
v.transpose(1, 2),
|
||||||
|
attn_mask=full_mask[:, None, None, :],
|
||||||
|
scale=self._D**-0.5,
|
||||||
|
).transpose(1, 2)
|
||||||
|
|
||||||
|
h_local = self._H // _SP
|
||||||
|
rows = self._NUM_REP + self._S_LOCAL
|
||||||
|
self.assertEqual(tuple(out.shape), (1, rows, self._H, self._D))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
out[:, :, :h_local, :],
|
||||||
|
ref[:, :rows, :h_local, :],
|
||||||
|
rtol=1e-5,
|
||||||
|
atol=1e-5,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_masked_prefix_still_rejects_ring_and_3d_masks(self):
|
||||||
|
obj = USPAttention.__new__(USPAttention)
|
||||||
|
obj.attn_impl = _CaptureAttn()
|
||||||
|
obj.skip_sequence_parallel = False
|
||||||
|
obj.sp_attention_mode = "ulysses"
|
||||||
|
obj.sp_attention_mode_is_auto = True
|
||||||
|
q = torch.randn(1, 6, 2, 4)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
f"{_LAYER}.get_forward_context",
|
||||||
|
return_value=MagicMock(attn_metadata=None),
|
||||||
|
),
|
||||||
|
patch(f"{_LAYER}.get_sequence_parallel_world_size", return_value=2),
|
||||||
|
patch(f"{_LAYER}.get_ring_parallel_world_size", return_value=2),
|
||||||
|
):
|
||||||
|
with self.assertRaisesRegex(NotImplementedError, "replicated"):
|
||||||
|
obj.forward(
|
||||||
|
q,
|
||||||
|
q,
|
||||||
|
q,
|
||||||
|
attn_mask=torch.ones(1, 6, dtype=torch.bool),
|
||||||
|
num_replicated_prefix=2,
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
f"{_LAYER}.get_forward_context",
|
||||||
|
return_value=MagicMock(attn_metadata=None),
|
||||||
|
),
|
||||||
|
patch(f"{_LAYER}.get_sequence_parallel_world_size", return_value=2),
|
||||||
|
patch(f"{_LAYER}.get_ring_parallel_world_size", return_value=1),
|
||||||
|
):
|
||||||
|
with self.assertRaisesRegex(NotImplementedError, "replicated"):
|
||||||
|
obj.forward(
|
||||||
|
q,
|
||||||
|
q,
|
||||||
|
q,
|
||||||
|
attn_mask=torch.ones(1, 6, 6, dtype=torch.bool),
|
||||||
|
num_replicated_prefix=2,
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user