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
50 changes: 25 additions & 25 deletions cosmos_framework/model/generator/mot/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,7 @@ def _is_split_info_compatible(attention_mask: object) -> bool:
get_causal_seq,
get_full_only_seq,
get_num_real_samples,
get_num_real_tokens,
sequence_pack_from_packed_sequence,
)

Expand Down Expand Up @@ -644,21 +645,25 @@ def multi_control_two_way_attention(
causal_out = causal_res.squeeze(0).flatten(-2, -1) # type: ignore # [N_text, Hq*D]

# ── 2. Extract unpadded full/gen tokens ──────────────────────────────────
full_q, full_q_offsets = get_full_only_seq(packed_query_states)
full_q, _ = get_full_only_seq(packed_query_states)
full_k, _ = get_full_only_seq(packed_key_states)
full_v, _ = get_full_only_seq(packed_value_states)

n_text = int(causal_k_offsets[-1])
n_full = int(full_q_offsets[-1])

# `n_full` comes from int(full_q_offsets[-1]) → an unbacked symint under
# torch.compile. The control ranges + noisy range partition the full/gen
# segment with noisy last, so `noisy_e` (a concrete int from SplitInfo) is
# exactly the number of valid gen tokens == n_full. Binding them lets Dynamo
# treat the per-segment `full_*_v[cs:ce]` slices below as concrete-length, so
# the in-place writes `full_out_v[cs:ce] = _sdpa(...)` don't raise
# data-dependent `Eq(slice_len, out_len)` guards.
torch._check(n_full == noisy_e)
# Read both counts off the pack rather than off the offsets tensors. `int(offsets[-1])` reads
# device memory, so under torch.compile it yields an unbacked symint and every slice below
# becomes data-dependent. That is fatal now that this path reaches cuDNN: its Inductor lowering
# compares the heads-first KV stride against the head dim, and for KV length `u + n` it cannot
# decide `128 * Max(1, u + n) > 128`. The pack already carries the same two counts as host-side
# ints, and they are exact rather than an approximation: `_compute_mode_indices_and_offsets`
# builds a mode's offsets and its index list in one pass over `split_lens`, so
# `offsets[-1] == len(indices) == num_<mode>_tokens`. Ulysses CP restores the full sequence on
# every rank before dispatch, so the batch-wide counts are the right ones here.
n_text, _ = get_num_real_tokens(packed_key_states)
_, n_full = get_num_real_tokens(packed_query_states)

# The control ranges + noisy range partition the full/gen segment with noisy last, so `noisy_e`
# is exactly the number of valid gen tokens.
assert n_full == noisy_e, f"gen stream holds {n_full} real tokens but the control split ends at {noisy_e}"

# Unpad to avoid padded rows entering the softmax denominator.
causal_k_v = causal_k[:n_text] # [N_text, Hkv, D]
Expand All @@ -673,22 +678,17 @@ def multi_control_two_way_attention(

def _sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
"""Maskless attention using cosmos_framework.model.attention() → [N_q, Hq*D]."""
# K and V are built by concatenating the SAME [text | ctrl_i | noisy]
# slices, so their sequence lengths are always equal. Under
# torch.compile (fullgraph=True) those lengths are unbacked symints
# (from data-dependent unpadding), and the attention frontend's
# `if key_shape[1] != value_shape[1]` guard (attention/checks.py) cannot
# be resolved symbolically. Assert the invariant so Dynamo can discharge
# the guard statically instead of raising a data-dependent error.
# K and V are built by concatenating the SAME [text | ctrl_i | noisy] slices, so their
# sequence lengths are always equal; the attention frontend still checks it
# (`key_shape[1] != value_shape[1]` in attention/checks.py). The padded stream length is a
# dynamic dimension under torch.compile, so state the invariants as facts rather than as
# plain asserts: that discharges the frontend's checks without specializing on it.
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. Backend validation checks require positive
# lengths, and cuDNN specifically rejects KV length 1. This path builds
# KV as [text | ctrl_i | noisy], where ctrl_i and noisy are non-empty for
# valid multi-control packs, so assert the stronger invariant. Without
# these, Dynamo cannot discharge them against unbacked symints.
# Backend validation requires positive lengths, and cuDNN specifically rejects KV length 1.
# This path builds KV as [text | ctrl_i | noisy], where ctrl_i and noisy are non-empty for
# valid multi-control packs, so assert the stronger invariant.
torch._check(n_q > 0)
torch._check(n_kv > 1)

Expand Down
68 changes: 68 additions & 0 deletions cosmos_framework/model/generator/mot/attention_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -783,6 +783,74 @@ def test_multi_control_range_annotation_rejects_inconsistent_token_count() -> No
_annotate_multi_control_ranges_for_test(attention_meta, packed_seq, n_gen=9)


@pytest.mark.L0
@torch.no_grad()
def test_multi_control_dynamic_compile_does_not_read_token_counts_from_device(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Keep the per-control slices free of data-dependent device-scalar reads.

The real Blackwell failure occurs later in cuDNN's Inductor lowering, when the unbacked
sequence-length symint from ``int(offsets[-1])`` reaches a stride comparison. Compiling the
whole path with a pure-Torch attention stub is enough to catch the source of that symint on
CPU: Dynamo rejects the device-scalar conversion before backend lowering. The pack already
carries the same counts as host-side ints, which is what this test requires the path to use.
"""
text_tokens = 3
control_tokens = (2, 3)
noisy_tokens = 4
full_tokens = sum(control_tokens) + noisy_tokens
total_tokens = text_tokens + full_tokens
und_indexes = torch.arange(text_tokens, dtype=torch.long)
gen_indexes = torch.arange(text_tokens, total_tokens, dtype=torch.long)

def make_pack(values: torch.Tensor) -> SequencePack:
return sequence_pack_from_packed_sequence(
packed_sequence=values,
attn_modes=["causal", "full"],
split_lens=[text_tokens, full_tokens],
sample_lens=[total_tokens],
packed_und_token_indexes=und_indexes,
packed_gen_token_indexes=gen_indexes,
)

query = make_pack(torch.randn(total_tokens, 4, 8))
key = make_pack(torch.randn(total_tokens, 2, 8))
value = make_pack(torch.randn(total_tokens, 2, 8))
split_info = attention.SplitInfo(
split_lens=[text_tokens, full_tokens],
attn_modes=["causal", "full"],
sample_lens=[total_tokens],
actual_len=total_tokens,
)
first_control_end = control_tokens[0]
controls_end = sum(control_tokens)
split_info.control_stream_token_ranges = [(0, first_control_end), (first_control_end, controls_end)]
split_info.noisy_token_range = (controls_end, full_tokens)
split_info.control_weights = [0.4, 0.6]

def fake_attention(
query_states: torch.Tensor,
key_states: torch.Tensor,
value_states: torch.Tensor,
**kwargs: Any,
) -> torch.Tensor:
del key_states, kwargs
return query_states.new_zeros((*query_states.shape[:-1], value_states.shape[-1]))

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

def run(query_pack: SequencePack, key_pack: SequencePack, value_pack: SequencePack) -> torch.Tensor:
result = attention.multi_control_two_way_attention(query_pack, key_pack, value_pack, split_info)
return result["full_only_seq"]

torch.compiler.reset()
compiled = torch.compile(run, fullgraph=True, dynamic=True, backend="eager")
output = compiled(query, key, value)

assert output.shape == (query["full_only_seq"].shape[0], 4 * 8)


# ── two_way_attention on the multiview FlexAttention mask ────────────────────
# The generator's full attention has two implementations of "every GEN token attends to
# its whole sample": the dense varlen kernel, and a single FlexAttention call over the
Expand Down