Fix scale_step_k computation in the fp8_kernel (#20819)
Co-authored-by: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com>
This commit is contained in:
@@ -759,7 +759,7 @@ def _w8a8_block_fp8_matmul(
|
|||||||
As_ptrs = As + offs_am * stride_As_m
|
As_ptrs = As + offs_am * stride_As_m
|
||||||
offs_bsn = offs_bn // group_n
|
offs_bsn = offs_bn // group_n
|
||||||
Bs_ptrs = Bs + offs_bsn * stride_Bs_n
|
Bs_ptrs = Bs + offs_bsn * stride_Bs_n
|
||||||
scale_step_k = BLOCK_SIZE_K // group_k
|
n_tiles_k_per_group_k = group_k // BLOCK_SIZE_K
|
||||||
|
|
||||||
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
||||||
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
|
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
|
||||||
@@ -773,6 +773,7 @@ def _w8a8_block_fp8_matmul(
|
|||||||
a_s = tl.load(As_ptrs)
|
a_s = tl.load(As_ptrs)
|
||||||
b_s = tl.load(Bs_ptrs)
|
b_s = tl.load(Bs_ptrs)
|
||||||
|
|
||||||
|
scale_step_k = tl.where((k + 1) % n_tiles_k_per_group_k == 0, 1, 0)
|
||||||
accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :]
|
accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :]
|
||||||
a_ptrs += BLOCK_SIZE_K * stride_ak
|
a_ptrs += BLOCK_SIZE_K * stride_ak
|
||||||
b_ptrs += BLOCK_SIZE_K * stride_bk
|
b_ptrs += BLOCK_SIZE_K * stride_bk
|
||||||
|
|||||||
Reference in New Issue
Block a user