Skip to content

vllm.model_executor.models.qwen3_dflash2

Functions:

dflash2_grouped_conv_impl(x, delta, base, block_size, group_size)

DFlash2 grouped convolution adapted from LightSeek TokenSpeed PR #1399.

Source code in vllm/model_executor/models/qwen3_dflash2.py
def dflash2_grouped_conv_impl(
    x: torch.Tensor,
    delta: torch.Tensor,
    base: torch.Tensor,
    block_size: int,
    group_size: int,
) -> torch.Tensor:
    """DFlash2 grouped convolution adapted from LightSeek TokenSpeed PR #1399."""
    num_rows, num_channels = x.shape
    output = torch.empty_like(x)
    if num_rows == 0:
        return output

    element_block = 1024 if num_rows >= 128 and num_channels % 1024 == 0 else 512
    grid = (num_rows * triton.cdiv(num_channels, element_block),)
    _dflash2_grouped_conv_kernel[grid](
        x,
        delta,
        base,
        output,
        x.stride(0),
        delta.stride(0),
        delta.stride(1),
        base.stride(0),
        output.stride(0),
        NUM_CHANNELS=num_channels,
        BLOCK_SIZE=block_size,
        GROUP_SIZE=group_size,
        TAPS=base.shape[0],
        ELEMENT_BLOCK=element_block,
        num_warps=4,
    )
    return output