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
70 changes: 48 additions & 22 deletions cosmos_framework/model/generator/mot/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,8 +353,15 @@ def multi_control_two_way_attention(
- ctrl_i output: from pass i only
- noisy output: w_1 * noisy_out_1 + ... + w_N * noisy_out_N (weighted sum)

All SDPA calls are maskless → Flash Attention is always active.
N=1, w=1.0 → identical to ``two_way_attention``.
All SDPA calls are maskless. N=1, w=1.0 → identical to ``two_way_attention``.

Dense vs varlen:
Single-sample inference uses the dense attention API (no cumulative/max
seqlen args), matching ``two_way_attention``. That lets the chooser keep
cuDNN on sm_120 (workstation Blackwell), where Cosmos prefers cuDNN but
the current cuDNN integration rejects varlen and would otherwise fall back
to NATTEN kernels that lack sm_120 images. Training / multi-sample packs
keep the varlen path.

Padding safety:
Both ``get_causal_seq`` and ``get_full_only_seq`` can return padded rows.
Expand All @@ -374,22 +381,34 @@ def multi_control_two_way_attention(
noisy_s, noisy_e = split_info.noisy_token_range
weights = split_info.control_weights

# ── 1. Text self-attention (causal, unchanged) ───────────────────────────
# ── 1. Text self-attention (causal) ──────────────────────────────────────
causal_q, causal_q_offsets = get_causal_seq(packed_query_states)
causal_k, causal_k_offsets = get_causal_seq(packed_key_states)
causal_v, _ = get_causal_seq(packed_value_states)

# Mirror two_way_attention: dense FMHA for single-sample inference so cuDNN
# remains eligible on sm_120; varlen only when training / multi-sample.
sample_offsets = packed_query_states["sample_offsets"]
num_samples = sample_offsets.shape[0] - 1
use_dense = num_samples == 1 and not torch.is_grad_enabled()

use_dont_care_mask = causal_q_offsets is causal_k_offsets
if use_dense:
causal_varlen_kwargs = {}
else:
causal_varlen_kwargs = dict(
cumulative_seqlen_Q=causal_q_offsets,
cumulative_seqlen_KV=causal_k_offsets,
max_seqlen_Q=packed_query_states["max_causal_len"],
max_seqlen_KV=packed_query_states["max_causal_len"],
)
causal_res = attention(
causal_q.unsqueeze(0),
causal_k.unsqueeze(0),
causal_v.unsqueeze(0),
cumulative_seqlen_Q=causal_q_offsets,
cumulative_seqlen_KV=causal_k_offsets,
max_seqlen_Q=packed_query_states["max_causal_len"],
max_seqlen_KV=packed_query_states["max_causal_len"],
is_causal=True,
causal_type=CausalType.DontCare if use_dont_care_mask else CausalType.TopLeft,
**causal_varlen_kwargs,
)
causal_out = causal_res.squeeze(0).flatten(-2, -1) # [N_text, Hq*D]

Expand Down Expand Up @@ -432,30 +451,37 @@ def _sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
torch._check(k.shape[0] == v.shape[0])
n_q, n_kv = q.shape[0], k.shape[0]
# These lengths come from data-dependent unpadding, so they are unbacked
# symints under torch.compile. The selected attention backend (NATTEN)
# validates varlen inputs with `max_seqlen == 0` / `max_seqlen < 1`
# guards; without a positivity fact Dynamo cannot discharge `Eq(n, 0)`.
# Every control/noisy segment always has at least one token, so assert it.
# symints under torch.compile. Varlen backends validate with
# `max_seqlen == 0` / `max_seqlen < 1` guards; without a positivity fact
# Dynamo cannot discharge `Eq(n, 0)`. Every control/noisy segment always
# has at least one token, so assert it (needed for the varlen path).
torch._check(n_q > 0)
torch._check(n_kv > 0)
# Pass cumulative_seqlen_{Q,KV} + max_seqlen_{Q,KV} directly instead of
# Dense (inference, single sample): omit varlen args so cuDNN stays
# eligible on sm_120. Varlen (training / multi-sample): pass
# cumulative_seqlen_{Q,KV} + max_seqlen_{Q,KV} directly instead of
# seqlens_{Q,KV}. The frontend derives cumulative offsets from seqlens via
# `generate_varlen_parameters`, which calls `.max().item()` (a device-host
# sync) and is explicitly disallowed inside a torch.compile region. Each
# pass here is a single (batch=1) packed sequence, so the cumulative
# offsets are simply [0, n]. Building them ourselves keeps the whole path
# inside the compiled graph.
zero = torch.zeros(1, dtype=torch.int32, device=q.device)
cu_seqlens_q = torch.cat([zero, torch.tensor([n_q], dtype=torch.int32, device=q.device)])
cu_seqlens_kv = torch.cat([zero, torch.tensor([n_kv], dtype=torch.int32, device=q.device)])
# pass here is a single packed sequence, so the cumulative offsets are
# simply [0, n]. Building them ourselves keeps the path compiled.
if use_dense:
varlen_kwargs = {}
else:
zero = torch.zeros(1, dtype=torch.int32, device=q.device)
cu_seqlens_q = torch.cat([zero, torch.tensor([n_q], dtype=torch.int32, device=q.device)])
cu_seqlens_kv = torch.cat([zero, torch.tensor([n_kv], dtype=torch.int32, device=q.device)])
varlen_kwargs = dict(
cumulative_seqlen_Q=cu_seqlens_q,
cumulative_seqlen_KV=cu_seqlens_kv,
max_seqlen_Q=n_q,
max_seqlen_KV=n_kv,
)
res = attention(
q.unsqueeze(0), # [1, N_q, Hq, D]
k.unsqueeze(0), # [1, N_kv, Hkv, D]
v.unsqueeze(0), # [1, N_kv, Hkv, D]
cumulative_seqlen_Q=cu_seqlens_q,
cumulative_seqlen_KV=cu_seqlens_kv,
max_seqlen_Q=n_q,
max_seqlen_KV=n_kv,
**varlen_kwargs,
) # [1, N_q, Hq, D]
return res.squeeze(0).flatten(-2, -1) # [N_q, Hq*D]

Expand Down
56 changes: 56 additions & 0 deletions cosmos_framework/model/generator/mot/attention_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,6 +400,62 @@ def fake_two_way_attention(*args: object, **kwargs: object) -> object:
assert kv_to_store is None


@pytest.mark.L0
@pytest.mark.CPU
def test_multi_control_matches_two_way_dense_inference_convention(monkeypatch: pytest.MonkeyPatch) -> None:
"""Single-sample inference is dense; grad-enabled execution remains varlen."""

def make_pack(x: torch.Tensor):
return build_packed_sequence(
"two_way",
packed_sequence=x,
attn_modes=["causal", "full"],
split_lens=[2, 4],
sample_lens=[6],
packed_und_token_indexes=torch.tensor([0, 1], dtype=torch.long),
packed_gen_token_indexes=torch.tensor([2, 3, 4, 5], dtype=torch.long),
num_heads=1,
head_dim=8,
num_layers=1,
)

q_pack, split_info, _ = make_pack(torch.arange(48, dtype=torch.float32).reshape(6, 1, 8))
raw_k_pack, _, _ = make_pack(torch.zeros(6, 1, 8))
v_pack, _, _ = make_pack(torch.ones(6, 1, 8))
split_info.control_stream_token_ranges = [(0, 2)]
split_info.noisy_token_range = (2, 4)
split_info.control_weights = [1.0]

calls: list[dict[str, object]] = []

def fake_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, **kwargs: object) -> torch.Tensor:
calls.append(kwargs)
return q

monkeypatch.setattr(attention, "attention", fake_attention)

with torch.no_grad():
attention.multi_control_two_way_attention(
q_pack,
raw_k_pack,
v_pack,
split_info,
)

assert len(calls) == 3 # causal text, control, noisy
assert all("cumulative_seqlen_Q" not in call for call in calls)

calls.clear()
attention.multi_control_two_way_attention(
q_pack,
raw_k_pack,
v_pack,
split_info,
)
assert len(calls) == 3
assert all("cumulative_seqlen_Q" in call for call in calls)


@pytest.mark.L0
def test_dispatch_attention_rejects_incomplete_split_info() -> None:
foreign_split_info = _foreign_split_info()
Expand Down
4 changes: 2 additions & 2 deletions cosmos_framework/model/generator/mot/cosmos3_vfm_network.py
Original file line number Diff line number Diff line change
Expand Up @@ -1010,8 +1010,8 @@ def forward(
# one per control. For each pass i, KV = [text | ctrl_i | noisy].
# The final noisy output is the weighted sum of the N pass outputs:
# noisy_out = w_1 * noisy_out_1 + ... + w_N * noisy_out_N
# All SDPA calls are maskless → Flash Attention always active.
# N=1, w=1.0 → identical to two_way_attention.
# Single-sample inference uses dense FMHA (cuDNN-eligible on sm_120);
# see multi_control_two_way_attention. N=1, w=1.0 → identical to two_way_attention.
#
# CP compatibility: control_stream_token_ranges are gen-relative global
# offsets computed here, before CP sharding. Ulysses CP restores the full
Expand Down
Loading