template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 16, trans_a, trans_b> {
    template<int scale_b=1> __device__ static inline void rt_st(
        rt<T_D, 16, 16, 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>),
            "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, %13, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7}, " \
                "{%8, %9, %10, %11}, " \
                "%12, " \
                "p, 1, %15, %14;\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)

            :   "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, %13, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7}, " \
                "{%8, %9, %10, %11}, " \
                "%12, " \
                "p, 1, %15, %14;\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)

            :   "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, %9, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
                "{%0, %1, %2, %3}, " \
                "{%4, %5, %6, %7}, " \
                "%8, " \
                "p, 1, %11, %10;\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*)&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, 16, 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>),
            "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, %10, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7}, " \
                "%8, " \
                "%9, " \
                "p, 1, %13, %11, %12;\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)

            :   "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, %10, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
                "{%0, %1, %2, %3, %4, %5, %6, %7}, " \
                "%8, " \
                "%9, " \
                "p, 1, %13, %11, %12;\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)

            :   "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, %6, 0;\n" \
                "wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
                "{%0, %1, %2, %3}, " \
                "%4, " \
                "%5, " \
                "p, 1, %9, %7, %8;\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])

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