Update nvvm API call from nvvm enum to str (#2985)
This commit is contained in:
@@ -975,7 +975,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized Schedule warp
|
||||
#
|
||||
if warp_idx == self.sched_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_sched_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_sched_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1008,10 +1008,7 @@ class BlockwiseGemmKernel:
|
||||
)
|
||||
|
||||
# fence view async shared
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.sched_sync_barrier.arrive_and_wait()
|
||||
# commit tile info pipeline
|
||||
tile_info_pipeline.producer_commit(tile_info_producer_state)
|
||||
@@ -1023,7 +1020,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized TMA load warp
|
||||
#
|
||||
if warp_idx == self.tma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1126,10 +1123,7 @@ class BlockwiseGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1142,7 +1136,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized Scale load warp
|
||||
#
|
||||
if warp_idx == self.scale_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1301,10 +1295,7 @@ class BlockwiseGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1317,7 +1308,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized MMA warp
|
||||
#
|
||||
if warp_idx == self.mma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared mem
|
||||
#
|
||||
@@ -1459,10 +1450,7 @@ class BlockwiseGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1475,7 +1463,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized acc update warps
|
||||
#
|
||||
if warp_idx <= self.acc_update_warp_id[-1]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_acc_update_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_acc_update_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared memory
|
||||
#
|
||||
@@ -1696,10 +1684,7 @@ class BlockwiseGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1707,7 +1692,7 @@ class BlockwiseGemmKernel:
|
||||
# Specialized epilogue warps
|
||||
#
|
||||
if warp_idx <= self.epilog_warp_id[-1] and warp_idx >= self.epilog_warp_id[0]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_epilogue_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_epilogue_warps)
|
||||
#
|
||||
# Alloc tensor memory buffer
|
||||
#
|
||||
@@ -1866,10 +1851,7 @@ class BlockwiseGemmKernel:
|
||||
tRS_sC[(None, None, None, c_buffer)],
|
||||
)
|
||||
# Fence and barrier to make sure shared memory store is visible to TMA store
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.epilog_sync_barrier.arrive_and_wait()
|
||||
|
||||
#
|
||||
@@ -1899,10 +1881,7 @@ class BlockwiseGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
|
||||
@@ -996,7 +996,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized Schedule warp
|
||||
#
|
||||
if warp_idx == self.sched_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_sched_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_sched_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1034,10 +1034,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
)
|
||||
|
||||
# fence view async shared
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.sched_sync_barrier.arrive_and_wait()
|
||||
# commit tile info pipeline
|
||||
tile_info_pipeline.producer_commit(tile_info_producer_state)
|
||||
@@ -1051,7 +1048,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized TMA load warp
|
||||
#
|
||||
if warp_idx == self.tma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1153,10 +1150,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1168,7 +1162,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized Scale load warp
|
||||
#
|
||||
if warp_idx == self.scale_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1328,10 +1322,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1344,7 +1335,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized MMA warp
|
||||
#
|
||||
if warp_idx == self.mma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared mem
|
||||
#
|
||||
@@ -1488,10 +1479,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1504,7 +1492,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized acc update warps
|
||||
#
|
||||
if warp_idx <= self.acc_update_warp_id[-1]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_acc_update_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_acc_update_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared memory
|
||||
#
|
||||
@@ -1727,10 +1715,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1738,7 +1723,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
# Specialized epilogue warps
|
||||
#
|
||||
if warp_idx <= self.epilog_warp_id[-1] and warp_idx >= self.epilog_warp_id[0]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_epilogue_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_epilogue_warps)
|
||||
#
|
||||
# Alloc tensor memory buffer
|
||||
#
|
||||
@@ -1899,10 +1884,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
tRS_sC[(None, None, None, c_buffer)],
|
||||
)
|
||||
# Fence and barrier to make sure shared memory store is visible to TMA store
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.epilog_sync_barrier.arrive_and_wait()
|
||||
|
||||
#
|
||||
@@ -1932,10 +1914,7 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
|
||||
@@ -995,7 +995,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized Schedule warp
|
||||
#
|
||||
if warp_idx == self.sched_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_sched_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_sched_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1041,10 +1041,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
)
|
||||
|
||||
# fence view async shared
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.sched_sync_barrier.arrive_and_wait()
|
||||
# commit tile info pipeline
|
||||
tile_info_pipeline.producer_commit(tile_info_producer_state)
|
||||
@@ -1056,7 +1053,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized TMA load warp
|
||||
#
|
||||
if warp_idx == self.tma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1159,10 +1156,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1175,7 +1169,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized Scale load warp
|
||||
#
|
||||
if warp_idx == self.scale_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
#
|
||||
@@ -1334,10 +1328,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1350,7 +1341,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized MMA warp
|
||||
#
|
||||
if warp_idx == self.mma_warp_id:
|
||||
cute.arch.warpgroup_reg_dealloc(self.num_regs_uniform_warps)
|
||||
cute.arch.setmaxregister_decrease(self.num_regs_uniform_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared mem
|
||||
#
|
||||
@@ -1492,10 +1483,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1508,7 +1496,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized acc update warps
|
||||
#
|
||||
if warp_idx <= self.acc_update_warp_id[-1]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_acc_update_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_acc_update_warps)
|
||||
#
|
||||
# Bar sync for retrieve tensor memory ptr from shared memory
|
||||
#
|
||||
@@ -1729,10 +1717,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
@@ -1740,7 +1725,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
# Specialized epilogue warps
|
||||
#
|
||||
if warp_idx <= self.epilog_warp_id[-1] and warp_idx >= self.epilog_warp_id[0]:
|
||||
cute.arch.warpgroup_reg_alloc(self.num_regs_epilogue_warps)
|
||||
cute.arch.setmaxregister_increase(self.num_regs_epilogue_warps)
|
||||
#
|
||||
# Alloc tensor memory buffer
|
||||
#
|
||||
@@ -1899,10 +1884,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
tRS_sC[(None, None, None, c_buffer)],
|
||||
)
|
||||
# Fence and barrier to make sure shared memory store is visible to TMA store
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
self.epilog_sync_barrier.arrive_and_wait()
|
||||
|
||||
#
|
||||
@@ -1932,10 +1914,7 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
for idx in cutlass.range(4, unroll_full=True):
|
||||
tile_info[idx] = sInfo[(idx, tile_info_consumer_state.index)]
|
||||
is_valid_tile = tile_info[3] == 1
|
||||
cute.arch.fence_proxy(
|
||||
cute.arch.ProxyKind.async_shared,
|
||||
space=cute.arch.SharedSpace.shared_cta,
|
||||
)
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
tile_info_pipeline.consumer_release(tile_info_consumer_state)
|
||||
tile_info_consumer_state.advance()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user