@@ -698,6 +698,70 @@ MTL::ComputePipelineState* get_steel_gemm_segmented_kernel(
698698 return d.get_kernel (kernel_name, lib, hash_name, func_consts);
699699}
700700
701+ MTL ::ComputePipelineState* get_gemv_kernel (
702+ metal::Device& d,
703+ const std::string& kernel_name,
704+ const array& out,
705+ bool transpose_mat,
706+ int bm,
707+ int bn,
708+ int sm,
709+ int sn,
710+ int tm,
711+ int tn,
712+ bool nc,
713+ bool axpby) {
714+ const auto & lib_name = kernel_name;
715+ auto lib = d.get_library (lib_name, [&]() {
716+ std::ostringstream kernel_source;
717+ kernel_source << metal::gemv ()
718+ << get_template_definition (
719+ lib_name,
720+ transpose_mat ? " gemv_t" : " gemv" ,
721+ get_type_string (out.dtype ()),
722+ bm,
723+ bn,
724+ sm,
725+ sn,
726+ tm,
727+ tn,
728+ nc ? 1 : 0 ,
729+ axpby ? 1 : 0 );
730+ return kernel_source.str ();
731+ });
732+ return d.get_kernel (kernel_name, lib);
733+ }
734+
735+ MTL ::ComputePipelineState* get_gemv_gather_kernel (
736+ metal::Device& d,
737+ const std::string& kernel_name,
738+ const array& out,
739+ bool transpose_mat,
740+ int bm,
741+ int bn,
742+ int sm,
743+ int sn,
744+ int tm,
745+ int tn) {
746+ const auto & lib_name = kernel_name;
747+ auto lib = d.get_library (lib_name, [&]() {
748+ std::ostringstream kernel_source;
749+ kernel_source << metal::gemv ()
750+ << get_template_definition (
751+ lib_name,
752+ transpose_mat ? " gemv_t_gather" : " gemv_gather" ,
753+ get_type_string (out.dtype ()),
754+ bm,
755+ bn,
756+ sm,
757+ sn,
758+ tm,
759+ tn);
760+ return kernel_source.str ();
761+ });
762+ return d.get_kernel (kernel_name, lib);
763+ }
764+
701765MTL ::ComputePipelineState* get_gemv_masked_kernel (
702766 metal::Device& d,
703767 const std::string& kernel_name,
0 commit comments