flashinfer.gemm.gemm_alpha_beta

flashinfer.gemm.gemm_alpha_beta(a: Tensor, b: Tensor, c: Tensor, trans_a: bool = False, trans_b: bool = True, alpha: float = 1.0, beta: float = 0.0, num_sms: int | None = None, backend: Literal['cutile'] = 'cutile') → Tensor

GEMM with alpha/beta scaling: C = alpha * (A @ B) + beta * C.

Parameters:
  • a (torch.Tensor) – Input, shape (M, K) (or (K, M) if trans_a).

  • b (torch.Tensor) – Input, shape (K, N) (or (N, K) if trans_b; default trans_b=True).

  • c (torch.Tensor) – Accumulator / output tensor, shape (M, N). Read when beta != 0 and written in place.

  • trans_a (bool) – Whether a / b are stored transposed.

  • trans_b (bool) – Whether a / b are stored transposed.

  • alpha (float) – Scaling factors.

  • beta (float) – Scaling factors.

  • num_sms (Optional[int]) – Optional override for the number of SMs used by the grid.

  • backend (str) – Implementation backend. Currently only "cutile" is supported.

Returns:

The output tensor C.

Return type:

torch.Tensor