From b165ad3f2a854ccba1f6cfb9e0d2b816a4b230c1 Mon Sep 17 00:00:00 2001 From: Sunil Srinivasa <106262814+sunil-srinivasa@users.noreply.github.com> Date: Tue, 22 Sep 2026 11:01:42 -0700 Subject: [PATCH] fix(inference): remove dynamic device scalar from multi-control attention Signed-off-by: Sunil Srinivasa <106262814+sunil-srinivasa@users.noreply.github.com> --- .../model/generator/mot/attention.py | 50 +++++++------- .../model/generator/mot/attention_test.py | 68 +++++++++++++++++++ 2 files changed, 93 insertions(+), 25 deletions(-) diff --git a/cosmos_framework/model/generator/mot/attention.py b/cosmos_framework/model/generator/mot/attention.py index 03e4c64a..014f6fad 100644 --- a/cosmos_framework/model/generator/mot/attention.py +++ b/cosmos_framework/model/generator/mot/attention.py @@ -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, ) @@ -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__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] @@ -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) diff --git a/cosmos_framework/model/generator/mot/attention_test.py b/cosmos_framework/model/generator/mot/attention_test.py index 9a23bb8b..4b81f3ef 100644 --- a/cosmos_framework/model/generator/mot/attention_test.py +++ b/cosmos_framework/model/generator/mot/attention_test.py @@ -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