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)iftrans_a).b (torch.Tensor) – Input, shape
(K, N)(or(N, K)iftrans_b; defaulttrans_b=True).c (torch.Tensor) – Accumulator / output tensor, shape
(M, N). Read whenbeta != 0and written in place.trans_a (bool) – Whether
a/bare stored transposed.trans_b (bool) – Whether
a/bare 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