Skip to content

vllm.models.dots3_note.nvidia.attention

Dots3 NOTE sliding-window MLA attention backends for Hopper.

Prefill and mixed batches expand the latent cache and use FlashAttention-3 varlen MHA. Decode-only batches use the Triton absorbed-MQA kernel.

Classes:

Dots3NoteFlashAttnPrefillBackend

Bases: FlashAttnPrefillBackend

FA3 varlen prefill for the NOTE SWA MLA dimensions.

Source code in vllm/models/dots3_note/nvidia/attention.py
class Dots3NoteFlashAttnPrefillBackend(FlashAttnPrefillBackend):
    """FA3 varlen prefill for the NOTE SWA MLA dimensions."""

    @classmethod
    def supports_mla_dimensions(cls, dims: MLADimensions) -> bool:
        return dims == MLADimensions(
            qk_nope_head_dim=192,
            qk_rope_head_dim=64,
            v_head_dim=128,
        )

    def __init__(self, *args, **kwargs) -> None:
        super().__init__(*args, **kwargs)
        # FA3 on SM90 does not support NOTE's Q/K=256, V=128 combination via
        # the different-head-dimension path. Padding V selects its supported
        # equal-dimension varlen kernel.
        self.requires_v_padding = True

    def supports_quant_output(self, quant_key) -> bool:
        return False

    def run_sliding_window(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        cu_seq_lens_q: torch.Tensor,
        cu_seq_lens_k: torch.Tensor,
        max_seq_len_q: int,
        max_seq_len_k: int,
        sliding_window: int,
    ) -> torch.Tensor:
        output = self._flash_attn_varlen_diff_headdims(
            q=q,
            k=k,
            v=v,
            cu_seqlens_q=cu_seq_lens_q,
            cu_seqlens_k=cu_seq_lens_k,
            max_seqlen_q=max_seq_len_q,
            max_seqlen_k=max_seq_len_k,
            softmax_scale=self.scale,
            causal=True,
            window_size=(sliding_window - 1, 0),
            return_softmax_lse=False,
        )
        assert isinstance(output, torch.Tensor)
        return output

Dots3NoteMLAMetadataBuilder

Bases: TritonMLAMetadataBuilder

Keep decode on MQA and route prefill/mixed batches through FA3.

Source code in vllm/models/dots3_note/nvidia/attention.py
class Dots3NoteMLAMetadataBuilder(TritonMLAMetadataBuilder):
    """Keep decode on MQA and route prefill/mixed batches through FA3."""

    query_len_support = QueryLenSupport.UNIFORM

    def __init__(self, kv_cache_spec, layer_names, vllm_config, device):
        self.sliding_window = kv_cache_spec.sliding_window
        super().__init__(kv_cache_spec, layer_names, vllm_config, device)

    def _reserve_attn_logits_workspace(self) -> None:
        if not is_workspace_manager_initialized():
            return
        batch = self.vllm_config.scheduler_config.max_num_seqs
        max_query_len = self.reorder_batch_threshold or 1
        gather_len = (self.sliding_window + max_query_len - 1 + 7) // 8 * 8
        current_workspace_manager().get_simultaneous(
            (
                (
                    batch,
                    gather_len,
                    self.mla_dims.kv_lora_rank + self.mla_dims.qk_rope_head_dim,
                ),
                self.model_config.dtype,
            ),
            ((batch, gather_len), torch.bool),
        )

    def _build_decode(
        self,
        block_table_tensor: torch.Tensor,
        seq_lens_device: torch.Tensor,
        max_seq_len: int,
        query_start_loc_cpu: torch.Tensor,
        query_start_loc_device: torch.Tensor,
        num_decode_tokens: int,
        dcp_tot_seq_lens_device: torch.Tensor | None,
    ) -> Dots3NoteDecodeMetadata:
        del max_seq_len, query_start_loc_device, num_decode_tokens
        query_len = int(query_start_loc_cpu[1] - query_start_loc_cpu[0])
        return Dots3NoteDecodeMetadata(
            block_table=block_table_tensor,
            seq_lens=seq_lens_device,
            dcp_tot_seq_lens=dcp_tot_seq_lens_device,
            query_len=query_len,
        )

    def build(
        self,
        common_prefix_len: int,
        common_attn_metadata: CommonAttentionMetadata,
        fast_build: bool = False,
    ) -> MLACommonMetadata:
        metadata = super().build(
            common_prefix_len,
            common_attn_metadata,
            fast_build=fast_build,
        )
        if metadata.prefill is None:
            return metadata
        if self.dcp_world_size > 1:
            raise NotImplementedError(
                "Dots3 NOTE SWA prefill does not support decode context parallelism"
            )

        reqs_start = metadata.num_decodes
        seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
        assert seq_lens_cpu is not None
        query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu[reqs_start:]
        query_start_loc_cpu = query_start_loc_cpu - query_start_loc_cpu[0]
        sliding_metadata = _build_sliding_window_metadata(
            seq_lens_cpu=seq_lens_cpu[reqs_start:],
            query_start_loc_cpu=query_start_loc_cpu,
            sliding_window=self.sliding_window,
            workspace=self.chunked_prefill_workspace,
            workspace_size=self.chunked_prefill_workspace_size,
            device=self.device,
        )
        metadata.prefill.chunked_context = None
        metadata.prefill.sliding_window = sliding_metadata  # type: ignore[attr-defined]
        if metadata.num_decodes > 0 and metadata.num_prefills > 0:
            query_lens_cpu = query_start_loc_cpu[1:] - query_start_loc_cpu[:-1]
            for chunk in sliding_metadata.chunks:
                req_start, req_end = chunk.req_start, chunk.req_end
                seq_lens = common_attn_metadata.seq_lens[
                    reqs_start + req_start : reqs_start + req_end
                ]
                query_lens = query_lens_cpu[req_start:req_end].to(
                    self.device, non_blocking=True
                )
                kv_lens = torch.minimum(seq_lens, query_lens + self.sliding_window - 1)
                chunk.cu_seq_lens_k.zero_()
                torch.cumsum(
                    kv_lens,
                    dim=0,
                    out=chunk.cu_seq_lens_k[1:],
                    dtype=torch.int32,
                )
                chunk.starts = seq_lens - kv_lens
                token_ids = torch.arange(
                    chunk.num_kv_tokens,
                    dtype=torch.int32,
                    device=self.device,
                )
                chunk.token_to_seq = torch.searchsorted(
                    chunk.cu_seq_lens_k[1:], token_ids, right=True
                ).to(torch.int32)
                chunk.token_to_seq.clamp_max_(req_end - req_start - 1)
        assert metadata.prefill.prefill_backend is not None
        metadata.prefill.prefill_backend.prepare_metadata(metadata.prefill)
        return metadata

Dots3NotePaddedSparseBackend

Bases: FlashAttnMLASparseBackend

NOTE DSA backend for cache rows padded to the SWA latent width.

Source code in vllm/models/dots3_note/nvidia/attention.py
class Dots3NotePaddedSparseBackend(FlashAttnMLASparseBackend):
    """NOTE DSA backend for cache rows padded to the SWA latent width."""

    @staticmethod
    def get_name() -> str:
        return "DOTS3_NOTE_PADDED_MLA_SPARSE"

    @staticmethod
    def get_impl_cls() -> type["Dots3NotePaddedSparseImpl"]:
        return Dots3NotePaddedSparseImpl

Dots3NotePaddedSparseImpl

Bases: FlashAttnMLASparseImpl

Read top-k KV directly from uniformly padded cache rows.

Source code in vllm/models/dots3_note/nvidia/attention.py
class Dots3NotePaddedSparseImpl(FlashAttnMLASparseImpl):
    """Read top-k KV directly from uniformly padded cache rows."""

    def _logical_cache(self, kv_cache: torch.Tensor) -> torch.Tensor:
        assert kv_cache.shape[-1] >= self.head_size
        return kv_cache[..., : self.head_size]

    def do_kv_cache_update(
        self,
        kv_c_normed: torch.Tensor,
        k_pe: torch.Tensor,
        kv_cache: torch.Tensor,
        slot_mapping: torch.Tensor,
        kv_cache_dtype: str,
        k_scale: torch.Tensor,
    ) -> None:
        super().do_kv_cache_update(
            kv_c_normed,
            k_pe,
            self._logical_cache(kv_cache),
            slot_mapping,
            kv_cache_dtype,
            k_scale,
        )

    def forward_mha(  # type: ignore[override]
        self,
        q: torch.Tensor,
        kv_c_normed: torch.Tensor,
        k_pe: torch.Tensor,
        kv_c_and_k_pe_cache: torch.Tensor,
        attn_metadata: FlashAttnMLASparseMetadata,
        k_scale: torch.Tensor,
        output: torch.Tensor,
        output_scale: torch.Tensor | None = None,
    ) -> None:
        super().forward_mha(
            q,
            kv_c_normed,
            k_pe,
            self._logical_cache(kv_c_and_k_pe_cache),
            attn_metadata,
            k_scale,
            output,
            output_scale,
        )

    def forward_mqa(
        self,
        q: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
        kv_c_and_k_pe_cache: torch.Tensor,
        attn_metadata: FlashAttnMLASparseMetadata,
        layer: AttentionLayer,
    ) -> tuple[torch.Tensor, torch.Tensor | None]:
        if not isinstance(q, tuple):
            raise NotImplementedError(
                "Dots3NotePaddedSparseImpl expects split (q_nope, q_rope) input."
            )
        q_nope, q_rope = q
        num_actual_toks = q_rope.shape[0]

        assert self.topk_indices_buffer is not None
        topk_indices = self.topk_indices_buffer[:num_actual_toks]
        topk_indices, valid_counts = triton_convert_req_index_to_global_index(
            attn_metadata.req_id_per_token[:num_actual_toks],
            attn_metadata.block_table,
            topk_indices,
            BLOCK_SIZE=attn_metadata.block_size,
            NUM_TOPK_TOKENS=topk_indices.shape[1],
            return_valid_counts=True,
        )
        physical_head_size = kv_c_and_k_pe_cache.shape[-1]
        block_stride, token_stride, head_stride = kv_c_and_k_pe_cache.stride()
        if (
            token_stride != physical_head_size
            or head_stride != 1
            or block_stride % physical_head_size != 0
        ):
            raise RuntimeError(
                "Dots3 NOTE padded DSA cache must contain contiguous token rows; "
                f"got shape={tuple(kv_c_and_k_pe_cache.shape)} and "
                f"stride={kv_c_and_k_pe_cache.stride()}"
            )
        block_row_stride = block_stride // physical_head_size
        block_indices = torch.div(
            topk_indices,
            attn_metadata.block_size,
            rounding_mode="floor",
        )
        block_indices.clamp_min_(0).mul_(block_row_stride - attn_metadata.block_size)
        topk_indices.add_(block_indices)

        num_cache_rows = (
            kv_c_and_k_pe_cache.shape[0] - 1
        ) * block_row_stride + attn_metadata.block_size
        kv_cache = kv_c_and_k_pe_cache.as_strided(
            (num_cache_rows, physical_head_size),
            (physical_head_size, 1),
        )
        cu_seqlens_q = torch.arange(
            num_actual_toks + 1,
            dtype=torch.int32,
            device=q_rope.device,
        )
        output = flash_attn_varlen_func(
            q=q_rope,
            k=kv_cache[:, self.kv_lora_rank : self.head_size].unsqueeze(1).unsqueeze(1),
            v=kv_cache[:, : self.kv_lora_rank].unsqueeze(1).unsqueeze(1),
            q_v=q_nope,
            max_seqlen_q=1,
            cu_seqlens_q=cu_seqlens_q,
            max_seqlen_k=topk_indices.shape[1],
            seqused_k=valid_counts,
            block_table=topk_indices,
            softmax_scale=self.scale,
            causal=True,
            fa_version=3,
        )
        return output, None

Dots3NoteTritonMLABackend

Bases: TritonMLABackend

Internal NOTE SWA specialization; not a user-selectable backend.

Source code in vllm/models/dots3_note/nvidia/attention.py
class Dots3NoteTritonMLABackend(TritonMLABackend):
    """Internal NOTE SWA specialization; not a user-selectable backend."""

    @classmethod
    def get_supported_head_sizes(cls) -> list[int]:
        return [1088]

    @classmethod
    def supports_sliding_window(cls) -> bool:
        return True

    @staticmethod
    def get_impl_cls() -> type["Dots3NoteTritonMLAImpl"]:
        return Dots3NoteTritonMLAImpl

    @staticmethod
    def get_builder_cls() -> type[Dots3NoteMLAMetadataBuilder]:
        return Dots3NoteMLAMetadataBuilder

_build_sliding_window_metadata(*, seq_lens_cpu, query_start_loc_cpu, sliding_window, workspace, workspace_size, device)

Plan per-request latent-cache gathers for SWA varlen attention.

Source code in vllm/models/dots3_note/nvidia/attention.py
def _build_sliding_window_metadata(
    *,
    seq_lens_cpu: torch.Tensor,
    query_start_loc_cpu: torch.Tensor,
    sliding_window: int,
    workspace: torch.Tensor,
    workspace_size: int,
    device: torch.device,
) -> _SlidingWindowMetadata:
    """Plan per-request latent-cache gathers for SWA varlen attention."""
    query_lens_cpu = (query_start_loc_cpu[1:] - query_start_loc_cpu[:-1]).to(
        dtype=torch.int32
    )
    seq_lens_cpu = seq_lens_cpu.to(dtype=torch.int32)
    kv_lens_cpu = torch.minimum(seq_lens_cpu, query_lens_cpu + sliding_window - 1)
    starts_cpu = seq_lens_cpu - kv_lens_cpu

    chunks: list[_SlidingWindowChunk] = []
    req_start = 0
    while req_start < query_lens_cpu.numel():
        req_end = req_start
        num_kv_tokens = 0
        while req_end < query_lens_cpu.numel():
            next_len = int(kv_lens_cpu[req_end].item())
            if num_kv_tokens and num_kv_tokens + next_len > workspace_size:
                break
            if next_len > workspace_size:
                raise ValueError(
                    "Dots3 NOTE SWA prefill window exceeds the MLA workspace: "
                    f"{next_len} > {workspace_size}"
                )
            num_kv_tokens += next_len
            req_end += 1

        query_lens = query_lens_cpu[req_start:req_end]
        kv_lens = kv_lens_cpu[req_start:req_end]
        num_reqs = req_end - req_start
        cu_seq_lens_q_cpu = torch.zeros(num_reqs + 1, dtype=torch.int32)
        cu_seq_lens_k_cpu = torch.zeros(num_reqs + 1, dtype=torch.int32)
        torch.cumsum(query_lens, 0, out=cu_seq_lens_q_cpu[1:])
        torch.cumsum(kv_lens, 0, out=cu_seq_lens_k_cpu[1:])
        token_to_seq_cpu = torch.repeat_interleave(
            torch.arange(num_reqs, dtype=torch.int32), kv_lens
        )
        query_start = int(query_start_loc_cpu[req_start].item())
        query_end = int(query_start_loc_cpu[req_end].item())
        chunks.append(
            _SlidingWindowChunk(
                req_start=req_start,
                req_end=req_end,
                query_start=query_start,
                query_end=query_end,
                cu_seq_lens_q=cu_seq_lens_q_cpu.to(device, non_blocking=True),
                cu_seq_lens_k=cu_seq_lens_k_cpu.to(device, non_blocking=True),
                starts=starts_cpu[req_start:req_end].to(device, non_blocking=True),
                token_to_seq=token_to_seq_cpu.to(device, non_blocking=True),
                num_kv_tokens=num_kv_tokens,
                max_seq_len_q=int(query_lens.max().item()),
                max_seq_len_k=int(kv_lens.max().item()),
            )
        )
        req_start = req_end

    return _SlidingWindowMetadata(chunks=chunks, workspace=workspace)