[Public release 26/04] Introducing Mega MoE, FP4 Indexer and other features/fixes (#304)
* Merge with private repo * Update README * Update README * Update README * Add PyTorch requirements * Fix sync scopes for MQA logits (#256) * Update README
This commit is contained in:
@@ -47,7 +47,7 @@ def a_fused_m_grouped_bf16_gemm_contiguous_tl_impl(a_ptr, b_ptr, d_ptr,
|
||||
# Compute
|
||||
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
for k in range(0, K, BLOCK_SIZE_K):
|
||||
k_range = (k + tl.arange(0, BLOCK_SIZE_K)).to(tl.int64)
|
||||
k_range = k.to(tl.int64) + tl.arange(0, BLOCK_SIZE_K).to(tl.int64)
|
||||
k_mask = k_range < K
|
||||
a_ptrs = a_ptr + rows[:, None] * K + k_range[None, :]
|
||||
b_ptrs = b_ptr + batch_id * K * N + k_range[:, None] * (1 if IS_B_K_MAJOR else N) + n_range[None, :].to(tl.int64) * (K if IS_B_K_MAJOR else 1)
|
||||
|
||||
@@ -50,7 +50,7 @@ def b_fused_k_grouped_bf16_gemm_contiguous_tl_impl(a_ptr, b_ptr, d_ptr,
|
||||
# Compute
|
||||
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
for k in range(k_start, k_end, BLOCK_SIZE_K):
|
||||
k_range = (k + tl.arange(0, BLOCK_SIZE_K)).to(tl.int64)
|
||||
k_range = k.to(tl.int64) + tl.arange(0, BLOCK_SIZE_K).to(tl.int64)
|
||||
rows = tl.load(k_indices_ptr + k_range).to(tl.int64)
|
||||
a_ptrs = a_ptr + m_range[:, None] + k_range[None, :] * M
|
||||
b_ptrs = b_ptr + rows[:, None] * N + n_range[None, :]
|
||||
|
||||
Reference in New Issue
Block a user