template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 64, trans_a, trans_b> {
    template<int scale_b=1> __device__ static inline void rt_st(
        rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
        const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
        const uint64_t b_st_desc,
        int scale_d = 1
    ) {
        static_assert(
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
            (std::is_same_v<T_D, half>  && std::is_same_v<T_AB, half>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
            (std::is_same_v<T_D, half>  && std::is_same_v<T_AB, fp8e4m3>) ||
            (std::is_same_v<T_D, half>  && std::is_same_v<T_AB, fp8e5m2>),
            "Invalid type combination for WGMMA."
        );
        static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
        // ----- BF16,BF16 -> FP32 ----- //
        if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %37, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "{%32, %33, %34, %35}, " \
                "%36, " \
                "p, 1, %39, %38;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
            );
        }
        // ----- FP16,FP16 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %37, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "{%32, %33, %34, %35}, " \
                "%36, " \
                "p, 1, %39, %38;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
            );
        }
        // ----- FP16,FP16 -> FP16 ----- //
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %21, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "{%16, %17, %18, %19}, " \
                "%20, " \
                "p, 1, %23, %22;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %37, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "{%32, %33, %34, %35}, " \
                "%36, " \
                "p, 1, %38;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), 
                "r"(scale_d), 
                // "n"(trans_b), 
                "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %37, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "{%32, %33, %34, %35}, " \
                "%36, " \
                "p, 1, %38;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), 
                "r"(scale_d), 
                // "n"(trans_b), 
                "n"(scale_b)
            );
        }
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %21, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "{%16, %17, %18, %19}, " \
                "%20, " \
                "p, 1, %22;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), 
                "r"(scale_d), 
                // "n"(trans_b), 
                "n"(scale_b)
            );
        }
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %21, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "{%16, %17, %18, %19}, " \
                "%20, " \
                "p, 1, %22;\n" \
                "}\n"
                // a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
                "r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
                
                "l"(b_st_desc), "r"(scale_d), 
                // "n"(trans_b), 
                "n"(scale_b)
            );
        }
    }
    template<int scale_b=1> __device__ static inline void st_st(
        rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
        const uint64_t a_st_desc,
        const uint64_t b_st_desc,
        int scale_d = 1
    ) {
        static_assert(
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
            (std::is_same_v<T_D, half>  && std::is_same_v<T_AB, half>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
            (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
            (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
            (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
            "Invalid type combination for WGMMA."
        );
        static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
        // ----- BF16,BF16 -> FP32 ----- //
        if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %34, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "%32, " \
                "%33, " \
                "p, 1, %37, %35, %36;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                "n"(trans_a),
                "n"(trans_b),
                "n"(scale_b)
            );
        }
        // ----- FP16,FP16 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %34, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "%32, " \
                "%33, " \
                "p, 1, %37, %35, %36;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                "n"(trans_a),
                "n"(trans_b),
                "n"(scale_b)
            );
        }
        // ----- FP16,FP16 -> FP16 ----- //
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %18, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "%16, " \
                "%17, " \
                "p, 1, %21, %19, %20;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                "n"(trans_a),
                "n"(trans_b),
                "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %34, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "%32, " \
                "%33, " \
                "p, 1, %35;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                // "n"(trans_a),
                // "n"(trans_b), // transpose is not supported for FP8
                "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP32 ----- //
        else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %34, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
                "%32, " \
                "%33, " \
                "p, 1, %35;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
                "+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
                "+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
                "+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
                "+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
                "+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
                "+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
                "+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
                "+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
                "+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
                "+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
                "+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
                "+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
                "+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
                "+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
                "+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                // "n"(trans_a),
                // "n"(trans_b), // transpose is not supported for FP8
                "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP16 ----- //
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %18, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "%16, " \
                "%17, " \
                "p, 1, %19;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                // "n"(trans_a),
                // "n"(trans_b),
                "n"(scale_b)
            );
        }
        // ----- FP8,FP8 -> FP16 ----- //
        else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
            asm volatile (
                "{\n"
                ".reg .pred p;\n" \
                "setp.ne.b32 p, %18, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
                "%16, " \
                "%17, " \
                "p, 1, %19;\n" \
                "}\n"
                // a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b

            :   "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
                "+r"(*(uint32_t*)&dst.tiles[0][3].data[3])

            :   "l"(a_st_desc),
                "l"(b_st_desc),
            
                "r"(scale_d),
                // "n"(trans_a),
                // "n"(trans_b),
                "n"(scale_b)
            );
        }
    }
};