Skip to content

vllm.models.glm5next.nvidia.ops.third_party.kda.kernels

Functions:

chunk_kda_scaled_dot_kkt_fwd(q, k, gk=None, beta=None, scale=None, cu_seqlens=None, chunk_indices=None, chunk_size=FLA_CHUNK_SIZE, output_dtype=torch.float32)

Compute beta * K * K^T.

Parameters:

  • q

    (Tensor) –

    The query tensor of shape [B, T, H, K].

  • k

    (Tensor) –

    The key tensor of shape [B, T, H, K].

  • beta

    (Tensor, default: None ) –

    The beta tensor of shape [B, T, H].

  • gk

    (Tensor, default: None ) –

    The cumulative sum of the gate tensor of shape [B, T, H, K] applied to the key tensor. Default: None.

  • scale

    (float, default: None ) –

    Scale applied to the query-key products. Default: None.

  • cu_seqlens

    (Tensor, default: None ) –

    The cumulative sequence lengths of the input tensor. Default: None

  • chunk_indices

    (Tensor, default: None ) –

    Precomputed chunk indices for cu_seqlens. Default: None.

  • chunk_size

    (int, default: FLA_CHUNK_SIZE ) –

    The chunk size. Default: 64.

  • output_dtype

    (dtype, default: float32 ) –

    The dtype of the output tensor. Default: torch.float32

Returns:

  • tuple[Tensor, Tensor] –

    beta * K * K^T of shape [B, T, H, BT] where BT is the chunk size.

Source code in vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py
def chunk_kda_scaled_dot_kkt_fwd(
    q: torch.Tensor,
    k: torch.Tensor,
    gk: torch.Tensor | None = None,
    beta: torch.Tensor | None = None,
    scale: float | None = None,
    cu_seqlens: torch.Tensor | None = None,
    chunk_indices: torch.Tensor | None = None,
    chunk_size: int = FLA_CHUNK_SIZE,
    output_dtype: torch.dtype = torch.float32,
) -> tuple[torch.Tensor, torch.Tensor]:
    r"""Compute beta * K * K^T.

    Args:
        q (torch.Tensor):
            The query tensor of shape `[B, T, H, K]`.
        k (torch.Tensor):
            The key tensor of shape `[B, T, H, K]`.
        beta (torch.Tensor):
            The beta tensor of shape `[B, T, H]`.
        gk (torch.Tensor):
            The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`.
        scale (float):
            Scale applied to the query-key products. Default: `None`.
        cu_seqlens (torch.Tensor):
            The cumulative sequence lengths of the input tensor.
            Default: None
        chunk_indices (torch.Tensor):
            Precomputed chunk indices for `cu_seqlens`. Default: `None`.
        chunk_size (int):
            The chunk size. Default: 64.
        output_dtype (torch.dtype):
            The dtype of the output tensor. Default: `torch.float32`

    Returns:
        beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.

    """
    B, T, H, K = k.shape
    assert K <= 256
    BT = chunk_size
    if chunk_indices is None and cu_seqlens is not None:
        chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
    NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)

    BC = min(16, BT)
    NC = cdiv(BT, BC)
    BK = max(next_power_of_2(K), 16)
    A = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype)
    Aqk = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype)
    grid = (NT, NC * NC, B * H)
    chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid](
        q=q,
        k=k,
        g=gk,
        beta=beta,
        A=A,
        Aqk=Aqk,
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices,
        T=T,
        H=H,
        K=K,
        BT=BT,
        BC=BC,
        NC=NC,
    )

    grid = (NT, NC, B * H)
    chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra[grid](
        q=q,
        k=k,
        g=gk,
        beta=beta,
        A=A,
        Aqk=Aqk,
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices,
        T=T,
        H=H,
        K=K,
        BT=BT,
        BC=BC,
        BK=BK,
    )
    return A, Aqk

chunk_kda_with_fused_gate(q, k, v, raw_g, beta, A_log, g_bias, scale=None, initial_state=None, output_final_state=False, use_qk_l2norm_in_kernel=False, cu_seqlens=None, safe_gate=False, lower_bound=-5.0, **kwargs)

Run chunk KDA from raw gate projection using fused gate+cumsum.

Source code in vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py
def chunk_kda_with_fused_gate(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    raw_g: torch.Tensor,
    beta: torch.Tensor,
    A_log: torch.Tensor,
    g_bias: torch.Tensor | None,
    scale: float | None = None,
    initial_state: torch.Tensor | None = None,
    output_final_state: bool = False,
    use_qk_l2norm_in_kernel: bool = False,
    cu_seqlens: torch.Tensor | None = None,
    safe_gate: bool = False,
    lower_bound: float = -5.0,
    **kwargs,
):
    """Run chunk KDA from raw gate projection using fused gate+cumsum."""
    if scale is None:
        scale = k.shape[-1] ** -0.5

    if use_qk_l2norm_in_kernel:
        q = l2norm_fwd(q.contiguous())
        k = l2norm_fwd(k.contiguous())

    o, final_state = chunk_kda_with_fused_gate_fwd(
        q=q,
        k=k,
        v=v.contiguous(),
        raw_g=raw_g.contiguous(),
        beta=beta.contiguous(),
        A_log=A_log,
        g_bias=g_bias,
        scale=scale,
        initial_state=initial_state.contiguous() if initial_state is not None else None,
        output_final_state=output_final_state,
        cu_seqlens=cu_seqlens,
        safe_gate=safe_gate,
        lower_bound=lower_bound,
    )
    return o, final_state

fused_kda_gate(g, A, head_k_dim, g_bias=None, beta=1.0, threshold=20.0, safe_gate=False, lower_bound=-5.0)

Forward pass for KDA gate: input g: [..., HD] param A: [H] or [1, 1, H, 1] beta: softplus beta parameter (softplus branch only) threshold: softplus threshold parameter (softplus branch only) safe_gate: when False (default) compute y = -exp(A)softplus(g+g_bias); when True compute the bounded y = lower_boundsigmoid(exp(A)(g+g_bias)) lower_bound: floor for the safe_gate branch (default -5.0) return : [..., H, D]

Source code in vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py
def fused_kda_gate(
    g: torch.Tensor,
    A: torch.Tensor,
    head_k_dim: int,
    g_bias: torch.Tensor | None = None,
    beta: float = 1.0,
    threshold: float = 20.0,
    safe_gate: bool = False,
    lower_bound: float | None = -5.0,
) -> torch.Tensor:
    """Forward pass for KDA gate:
    input g: [..., H*D]
    param A: [H] or [1, 1, H, 1]
    beta: softplus beta parameter (softplus branch only)
    threshold: softplus threshold parameter (softplus branch only)
    safe_gate: when False (default) compute y = -exp(A)*softplus(g+g_bias);
      when True compute the bounded y = lower_bound*sigmoid(exp(A)*(g+g_bias))
    lower_bound: floor for the safe_gate branch (default -5.0)
    return  : [..., H, D]
    """
    orig_shape = g.shape[:-1]

    g = g.view(-1, g.shape[-1])
    T = g.shape[0]
    HD = g.shape[1]
    H = A.numel()
    assert H * head_k_dim == HD

    y = torch.empty_like(g, dtype=torch.float32)

    def grid(meta):
        return (cdiv(T, meta["BT"]), H)

    kda_gate_fwd_kernel[grid](
        g,
        A,
        y,
        g_bias,
        beta,
        threshold,
        safe_gate,
        lower_bound if lower_bound is not None else -5.0,
        T,
        H,
        head_k_dim,
        BD=next_power_of_2(head_k_dim),
        HAS_BIAS=g_bias is not None,
    )

    y = y.view(*orig_shape, H, head_k_dim)
    return y