v4.5 dev update. (#3153)
This commit is contained in:
@@ -603,8 +603,6 @@ class MixedInputFusedMultiHeadAttentionDecode:
|
||||
p_pipeline_ptr = smem.allocate_array(Int64, self.sp_stages * 2)
|
||||
o_pipeline_ptr = smem.allocate_array(Int64, self.o_stages * 2)
|
||||
|
||||
assert smem._allocated_bytes <= self.mbarrier_reserved_bytes
|
||||
|
||||
# Declare named barriers
|
||||
softmax_nbar_id = 1
|
||||
mma_kq_nbar_id = 2
|
||||
|
||||
@@ -403,7 +403,7 @@ class MixedInputFusedMultiHeadAttentionPrefillD256:
|
||||
s_corr_mbar_ptr: cute.struct.MemRange[Int64, self.qk_acc_stage * 2]
|
||||
sum_mbar_ptr: cute.struct.MemRange[Int64, 2]
|
||||
mma_o_mbar_ptr: cute.struct.MemRange[Int64, self.pv_acc_stage * 2]
|
||||
tmem_dealloc_mbar_ptr: Int64
|
||||
tmem_dealloc_mbar: Int64
|
||||
tmem_holding_buf: Int32
|
||||
|
||||
self.shared_storage = SharedStorage
|
||||
@@ -654,11 +654,11 @@ class MixedInputFusedMultiHeadAttentionPrefillD256:
|
||||
)
|
||||
# Tensor memory dealloc barrier init
|
||||
tmem = utils.TmemAllocator(
|
||||
storage.tmem_holding_buf,
|
||||
storage.tmem_holding_buf.ptr,
|
||||
barrier_for_retrieve=tmem_alloc_barrier,
|
||||
allocator_warp_id=self.correction_warp_ids[0],
|
||||
is_two_cta=True,
|
||||
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
|
||||
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
|
||||
)
|
||||
# Cluster arrive after barrier init
|
||||
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
|
||||
|
||||
@@ -390,7 +390,7 @@ class MixedInputFusedMultiHeadAttentionPrefillD512:
|
||||
p_mma_mbar_ptr: cute.struct.MemRange[Int64, self.qk_acc_stage * 2]
|
||||
mma_o_mbar_ptr: cute.struct.MemRange[Int64, self.pv_acc_stage * 2]
|
||||
swap_mbar_ptr: cute.struct.MemRange[Int64, self.swap_stage * 2]
|
||||
tmem_dealloc_mbar_ptr: Int64
|
||||
tmem_dealloc_mbar: Int64
|
||||
tmem_holding_buf: Int32
|
||||
|
||||
self.shared_storage = SharedStorage
|
||||
@@ -627,11 +627,11 @@ class MixedInputFusedMultiHeadAttentionPrefillD512:
|
||||
)
|
||||
# Tensor memory dealloc barrier init
|
||||
tmem = utils.TmemAllocator(
|
||||
storage.tmem_holding_buf,
|
||||
storage.tmem_holding_buf.ptr,
|
||||
barrier_for_retrieve=tmem_alloc_barrier,
|
||||
allocator_warp_id=self.softmax_warp_ids[0],
|
||||
is_two_cta=True,
|
||||
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
|
||||
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
|
||||
)
|
||||
# Cluster arrive after barrier init
|
||||
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
|
||||
|
||||
Reference in New Issue
Block a user