Update blackwell tutorial to be compatible with 4.5-dev version (#3130)

* Update blackwell tutorial to be compatible with 4.5-dev version

* update example for reverted changes

* add more example fix
This commit is contained in:
Longsheng Du
2026-04-09 14:40:33 +08:00
committed by GitHub
parent bd01dd3651
commit 08185b9c3e
12 changed files with 29 additions and 29 deletions

View File

@@ -219,11 +219,11 @@ def kernel(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True if use_2cta_instrs else False,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Partition tensors for TMA; This requires the tensors partitioned for MMA

View File

@@ -152,11 +152,11 @@ def kernel(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tma_copy_bytes = (

View File

@@ -159,11 +159,11 @@ def kernel(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tma_copy_bytes = (

View File

@@ -184,11 +184,11 @@ def cluster_specific_kernel(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tma_copy_bytes = (

View File

@@ -171,11 +171,11 @@ def kernel(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tma_copy_bytes = (

View File

@@ -214,11 +214,11 @@ def gemm(
* len((mma_warp_id, *epilogue_warp_ids)), # 5 warps = 160 threads
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buffer,
storage.tmem_holding_buffer.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=epilogue_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tma_copy_bytes = (