Skip to content

vllm.models.kimi_k3.nvidia.ops.cute_dsl.kda_skinny_gemm

Kimi-K3 TP8 skinny GEMMs for the KDA F_A/beta and F_B projections.

_fma_f32_bf16_portable(a, b, acc, *, loc=None, ip=None)

BF16 multiply with FP32 accumulation for pre-SM100 GPUs.

PTX fma.f32.bf16 requires SM100 or newer. Converting the operands to FP32 first keeps this kernel available on Hopper, where an unconditional mixed-precision FMA fails libNVVM compilation for sm_90a.

Source code in vllm/models/kimi_k3/nvidia/ops/cute_dsl/kda_skinny_gemm.py
@dsl_user_op
def _fma_f32_bf16_portable(
    a: BFloat16,
    b: BFloat16,
    acc: Float32,
    *,
    loc=None,
    ip=None,
) -> Float32:
    """BF16 multiply with FP32 accumulation for pre-SM100 GPUs.

    PTX ``fma.f32.bf16`` requires SM100 or newer. Converting the operands to
    FP32 first keeps this kernel available on Hopper, where an unconditional
    mixed-precision FMA fails libNVVM compilation for ``sm_90a``.
    """
    a_bits = llvm.bitcast(T.i16(), a.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
    b_bits = llvm.bitcast(T.i16(), b.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
    result = llvm.inline_asm(
        T.f32(),
        [a_bits, b_bits, acc.ir_value(loc=loc, ip=ip)],
        "{\n\t"
        ".reg .f32 a_f32, b_f32;\n\t"
        "cvt.f32.bf16 a_f32, $1;\n\t"
        "cvt.f32.bf16 b_f32, $2;\n\t"
        "fma.rn.f32 $0, a_f32, b_f32, $3;\n\t"
        "}",
        "=f,h,h,f",
        has_side_effects=False,
        is_align_stack=False,
        asm_dialect=llvm.AsmDialect.AD_ATT,
        loc=loc,
        ip=ip,
    )
    return Float32(result)

_has_mixed_precision_bf16_fma()

Whether PTX fma.f32.bf16 is supported by the current GPU.

Source code in vllm/models/kimi_k3/nvidia/ops/cute_dsl/kda_skinny_gemm.py
def _has_mixed_precision_bf16_fma() -> bool:
    """Whether PTX ``fma.f32.bf16`` is supported by the current GPU."""
    return current_platform.has_device_capability(100)