* Merge with private repo * Update README * Update README * Update README * Add PyTorch requirements * Fix sync scopes for MQA logits (#256) * Update README
169 lines
5.3 KiB
Plaintext
169 lines
5.3 KiB
Plaintext
#pragma once
|
|
|
|
namespace deep_gemm::ptx {
|
|
|
|
/// UMMA versions with relaxed assertions
|
|
struct SM100_MMA_F16BF16_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, p; \n\t"
|
|
"}\n"
|
|
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_F16BF16_2x1SM_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::2.kind::f16 [%0], %1, %2, %3, p; \n\t"
|
|
"}\n"
|
|
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_MXF8F6F4_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc,
|
|
uint32_t const& tmem_sfa,
|
|
uint32_t const& tmem_sfb) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::1.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
|
"}\n"
|
|
:
|
|
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
|
|
"r"(tmem_sfa), "r"(tmem_sfb));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_MXF8F6F4_2x1SM_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc,
|
|
uint32_t const& tmem_sfa,
|
|
uint32_t const& tmem_sfb) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::2.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
|
"}\n"
|
|
:
|
|
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
|
|
"r"(tmem_sfa), "r"(tmem_sfb));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_F8F6F4_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::1.kind::f8f6f4 [%0], %1, %2, %3, p; \n\t"
|
|
"}\n"
|
|
:
|
|
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_F8F6F4_2x1SM_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.cta_group::2.kind::f8f6f4 [%0], %1, %2, %3, p; \n\t"
|
|
"}\n"
|
|
:
|
|
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_MXF4_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc,
|
|
uint32_t const& tmem_sfa,
|
|
uint32_t const& tmem_sfb) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 9)
|
|
"tcgen05.mma.cta_group::1.kind::mxf4.block_scale.block32 [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
|
#else
|
|
"tcgen05.mma.cta_group::1.kind::mxf4.block_scale.scale_vec::2X [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
|
#endif
|
|
"}\n"
|
|
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
|
|
"r"(tmem_sfa), "r"(tmem_sfb));
|
|
}
|
|
};
|
|
|
|
struct SM100_MMA_F16BF16_WS_SS {
|
|
CUTLASS_DEVICE static void
|
|
fma(uint64_t const& desc_a,
|
|
uint64_t const& desc_b,
|
|
uint32_t const& tmem_c,
|
|
uint32_t const& scale_c,
|
|
uint64_t const& desc) {
|
|
asm volatile(
|
|
"{\n\t"
|
|
".reg .pred p;\n\t"
|
|
"setp.ne.b32 p, %4, 0;\n\t"
|
|
"tcgen05.mma.ws.cta_group::1.kind::f16 [%0], %1, %2, %3, p; \n\t"
|
|
"}\n"
|
|
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
|
}
|
|
};
|
|
|
|
/// Tensor memory operations
|
|
CUTLASS_DEVICE void tcgen05_before_thread_sync() {
|
|
asm volatile("tcgen05.fence::before_thread_sync;");
|
|
}
|
|
|
|
CUTLASS_DEVICE void tcgen05_after_thread_sync() {
|
|
asm volatile("tcgen05.fence::after_thread_sync;");
|
|
}
|
|
|
|
} // namespace deep_gemm::ptx
|