Brunobkr/llama.cpp_AlgMor24_github
ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.
03.1k
1#include "ggml-vulkan.h"2#include <vulkan/vulkan_core.h>3#if defined(GGML_VULKAN_RUN_TESTS) || defined(GGML_VULKAN_CHECK_RESULTS)4#include <chrono>5#include "ggml-cpu.h"6#endif7 8// See https://github.com/KhronosGroup/Vulkan-Hpp?tab=readme-ov-file#extensions--per-device-function-pointers-9#define VULKAN_HPP_DISPATCH_LOADER_DYNAMIC 110// We use VULKAN_HPP_DEFAULT_DISPATCHER, but not VULKAN_HPP_DEFAULT_DISPATCH_LOADER_DYNAMIC_STORAGE11// to avoid conflicts with applications or other libraries who might use it.12#if VK_HEADER_VERSION >= 30113namespace vk::detail { class DispatchLoaderDynamic; }14using vk::detail::DispatchLoaderDynamic;15#else16namespace vk { class DispatchLoaderDynamic; }17using vk::DispatchLoaderDynamic;18#endif19DispatchLoaderDynamic & ggml_vk_default_dispatcher();20#define VULKAN_HPP_DEFAULT_DISPATCHER ggml_vk_default_dispatcher()21 22#include <vulkan/vulkan.hpp>23 24// Fallback definitions for VK_NV_cooperative_matrix_decode_vector in case the25// installed Vulkan headers predate the extension.26#ifndef VK_NV_cooperative_matrix_decode_vector27#define VK_NV_cooperative_matrix_decode_vector 128#define VK_NV_COOPERATIVE_MATRIX_DECODE_VECTOR_EXTENSION_NAME "VK_NV_cooperative_matrix_decode_vector"29#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_COOPERATIVE_MATRIX_DECODE_VECTOR_FEATURES_NV ((VkStructureType)1000689000)30typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV {31 VkStructureType sType;32 void* pNext;33 VkBool32 cooperativeMatrixDecodeVector;34} VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV;35#endif36 37// SPIR-V Headers: different SDK installations expose different include paths.38// LunarG Vulkan SDK on Windows typically provides <spirv-headers/spirv.hpp>.39// Linux packages, MSYS2 and MinGW often use the Khronos layout <spirv/unified1/spirv.hpp>.40#if __has_include(<spirv/unified1/spirv.hpp>)41# include <spirv/unified1/spirv.hpp>42#elif __has_include(<spirv-headers/spirv.hpp>)43# include <spirv-headers/spirv.hpp>44#elif __has_include(<spirv.hpp>)45# include <spirv.hpp>46#else47 // Fallback to let the compiler throw a standard "file not found" error48# include <spirv/unified1/spirv.hpp>49#endif50 51#include <algorithm>52#include <cmath>53#include <iomanip>54#include <iostream>55#include <tuple>56#include <vector>57#include <deque>58#include <sstream>59#include <utility>60#include <memory>61#include <limits>62#include <map>63#include <set>64#include <unordered_map>65#include <shared_mutex>66#include <mutex>67#include <future>68#include <condition_variable>69#include <thread>70 71#if defined(_MSC_VER)72# define NOMINMAX 173# include <windows.h>74# define YIELD() YieldProcessor()75#elif defined(__clang__) || defined(__GNUC__)76# if defined(__x86_64__) ||defined(__i386__)77# include <immintrin.h>78# define YIELD() _mm_pause()79# elif defined(__arm__) || defined(__aarch64__)80# if defined(__clang__)81# include <arm_acle.h>82# define YIELD() __yield()83# else84# define YIELD() asm volatile("yield")85# endif86# endif87#endif88 89#if !defined(YIELD)90#define YIELD()91#endif92 93#include "ggml-impl.h"94#include "ggml-backend-impl.h"95 96#include "ggml-vulkan-shaders.hpp"97 98// remove this once it's more widely available in the SDK99#if !defined(VK_KHR_shader_bfloat16)100 101#define VK_KHR_shader_bfloat16 1102#define VK_KHR_SHADER_BFLOAT16_SPEC_VERSION 1103#define VK_KHR_SHADER_BFLOAT16_EXTENSION_NAME "VK_KHR_shader_bfloat16"104#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_BFLOAT16_FEATURES_KHR ((VkStructureType)1000141000)105#define VK_COMPONENT_TYPE_BFLOAT16_KHR ((VkComponentTypeKHR)1000141000)106 107typedef struct VkPhysicalDeviceShaderBfloat16FeaturesKHR {108 VkStructureType sType;109 void* pNext;110 VkBool32 shaderBFloat16Type;111 VkBool32 shaderBFloat16DotProduct;112 VkBool32 shaderBFloat16CooperativeMatrix;113} VkPhysicalDeviceShaderBfloat16FeaturesKHR;114#endif115 116#if !defined(VK_VALVE_shader_mixed_float_dot_product)117#define VK_VALVE_shader_mixed_float_dot_product 1118#define VK_VALVE_SHADER_MIXED_FLOAT_DOT_PRODUCT_SPEC_VERSION 1119#define VK_VALVE_SHADER_MIXED_FLOAT_DOT_PRODUCT_EXTENSION_NAME "VK_VALVE_shader_mixed_float_dot_product"120#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_MIXED_FLOAT_DOT_PRODUCT_FEATURES_VALVE ((VkStructureType)1000673000)121typedef struct VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE {122 VkStructureType sType;123 void* pNext;124 VkBool32 shaderMixedFloatDotProductFloat16AccFloat32;125 VkBool32 shaderMixedFloatDotProductFloat16AccFloat16;126 VkBool32 shaderMixedFloatDotProductBFloat16Acc;127 VkBool32 shaderMixedFloatDotProductFloat8AccFloat32;128} VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE;129#endif130 131#if !defined(VK_EXT_shader_ocp_microscaling_types)132#define VK_EXT_shader_ocp_microscaling_types 1133#define VK_EXT_SHADER_OCP_MICROSCALING_TYPES_SPEC_VERSION 1134#define VK_EXT_SHADER_OCP_MICROSCALING_TYPES_EXTENSION_NAME "VK_EXT_shader_ocp_microscaling_types"135#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_OCP_MICROSCALING_TYPES_FEATURES_EXT ((VkStructureType)1000672000)136typedef struct VkPhysicalDeviceShaderOCPMicroscalingTypesFeaturesEXT {137 VkStructureType sType;138 void* pNext;139 VkBool32 shaderFloat4;140 VkBool32 shaderFloat6;141 VkBool32 shaderFloat8UnsignedE8M0;142 VkBool32 shaderMXInt8;143} VkPhysicalDeviceShaderOCPMicroscalingTypesFeaturesEXT;144#endif145 146#if !defined(VK_EXT_shader_float8)147#define VK_EXT_shader_float8 1148#define VK_EXT_SHADER_FLOAT8_SPEC_VERSION 1149#define VK_EXT_SHADER_FLOAT8_EXTENSION_NAME "VK_EXT_shader_float8"150#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_FLOAT8_FEATURES_EXT ((VkStructureType)1000567000)151typedef struct VkPhysicalDeviceShaderFloat8FeaturesEXT {152 VkStructureType sType;153 void* pNext;154 VkBool32 shaderFloat8;155 VkBool32 shaderFloat8CooperativeMatrix;156} VkPhysicalDeviceShaderFloat8FeaturesEXT;157#endif158 159#ifndef VK_KHR_INTERNALLY_SYNCHRONIZED_QUEUES_EXTENSION_NAME160#define VK_KHR_INTERNALLY_SYNCHRONIZED_QUEUES_EXTENSION_NAME "VK_KHR_internally_synchronized_queues"161#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_INTERNALLY_SYNCHRONIZED_QUEUES_FEATURES_KHR ((VkStructureType)1000504000)162#define VK_DEVICE_QUEUE_CREATE_INTERNALLY_SYNCHRONIZED_BIT_KHR ((VkDeviceQueueCreateFlagBits)0x00000004)163 164// Compile-time constant guaranteed; no runtime initialization overhead165static constexpr vk::DeviceQueueCreateFlagBits eInternallySynchronizedKHR =166 static_cast<vk::DeviceQueueCreateFlagBits>(0x00000004);167 168typedef struct VkPhysicalDeviceInternallySynchronizedQueuesFeaturesKHR {169 VkStructureType sType;170 void* pNext;171 VkBool32 internallySynchronizedQueues;172} VkPhysicalDeviceInternallySynchronizedQueuesFeaturesKHR;173#else174static constexpr vk::DeviceQueueCreateFlagBits eInternallySynchronizedKHR = vk::DeviceQueueCreateFlagBits::eInternallySynchronizedKHR;175#endif176 177#define ROUNDUP_POW2(M, N) (((M) + (N) - 1) & ~((N) - 1))178#define CEIL_DIV(M, N) (((M) / (N)) + (((M) % (N)) != 0))179static bool is_pow2(uint32_t x) { return x > 1 && (x & (x-1)) == 0; }180 181#define VK_VENDOR_ID_AMD 0x1002182#define VK_VENDOR_ID_APPLE 0x106b183#define VK_VENDOR_ID_INTEL 0x8086184#define VK_VENDOR_ID_NVIDIA 0x10de185#define VK_VENDOR_ID_QUALCOMM 0x5143186 187#define VK_DEVICE_DESCRIPTOR_POOL_SIZE 256188 189#define VK_CHECK(err, msg, dev) \190 do { \191 vk::Result err_; \192 try { \193 err_ = (err); \194 } catch (vk::DeviceLostError &) { \195 ggml_vk_print_device_lost_info(dev); \196 GGML_LOG_ERROR("ggml_vulkan: %s at %s:%d\n", \197 #err, __FILE__, __LINE__); \198 throw; \199 } \200 if (err_ != vk::Result::eSuccess) { \201 GGML_LOG_ERROR("ggml_vulkan: %s error %s at %s:%d\n", \202 #err, to_string(err_).c_str(), __FILE__, __LINE__); \203 throw vk::SystemError(vk::make_error_code(err_), \204 "ggml_vulkan: " msg); \205 } \206 } while (0)207 208#ifdef GGML_VULKAN_DEBUG209#define VK_LOG_DEBUG(msg) std::cerr << msg << std::endl210#else211#define VK_LOG_DEBUG(msg) ((void) 0)212#endif // GGML_VULKAN_DEBUG213 214struct ggml_backend_vk_context;215 216#define MAX_PARAMETER_COUNT 12217// Max number of adds that can be fused without exceeding MAX_PARAMETER_COUNT.218#define MAX_FUSED_ADDS (MAX_PARAMETER_COUNT - 3)219 220typedef std::shared_ptr<struct vk_pipeline_struct> vk_pipeline;221 222struct vk_pipeline_struct {223 std::string name;224 vk::ShaderModule shader_module;225 vk::PipelineLayout layout;226 vk::Pipeline pipeline;227 uint32_t push_constant_size;228 uint32_t parameter_count;229 std::array<uint32_t, 3> wg_denoms;230 uint32_t align;231 // true if fields have been set by ggml_vk_create_pipeline232 bool initialized {};233 // true while a compile is in flight, used to dedupe concurrent claims.234 // Protected by device->compile_mutex.235 bool compile_pending {};236 // set to true when the shader has been compiled237 std::atomic<bool> compiled {};238 // number of registers used, extracted from pipeline executable properties239 uint32_t register_count {};240 241#if defined(VK_EXT_shader_64bit_indexing)242 bool is_64b_indexing {};243#endif244 // linked list of pipelines for multiple compilation variants.245 // currently only used to compile a 64-bit indexing variant.246 vk_pipeline next;247};248 249typedef std::weak_ptr<vk_pipeline_struct> vk_pipeline_ref;250 251static void ggml_vk_destroy_pipeline(vk::Device& device, vk_pipeline& pipeline);252 253struct vk_matmul_pipeline_struct {254 vk_pipeline l, m, s;255 vk_pipeline a_l, a_m, a_s;256 // Returns true when all unaligned pipelines are null.257 // We only check for unaligned variants since one of the unaligned pipelines must exist258 // while aligned pipelines are optional259 bool is_empty() const {260 return l == nullptr && m == nullptr && s == nullptr;261 }262};263typedef std::shared_ptr<vk_matmul_pipeline_struct> vk_matmul_pipeline;264 265struct vk_matmul_pipeline2 {266 vk_matmul_pipeline2() {267 f16acc = std::make_shared<vk_matmul_pipeline_struct>();268 f32acc = std::make_shared<vk_matmul_pipeline_struct>();269 }270 vk_matmul_pipeline f32acc;271 vk_matmul_pipeline f16acc;272};273 274struct vk_device_struct;275typedef std::shared_ptr<vk_device_struct> vk_device;276typedef std::weak_ptr<vk_device_struct> vk_device_ref;277 278struct vk_buffer_struct;279typedef std::shared_ptr<vk_buffer_struct> vk_buffer;280typedef std::weak_ptr<vk_buffer_struct> vk_buffer_ref;281 282struct ggml_backend_vk_buffer_type_context {283 std::string name;284 vk_device device;285};286 287struct vk_queue;288 289struct vk_command_buffer {290 vk::CommandBuffer buf;291 uint64_t use_counter = 0;292 bool in_use = false;293};294 295// Stores command pool/buffers. There's an instance of this296// for each (context,queue) pair and for each (device,queue) pair.297struct vk_command_pool {298 void init(vk_device& device, vk_queue *q_);299 void destroy(vk::Device& device);300 301 vk::CommandPool pool;302 // Using deque so the pointers to command buffers303 // remain valid even if we add more304 std::deque<vk_command_buffer> cmd_buffers;305 306 vk_queue *q;307 308 size_t buffers_in_use() const {309 return std::count_if(cmd_buffers.begin(), cmd_buffers.end(),310 [](const auto& cb) { return cb.in_use; });311 }312};313 314static void ggml_vk_print_device_fault_info(const vk_device& device);315static void ggml_vk_print_device_lost_info(const vk_device& device);316 317// Prevent simultaneous submissions to the same queue.318struct vk_queue_handle {319 vk::Queue queue;320 vk_device_ref device;321 virtual void submit(vk::ArrayProxy<const vk::SubmitInfo> submits, vk::Fence fence) = 0;322 virtual void lock() {} // no-op by default (internally synchronized case)323 virtual void unlock() {}324 virtual ~vk_queue_handle() = default;325};326 327struct vk_queue_handle_synchronized : vk_queue_handle {328 std::mutex mutex;329 void submit(vk::ArrayProxy<const vk::SubmitInfo> submits, vk::Fence fence) override {330 std::lock_guard<std::mutex> guard(mutex);331 try {332 queue.submit(submits, fence);333 } catch (vk::DeviceLostError &) {334 if (auto dev = device.lock()) {335 ggml_vk_print_device_lost_info(dev);336 }337 throw;338 }339 }340 void lock() override { mutex.lock(); }341 void unlock() override { mutex.unlock(); }342};343 344struct vk_queue_handle_unsynchronized : vk_queue_handle {345 void submit(vk::ArrayProxy<const vk::SubmitInfo> submits, vk::Fence fence) override {346 // Driver guarantees internal synchronization via VK_KHR_internally_synchronized_queues347 try {348 queue.submit(submits, fence);349 } catch (vk::DeviceLostError &) {350 if (auto dev = device.lock()) {351 ggml_vk_print_device_lost_info(dev);352 }353 throw;354 }355 }356 // lock()/unlock() inherited no-ops357};358 359struct vk_queue {360 uint32_t queue_family_index;361 std::shared_ptr<vk_queue_handle> handle;362 363 vk_command_pool cmd_pool;364 365 vk::PipelineStageFlags stage_flags;366 367 bool transfer_only;368};369 370static const char * ggml_backend_vk_buffer_type_name(ggml_backend_buffer_type_t buft);371static ggml_backend_buffer_t ggml_backend_vk_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size);372static size_t ggml_backend_vk_buffer_type_get_alignment(ggml_backend_buffer_type_t buft);373static size_t ggml_backend_vk_buffer_type_get_max_size(ggml_backend_buffer_type_t buft);374static size_t ggml_backend_vk_buffer_type_get_alloc_size(ggml_backend_buffer_type_t buft, const ggml_tensor * tensor);375static ggml_backend_buffer_type_i ggml_backend_vk_buffer_type_interface = {376 /* .get_name = */ ggml_backend_vk_buffer_type_name,377 /* .alloc_buffer = */ ggml_backend_vk_buffer_type_alloc_buffer,378 /* .get_alignment = */ ggml_backend_vk_buffer_type_get_alignment,379 /* .get_max_size = */ ggml_backend_vk_buffer_type_get_max_size,380 /* .get_alloc_size = */ ggml_backend_vk_buffer_type_get_alloc_size,381 /* .is_host = */ NULL,382};383 384class vk_memory_logger;385class vk_perf_logger;386static void ggml_vk_destroy_buffer(vk_buffer& buf);387static void ggml_vk_synchronize(ggml_backend_vk_context * ctx);388 389static constexpr uint32_t mul_mat_vec_max_cols = 8;390static constexpr uint32_t p021_max_gqa_ratio = 8;391 392enum vk_device_architecture {393 OTHER,394 AMD_GCN,395 AMD_RDNA1,396 AMD_RDNA2,397 AMD_RDNA3,398 INTEL_XE1,399 INTEL_XE2,400 NVIDIA_PRE_TURING,401 NVIDIA_TURING,402};403 404static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& device) {405 vk::PhysicalDeviceProperties props = device.getProperties();406 407 if (props.vendorID == VK_VENDOR_ID_AMD) {408 const std::vector<vk::ExtensionProperties> ext_props = device.enumerateDeviceExtensionProperties();409 410 bool amd_shader_core_properties = false;411 bool integer_dot_product = false;412 bool subgroup_size_control = false;413 414 for (const auto& properties : ext_props) {415 if (strcmp("VK_AMD_shader_core_properties", properties.extensionName) == 0) {416 amd_shader_core_properties = true;417 } else if (strcmp("VK_KHR_shader_integer_dot_product", properties.extensionName) == 0) {418 integer_dot_product = true;419 } else if (strcmp("VK_EXT_subgroup_size_control", properties.extensionName) == 0) {420 subgroup_size_control = true;421 }422 }423 424 if (!amd_shader_core_properties || !integer_dot_product || !subgroup_size_control) {425 return vk_device_architecture::OTHER;426 }427 428 vk::PhysicalDeviceProperties2 props2;429 vk::PhysicalDeviceShaderCorePropertiesAMD shader_core_props_amd;430 vk::PhysicalDeviceShaderIntegerDotProductPropertiesKHR integer_dot_props;431 vk::PhysicalDeviceSubgroupSizeControlPropertiesEXT subgroup_size_control_props;432 433 props2.pNext = &shader_core_props_amd;434 shader_core_props_amd.pNext = &integer_dot_props;435 integer_dot_props.pNext = &subgroup_size_control_props;436 437 device.getProperties2(&props2);438 439 if (subgroup_size_control_props.maxSubgroupSize == 64 && subgroup_size_control_props.minSubgroupSize == 64) {440 return vk_device_architecture::AMD_GCN;441 }442 if (subgroup_size_control_props.maxSubgroupSize == 64 && subgroup_size_control_props.minSubgroupSize == 32) {443 // RDNA444 if (shader_core_props_amd.wavefrontsPerSimd == 20) {445 return vk_device_architecture::AMD_RDNA1;446 }447 if (integer_dot_props.integerDotProduct4x8BitPackedMixedSignednessAccelerated) {448 return vk_device_architecture::AMD_RDNA3;449 }450 return vk_device_architecture::AMD_RDNA2;451 }452 } else if (props.vendorID == VK_VENDOR_ID_INTEL) {453 const std::vector<vk::ExtensionProperties> ext_props = device.enumerateDeviceExtensionProperties();454 455 bool subgroup_size_control = false;456 bool integer_dot_product = false;457 458 for (const auto& properties : ext_props) {459 if (strcmp("VK_EXT_subgroup_size_control", properties.extensionName) == 0) {460 subgroup_size_control = true;461 } else if (strcmp("VK_KHR_shader_integer_dot_product", properties.extensionName) == 0) {462 integer_dot_product = true;463 }464 }465 466 if (!subgroup_size_control || !integer_dot_product) {467 return vk_device_architecture::OTHER;468 }469 470 vk::PhysicalDeviceProperties2 props2;471 vk::PhysicalDeviceSubgroupSizeControlPropertiesEXT subgroup_size_control_props;472 vk::PhysicalDeviceShaderIntegerDotProductPropertiesKHR integer_dot_props;473 474 props2.pNext = &subgroup_size_control_props;475 subgroup_size_control_props.pNext = &integer_dot_props;476 device.getProperties2(&props2);477 478 if (subgroup_size_control_props.minSubgroupSize == 16) {479 // Xe2 architecture uses SIMD16 while previous Xe and Gen architecture uses SIMD8.480 // Minimum subgroup size matches the SIMD width so we distinguish architecture by checking this value.481 // https://www.intel.com/content/www/us/en/content-details/824434/2024-intel-tech-tour-xe2-and-lunar-lake-s-gpu.html482 // https://www.intel.com/content/www/us/en/docs/oneapi/optimization-guide-gpu/2025-0/intel-xe-gpu-architecture.html483 return vk_device_architecture::INTEL_XE2;484 } else if (subgroup_size_control_props.minSubgroupSize == 8 &&485 integer_dot_product && integer_dot_props.integerDotProduct4x8BitPackedSignedAccelerated) {486 return vk_device_architecture::INTEL_XE1;487 }488 } else if (props.vendorID == VK_VENDOR_ID_NVIDIA) {489 const std::vector<vk::ExtensionProperties> ext_props = device.enumerateDeviceExtensionProperties();490 491 bool cooperative_matrix = false;492 bool sm_builtins = false;493 494 // Detect "pre-turing" based on lack of coopmat support.495 for (const auto& properties : ext_props) {496 if (strcmp("VK_KHR_cooperative_matrix", properties.extensionName) == 0) {497 cooperative_matrix = true;498 } else if (strcmp("VK_NV_shader_sm_builtins", properties.extensionName) == 0) {499 sm_builtins = true;500 }501 }502 503 if (!cooperative_matrix) {504 return vk_device_architecture::NVIDIA_PRE_TURING;505 }506 507 if (sm_builtins) {508 vk::PhysicalDeviceProperties2 props2;509 vk::PhysicalDeviceShaderSMBuiltinsPropertiesNV sm_props;510 511 props2.pNext = &sm_props;512 513 device.getProperties2(&props2);514 515 // Turing has 32, following architectures have 48516 if (sm_props.shaderWarpsPerSM == 32) {517 return vk_device_architecture::NVIDIA_TURING;518 }519 }520 }521 return vk_device_architecture::OTHER;522}523 524enum vk_conv_shapes {525 CONV_SHAPE_128x128,526 CONV_SHAPE_64x32,527 CONV_SHAPE_32x256,528 CONV_SHAPE_64x128,529 CONV_SHAPE_COUNT,530};531 532struct vk_conv_block_size {533 uint32_t K;534 uint32_t NPQ;535 uint32_t CRS;536};537 538vk_conv_block_size vk_conv_block_sizes[CONV_SHAPE_COUNT] = {539 // K NPQ CRS540 { 128, 128, 16 }, // CONV_SHAPE_128x128541 { 64, 32, 32 }, // CONV_SHAPE_64x32542 { 32, 256, 16 }, // CONV_SHAPE_32x256543 { 64, 128, 16 }, // CONV_SHAPE_64x128544};545 546enum dmmv_wg_sizes {547 DMMV_WG_SIZE_SUBGROUP,548 DMMV_WG_SIZE_LARGE,549 DMMV_WG_SIZE_COUNT,550};551 552enum FaCodePath {553 FA_SCALAR,554 FA_COOPMAT1,555 FA_COOPMAT2,556};557 558struct vk_fa_pipeline_state {559 uint32_t HSK, HSV;560 uint32_t Br, Bc;561 uint32_t D_split, row_split;562 bool shmem_staging;563 FaCodePath path;564 uint32_t workgroup_size, subgroup_size;565 bool aligned;566 bool f32acc;567 uint32_t flags;568 uint32_t limit_occupancy_shmem;569 ggml_type k_type;570 ggml_type v_type;571 572 bool operator<(const vk_fa_pipeline_state &b) const {573 return std::tie(HSK, HSV, Br, Bc, D_split, row_split, shmem_staging, path, workgroup_size, subgroup_size, aligned, f32acc, flags, limit_occupancy_shmem, k_type, v_type) <574 std::tie(b.HSK, b.HSV, b.Br, b.Bc, b.D_split, b.row_split, b.shmem_staging, b.path, b.workgroup_size, b.subgroup_size, b.aligned, b.f32acc, b.flags, b.limit_occupancy_shmem, b.k_type, b.v_type);575 }576};577 578struct vk_conv2d_pipeline_state {579 vk_conv2d_pipeline_state(uint32_t s0, uint32_t s1, uint32_t p0, uint32_t p1, uint32_t d0, uint32_t d1, uint32_t KW, uint32_t KH, uint32_t aligned)580 : s0(s0), s1(s1), p0(p0), p1(p1), d0(d0), d1(d1), KW(KW), KH(KH), aligned(aligned) {}581 582 uint32_t s0, s1, p0, p1, d0, d1, KW, KH;583 // when set, shader can skip K/CRS/NPQ bounds checks and address clamps584 uint32_t aligned;585 586 bool operator<(const vk_conv2d_pipeline_state &b) const {587 return std::tie(s0, s1, p0, p1, d0, d1, KW, KH, aligned) <588 std::tie(b.s0, b.s1, b.p0, b.p1, b.d0, b.d1, b.KW, b.KH, b.aligned);589 }590};591 592struct vk_conv3d_pipeline_state {593 vk_conv3d_pipeline_state(uint32_t s0, uint32_t s1, uint32_t s2, uint32_t p0, uint32_t p1, uint32_t p2,594 uint32_t d0, uint32_t d1, uint32_t d2, uint32_t KW, uint32_t KH, uint32_t KD, uint32_t aligned)595 : s0(s0), s1(s1), s2(s2), p0(p0), p1(p1), p2(p2), d0(d0), d1(d1), d2(d2), KW(KW), KH(KH), KD(KD), aligned(aligned) {}596 597 uint32_t s0, s1, s2, p0, p1, p2, d0, d1, d2, KW, KH, KD;598 uint32_t aligned;599 600 bool operator<(const vk_conv3d_pipeline_state &b) const {601 return std::tie(s0, s1, s2, p0, p1, p2, d0, d1, d2, KW, KH, KD, aligned) <602 std::tie(b.s0, b.s1, b.s2, b.p0, b.p1, b.p2, b.d0, b.d1, b.d2, b.KW, b.KH, b.KD, b.aligned);603 }604};605 606struct vk_solve_tri_pipeline_state {607 vk_solve_tri_pipeline_state(uint32_t N, uint32_t K)608 : N(N), K(K) {}609 610 uint32_t N, K;611 612 bool operator<(const vk_solve_tri_pipeline_state &b) const {613 return std::tie(N, K) <614 std::tie(b.N, b.K);615 }616};617 618enum shader_reduction_mode {619 SHADER_REDUCTION_MODE_SHMEM,620 SHADER_REDUCTION_MODE_HYBRID,621 SHADER_REDUCTION_MODE_SUBGROUP,622 SHADER_REDUCTION_MODE_COUNT,623};624 625// argsort pipelines for up to 1<<10 invocations per workgroup626static constexpr uint32_t num_argsort_pipelines = 11;627static constexpr uint32_t num_topk_moe_pipelines = 10;628static constexpr uint32_t num_topk_pipelines = 11;629 630static constexpr std::initializer_list<ggml_op> topk_moe_early_softmax_norm{ GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT,631 GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_RESHAPE,632 GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV,633 GGML_OP_RESHAPE };634 635static constexpr std::initializer_list<ggml_op> topk_moe_sigmoid_norm_bias{ GGML_OP_UNARY, GGML_OP_RESHAPE, GGML_OP_ADD,636 GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS,637 GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP,638 GGML_OP_DIV, GGML_OP_RESHAPE };639 640static constexpr std::initializer_list<ggml_op> topk_moe_sqrt_softplus_norm_bias{ GGML_OP_UNARY, GGML_OP_SQRT,641 GGML_OP_RESHAPE, GGML_OP_ADD,642 GGML_OP_ARGSORT, GGML_OP_VIEW,643 GGML_OP_GET_ROWS, GGML_OP_RESHAPE,644 GGML_OP_SUM_ROWS, GGML_OP_CLAMP,645 GGML_OP_DIV, GGML_OP_RESHAPE };646 647static constexpr std::initializer_list<ggml_op> topk_moe_early_softmax { GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT,648 GGML_OP_VIEW, GGML_OP_GET_ROWS };649 650static constexpr std::initializer_list<ggml_op> topk_moe_late_softmax { GGML_OP_ARGSORT, GGML_OP_VIEW,651 GGML_OP_GET_ROWS, GGML_OP_RESHAPE,652 GGML_OP_SOFT_MAX, GGML_OP_RESHAPE };653 654// Snake activation: y = x + sin(a*x)^2 * inv_b. Used by the optimize_graph reorder655// pass so it keeps the chain contiguous and by the dispatcher to detect the fusion.656static constexpr std::initializer_list<ggml_op> snake_pattern { GGML_OP_MUL, GGML_OP_SIN,657 GGML_OP_SQR, GGML_OP_MUL,658 GGML_OP_ADD };659 660//node #978 ( SOFT_MAX): ffn_moe_probs-15 ( 0K) [Vulka ] use=2: ffn_moe_logits-15 ( 0K) [Vulka ]661//node #979 ( RESHAPE): ffn_moe_probs-15 (re ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ]662//node #980 ( ARGSORT): ffn_moe_argsort-15 ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ]663//node #981 ( VIEW): ffn_moe_topk-15 ( 0K) [Vulka ] use=4: ffn_moe_argsort-15 ( 0K) [Vulka ]664//node #982 ( GET_ROWS): ffn_moe_weights-15 ( 0K) [Vulka ] use=1: ffn_moe_probs-15 (re ( 0K) [Vulka ] ffn_moe_topk-15 ( 0K) [Vulka ]665//node #983 ( RESHAPE): ffn_moe_weights-15 ( ( 0K) [Vulka ] use=2: ffn_moe_weights-15 ( 0K) [Vulka ]666//node #984 ( SUM_ROWS): ffn_moe_weights_sum- ( 0K) [Vulka ] use=1: ffn_moe_weights-15 ( ( 0K) [Vulka ]667//node #985 ( CLAMP): ffn_moe_weights_sum_ ( 0K) [Vulka ] use=1: ffn_moe_weights_sum- ( 0K) [Vulka ]668//node #986 ( DIV): ffn_moe_weights_norm ( 0K) [Vulka ] use=1: ffn_moe_weights-15 ( ( 0K) [Vulka ] ffn_moe_weights_sum_ ( 0K) [Vulka ]669//node #987 ( RESHAPE): ffn_moe_weights_norm ( 0K) [Vulka ] use=1: ffn_moe_weights_norm ( 0K) [Vulka ]670static constexpr std::initializer_list<std::array<int, 3>> topk_moe_early_softmax_norm_edges {671 { 1, 0, 0 }, // reshape->src[0] == softmax672 { 2, 0, 0 }, // argsort->src[0] == softmax673 { 3, 0, 2 }, // view->src[0] == argsort674 { 4, 0, 1 }, // get_rows->src[0] == reshape675 { 4, 1, 3 }, // get_rows->src[1] == view676 { 5, 0, 4 }, // reshape->src[0] == get_rows677 { 6, 0, 5 }, // sum_rows->src[0] == reshape678 { 7, 0, 6 }, // clamp->src[0] == sum_rows679 { 8, 0, 5 }, // div->src[0] == reshape680 { 8, 1, 7 }, // div->src[1] == clamp681 { 9, 0, 8 }, // reshape->src[0] == div682};683 684//node #436 ( UNARY): ffn_moe_probs-10 ( 256K) [Vulka ] use=2: ffn_moe_logits-10 ( 256K) [Vulka ]685//node #437 ( RESHAPE): ffn_moe_probs-10 (re ( 256K) [Vulka ] use=1: ffn_moe_probs-10 ( 256K) [Vulka ]686//node #438 ( ADD): ffn_moe_probs_biased ( 256K) [Vulka ] use=1: ffn_moe_probs-10 ( 256K) [Vulka ] blk.10.exp_probs_b.b ( 0K) [Vulka ]687//node #439 ( ARGSORT): ffn_moe_argsort-10 ( 256K) [Vulka ] use=1: ffn_moe_probs_biased ( 256K) [Vulka ]688//node #440 ( VIEW): ffn_moe_topk-10 ( 255K) [Vulka ] use=3: ffn_moe_argsort-10 ( 256K) [Vulka ]689//node #441 ( GET_ROWS): ffn_moe_weights-10 ( 12K) [Vulka ] use=1: ffn_moe_probs-10 (re ( 256K) [Vulka ] ffn_moe_topk-10 ( 255K) [Vulka ]690//node #442 ( RESHAPE): ffn_moe_weights-10 ( ( 12K) [Vulka ] use=2: ffn_moe_weights-10 ( 12K) [Vulka ]691//node #443 ( SUM_ROWS): ffn_moe_weights_sum- ( 2K) [Vulka ] use=1: ffn_moe_weights-10 ( ( 12K) [Vulka ]692//node #444 ( CLAMP): ffn_moe_weights_sum_ ( 2K) [Vulka ] use=1: ffn_moe_weights_sum- ( 2K) [Vulka ]693//node #445 ( DIV): ffn_moe_weights_norm ( 12K) [Vulka ] use=1: ffn_moe_weights-10 ( ( 12K) [Vulka ] ffn_moe_weights_sum_ ( 2K) [Vulka ]694//node #446 ( RESHAPE): ffn_moe_weights_norm ( 12K) [Vulka ] use=1: ffn_moe_weights_norm ( 12K) [Vulka ]695static constexpr std::initializer_list<std::array<int, 3>> topk_moe_sigmoid_norm_bias_edges {696 { 1, 0, 0 }, // reshape->src[0] == sigmoid697 { 2, 0, 0 }, // add->src[0] == sigmoid698 { 3, 0, 2 }, // argsort->src[0] == add699 { 4, 0, 3 }, // view->src[0] == argsort700 { 5, 0, 1 }, // get_rows->src[0] == reshape701 { 5, 1, 4 }, // get_rows->src[1] == view702 { 6, 0, 5 }, // reshape->src[0] == get_rows703 { 7, 0, 6 }, // sum_rows->src[0] == reshape704 { 8, 0, 7 }, // clamp->src[0] == sum_rows705 { 9, 0, 6 }, // div->src[0] == reshape706 { 9, 1, 8 }, // div->src[1] == clamp707 {10, 0, 9 }, // reshape->src[0] == div708};709 710static constexpr std::initializer_list<std::array<int, 3>> topk_moe_sqrt_softplus_norm_bias_edges {711 { 1, 0, 0 }, // sqrt->src[0] == softplus712 { 2, 0, 1 }, // reshape->src[0] == sqrt713 { 3, 0, 1 }, // add->src[0] == sqrt714 { 4, 0, 3 }, // argsort->src[0] == add715 { 5, 0, 4 }, // view->src[0] == argsort716 { 6, 0, 2 }, // get_rows->src[0] == reshape717 { 6, 1, 5 }, // get_rows->src[1] == view718 { 7, 0, 6 }, // reshape->src[0] == get_rows719 { 8, 0, 7 }, // sum_rows->src[0] == reshape720 { 9, 0, 8 }, // clamp->src[0] == sum_rows721 {10, 0, 7 }, // div->src[0] == reshape722 {10, 1, 9 }, // div->src[1] == clamp723 {11, 0,10 }, // reshape->src[0] == div724};725 726// same as early_softmax_norm but ending after the get_rows727static constexpr std::initializer_list<std::array<int, 3>> topk_moe_early_softmax_edges {728 { 1, 0, 0 }, // reshape->src[0] == softmax729 { 2, 0, 0 }, // argsort->src[0] == softmax730 { 3, 0, 2 }, // view->src[0] == argsort731 { 4, 0, 1 }, // get_rows->src[0] == reshape732 { 4, 1, 3 }, // get_rows->src[1] == view733};734 735//node #652 ( ARGSORT): ffn_moe_argsort-11 ( 0K) [Vulka ] use=1: ffn_moe_probs-11 ( 0K) [Vulka ]736//node #653 ( VIEW): ffn_moe_topk-11 ( 0K) [Vulka ] use=7: ffn_moe_argsort-11 ( 0K) [Vulka ]737//node #654 ( GET_ROWS): ffn_moe_weights-11 ( 0K) [Vulka ] use=1: ffn_moe_probs-11 (re ( 0K) [Vulka ] ffn_moe_topk-11 ( 0K) [Vulka ]738//node #655 ( RESHAPE): ffn_moe_weights-11 ( ( 0K) [Vulka ] use=1: ffn_moe_weights-11 ( 0K) [Vulka ]739//node #656 ( SOFT_MAX): node_656 ( 0K) [Vulka ] use=1: ffn_moe_weights-11 ( ( 0K) [Vulka ]740//node #657 ( RESHAPE): ffn_moe_weights_soft ( 0K) [Vulka ] use=1: node_656 ( 0K) [Vulka ]741static constexpr std::initializer_list<std::array<int, 3>> topk_moe_late_softmax_edges {742 { 1, 0, 0 }, // view->src[0] == argsort743 { 2, 1, 1 }, // get_rows->src[1] == view744 { 3, 0, 2 }, // reshape->src[0] == get_rows745 { 4, 0, 3 }, // soft_max->src[0] == reshape746 { 5, 0, 4 }, // reshape->src[0] == soft_max747};748 749enum topk_moe_mode {750 TOPK_MOE_EARLY_SOFTMAX,751 TOPK_MOE_EARLY_SOFTMAX_NORM,752 TOPK_MOE_LATE_SOFTMAX,753 TOPK_MOE_SIGMOID_NORM_BIAS,754 TOPK_MOE_SQRT_SOFTPLUS_NORM_BIAS,755 TOPK_MOE_COUNT,756};757 758static constexpr std::initializer_list<std::array<int, 3>> rope_view_set_rows_edges {759 { 1, 0, 0 }, // view->src[0] == rope760 { 2, 0, 1 }, // set_rows->src[0] == view761};762 763static constexpr std::initializer_list<std::array<int, 3>> rms_norm_mul_rope_view_set_rows_edges {764 { 1, 0, 0 }, // mul->src[0] == rms765 { 2, 0, 1 }, // rope->src[0] == mul766 { 3, 0, 2 }, // view->src[0] == rope767 { 4, 0, 3 }, // set_rows->src[0] == view768};769 770 771struct vk_device_struct {772 std::recursive_mutex mutex;773 mutable std::shared_mutex pinned_memory_mutex;774 775 // Guards compile_pending, all_pipelines, and the dynamic pipeline maps776 // (flash_attn, fa_mask_opt, solve_tri, conv2d, etc). The actual compile777 // runs with no lock held, so different pipelines can compile in parallel.778 // Lock order is device->mutex -> compile_mutex, never the reverse.779 std::mutex compile_mutex;780 std::condition_variable compile_cv;781 782 vk::PhysicalDevice physical_device;783 vk::PhysicalDeviceProperties properties;784 std::string name;785 uint64_t max_memory_allocation_size;786 uint64_t max_buffer_size;787 uint64_t suballocation_block_size;788 uint64_t min_imported_host_pointer_alignment;789 bool external_memory_host {};790 bool fp16;791 bool bf16;792 bool pipeline_robustness;793 bool memory_priority;794 vk::Device device;795 uint32_t vendor_id;796 vk::DriverId driver_id;797 vk_device_architecture architecture;798 std::unique_ptr<vk_queue> compute_queue;799 std::unique_ptr<vk_queue> transfer_queue;800 bool single_queue;801 bool support_async;802 bool async_use_transfer_queue;803 bool has_internally_synchronized_queues = false;804 uint32_t subgroup_size;805 uint32_t subgroup_size_log2;806 uint32_t shader_core_count;807 bool uma;808 bool prefer_host_memory;809 bool float_controls_rte_fp16;810 bool float_controls_denorm_preserve_fp16;811 bool subgroup_basic;812 bool subgroup_arithmetic;813 bool subgroup_shuffle;814 bool subgroup_ballot;815 bool subgroup_clustered;816 bool subgroup_vote;817 bool multi_add;818 bool shader_int64;819 bool buffer_device_address;820 bool vulkan_memory_model;821 822 bool add_rms_fusion;823 uint32_t partials_binding_alignment;824 uint32_t max_nodes_per_submit;825 826 bool shader_64b_indexing;827 828 bool integer_dot_product;829 // 0: default, 1: force mmvq, -1: disable mmvq830 int32_t mmvq_mode;831 832 bool subgroup_size_control;833 uint32_t subgroup_min_size;834 uint32_t subgroup_max_size;835 bool subgroup_require_full_support;836 837 // floor(log2(maxComputeWorkGroupInvocations))838 uint32_t max_workgroup_size_log2 {};839 840 bool coopmat_support;841 bool coopmat_acc_f32_support {};842 bool coopmat_acc_f16_support {};843 bool coopmat_bf16_support {};844 bool coopmat_support_16x16x16_f16acc {};845 bool coopmat_support_16x16x16_f32acc {};846 bool coopmat1_fa_support {};847 uint32_t coopmat_m;848 uint32_t coopmat_n;849 uint32_t coopmat_k;850 851 bool coopmat_int_support;852 uint32_t coopmat_int_m;853 uint32_t coopmat_int_n;854 uint32_t coopmat_int_k;855 856 bool coopmat2;857 bool coopmat2_bf16_support {};858 bool coopmat2_decode_vector;859 860 bool dot2_f16 {};861 bool ocp_fp4 {};862 863 bool pipeline_executable_properties_support {};864 865 bool device_fault {};866 PFN_vkGetDeviceFaultInfoEXT pfn_vkGetDeviceFaultInfoEXT {};867 868 bool serialize_submissions {};869 870 const ggml_cgraph * diag_cgraph {};871 int diag_prev_start = -1;872 int diag_prev_end = -1;873 874 size_t idx;875 876 bool mul_mat_l[GGML_TYPE_COUNT];877 bool mul_mat_m[GGML_TYPE_COUNT];878 bool mul_mat_s[GGML_TYPE_COUNT];879 bool mul_mat_id_l[GGML_TYPE_COUNT];880 bool mul_mat_id_m[GGML_TYPE_COUNT];881 bool mul_mat_id_s[GGML_TYPE_COUNT];882 883 // Separate flags for the q8_1 (integer dot) mmq path, whose shader uses884 // a different shared-memory layout than the float matmul shaders.885 bool mul_mat_l_int[GGML_TYPE_COUNT];886 bool mul_mat_m_int[GGML_TYPE_COUNT];887 bool mul_mat_s_int[GGML_TYPE_COUNT];888 bool mul_mat_id_l_int[GGML_TYPE_COUNT];889 bool mul_mat_id_m_int[GGML_TYPE_COUNT];890 bool mul_mat_id_s_int[GGML_TYPE_COUNT];891 892 vk::DescriptorSetLayout dsl;893 894 vk_matmul_pipeline pipeline_matmul_f32 {};895 vk_matmul_pipeline pipeline_matmul_f32_f16 {};896 vk_matmul_pipeline pipeline_matmul_bf16 {};897 vk_matmul_pipeline2 pipeline_matmul_f16;898 vk_matmul_pipeline2 pipeline_matmul_f16_f32;899 900 vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat[GGML_TYPE_COUNT];901 vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_COUNT];902 vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_COUNT];903 904 vk_matmul_pipeline pipeline_matmul_id_f32 {};905 vk_matmul_pipeline pipeline_matmul_id_bf16 {};906 vk_matmul_pipeline2 pipeline_matmul_id_f16;907 vk_matmul_pipeline2 pipeline_matmul_id_f16_f32;908 909 vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id[GGML_TYPE_COUNT];910 vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_COUNT];911 912 vk_pipeline pipeline_matmul_split_k_reduce;913 vk_pipeline pipeline_quantize_q8_1_x4;914 915 vk_pipeline pipeline_dequant[GGML_TYPE_COUNT];916 vk_pipeline pipeline_dequant_mul_mat_vec_f32_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT][mul_mat_vec_max_cols];917 vk_pipeline pipeline_dequant_mul_mat_vec_f16_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT][mul_mat_vec_max_cols];918 vk_pipeline pipeline_dequant_mul_mat_vec_id_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT];919 920 vk_pipeline pipeline_dequant_mul_mat_vec_q8_1_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT][mul_mat_vec_max_cols];921 vk_pipeline pipeline_dequant_mul_mat_vec_id_q8_1_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT];922 923 vk_pipeline pipeline_mul_mat_vec_p021_f16_f32[p021_max_gqa_ratio];924 vk_pipeline pipeline_mul_mat_vec_nc_f16_f32;925 vk_pipeline pipeline_get_rows[GGML_TYPE_COUNT];926 vk_pipeline pipeline_get_rows_f32[GGML_TYPE_COUNT];927 vk_pipeline pipeline_get_rows_back_f32;928 vk_pipeline pipeline_acc_f32;929 vk_pipeline pipeline_set_f32;930 931 // [src0 0=fp32,1=fp16][src1 0=fp32,1=fp16][dst 0=fp32,1=fp16]932 vk_pipeline pipeline_add[2][2][2];933 vk_pipeline pipeline_add_norepeat[2][2][2];934 vk_pipeline pipeline_sub[2][2][2];935 vk_pipeline pipeline_sub_norepeat[2][2][2];936 vk_pipeline pipeline_mul[2][2][2];937 vk_pipeline pipeline_mul_norepeat[2][2][2];938 vk_pipeline pipeline_div[2][2][2];939 vk_pipeline pipeline_div_norepeat[2][2][2];940 vk_pipeline pipeline_add_rms[2][2][2];941 vk_pipeline pipeline_add_rms_norepeat[2][2][2];942 943 // indexed by num_additional_fused_ops == num_adds - 1944 vk_pipeline pipeline_multi_add[MAX_FUSED_ADDS];945 vk_pipeline pipeline_multi_add_rms[MAX_FUSED_ADDS];946 947 vk_pipeline pipeline_add_id_f32;948 949 vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64;950 vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32;951 vk_pipeline pipeline_scale_f32;952 vk_pipeline pipeline_log[2];953 vk_pipeline pipeline_tri[2];954 vk_pipeline pipeline_diag[2];955 vk_pipeline pipeline_clamp[2];956 vk_pipeline pipeline_pad_f32;957 vk_pipeline pipeline_roll_f32;958 vk_pipeline pipeline_repeat_i32, pipeline_repeat_back_f32;959 vk_pipeline pipeline_repeat_i16;960 vk_pipeline pipeline_cpy_f32_f32, pipeline_cpy_f32_f16, pipeline_cpy_f16_f16, pipeline_cpy_f16_f32, pipeline_cpy_f32_bf16, pipeline_cpy_bf16_f32, pipeline_cpy_f32_i32, pipeline_cpy_i32_f32;961 vk_pipeline pipeline_contig_cpy_f32_f32, pipeline_contig_cpy_f32_f16, pipeline_contig_cpy_f16_f16, pipeline_contig_cpy_f16_f32, pipeline_contig_cpy_f32_bf16, pipeline_contig_cpy_bf16_f32, pipeline_contig_cpy_f32_i32, pipeline_contig_cpy_i32_f32;962 vk_pipeline pipeline_cpy_f32_quant[GGML_TYPE_COUNT];963 vk_pipeline pipeline_cpy_quant_f32[GGML_TYPE_COUNT];964 vk_pipeline pipeline_cpy_transpose_16, pipeline_cpy_transpose_32;965 // [src0 0=fp32,1=fp16][dst]966 vk_pipeline pipeline_set_rows_i32[2][GGML_TYPE_COUNT];967 vk_pipeline pipeline_set_rows_i64[2][GGML_TYPE_COUNT];968 vk_pipeline pipeline_norm_f32;969 vk_pipeline pipeline_group_norm_f32;970 vk_pipeline pipeline_rms_norm_f32;971 vk_pipeline pipeline_rms_norm_mul_f32;972 vk_pipeline pipeline_rms_norm_partials_f32;973 vk_pipeline pipeline_rms_norm_mul_partials_f32;974 vk_pipeline pipeline_rms_norm_mul_rope_f32_f32;975 vk_pipeline pipeline_rms_norm_mul_rope_f32_f16;976 vk_pipeline pipeline_rms_norm_back_f32;977 vk_pipeline pipeline_l2_norm_f32;978 979 // [src/dst 0=fp32,1=fp16]980 vk_pipeline pipeline_exp[2];981 vk_pipeline pipeline_expm1[2];982 vk_pipeline pipeline_elu[2];983 vk_pipeline pipeline_gelu[2];984 vk_pipeline pipeline_gelu_erf[2];985 vk_pipeline pipeline_gelu_quick[2];986 vk_pipeline pipeline_silu[2];987 vk_pipeline pipeline_relu[2];988 vk_pipeline pipeline_sqr[2];989 vk_pipeline pipeline_sqrt[2];990 vk_pipeline pipeline_sin[2];991 vk_pipeline pipeline_cos[2];992 vk_pipeline pipeline_xielu[2];993 vk_pipeline pipeline_neg[2];994 vk_pipeline pipeline_tanh[2];995 vk_pipeline pipeline_sigmoid[2];996 vk_pipeline pipeline_hardsigmoid[2];997 vk_pipeline pipeline_hardswish[2];998 vk_pipeline pipeline_abs[2];999 vk_pipeline pipeline_softplus[2];1000 vk_pipeline pipeline_step[2];1001 vk_pipeline pipeline_round[2];1002 vk_pipeline pipeline_ceil[2];1003 vk_pipeline pipeline_floor[2];1004 vk_pipeline pipeline_trunc[2];1005 vk_pipeline pipeline_sgn[2];1006 1007 vk_pipeline pipeline_add1_f16_f16;1008 vk_pipeline pipeline_add1_f16_f32;1009 vk_pipeline pipeline_add1_f32_f32;1010 1011 vk_pipeline pipeline_arange_f32;1012 1013 vk_pipeline pipeline_fill_f32;1014 vk_pipeline pipeline_fill_f16;1015 1016 vk_pipeline pipeline_geglu[2];1017 vk_pipeline pipeline_reglu[2];1018 vk_pipeline pipeline_swiglu[2];1019 vk_pipeline pipeline_swiglu_oai[2];1020 vk_pipeline pipeline_geglu_erf[2];1021 vk_pipeline pipeline_geglu_quick[2];1022 1023 vk_pipeline pipeline_leaky_relu[2];1024 vk_pipeline pipeline_silu_back_f32;1025 vk_pipeline pipeline_diag_mask_inf_f32;1026 vk_pipeline pipeline_soft_max_f32, pipeline_soft_max_f32_f16;1027 vk_pipeline pipeline_soft_max_f32_wg512, pipeline_soft_max_f32_f16_wg512;1028 vk_pipeline pipeline_soft_max_back_f32;1029 1030 vk_pipeline pipeline_soft_max_large1_f32, pipeline_soft_max_large1_f32_f16;1031 vk_pipeline pipeline_soft_max_large2_f32, pipeline_soft_max_large2_f32_f16;1032 vk_pipeline pipeline_soft_max_large3_f32, pipeline_soft_max_large3_f32_f16;1033 1034 vk_pipeline pipeline_rope_norm_f32, pipeline_rope_norm_f16, pipeline_rope_norm_f32_f16;1035 vk_pipeline pipeline_rope_neox_f32, pipeline_rope_neox_f16, pipeline_rope_neox_f32_f16;1036 vk_pipeline pipeline_rope_multi_f32, pipeline_rope_multi_f16, pipeline_rope_multi_f32_f16;1037 vk_pipeline pipeline_rope_vision_f32, pipeline_rope_vision_f16;1038 vk_pipeline pipeline_argsort_f32[num_argsort_pipelines];1039 vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines];1040 vk_pipeline pipeline_topk_f32[num_topk_pipelines];1041 vk_pipeline pipeline_sum_rows_f32;1042 vk_pipeline pipeline_fwht_f32[4];1043 vk_pipeline pipeline_cumsum_f32;1044 vk_pipeline pipeline_cumsum_small_f32;1045 vk_pipeline pipeline_cumsum_multipass1_f32;1046 vk_pipeline pipeline_cumsum_multipass2_f32;1047 vk_pipeline pipeline_argmax_f32;1048 vk_pipeline pipeline_count_equal_i32;1049 std::map<vk_solve_tri_pipeline_state, vk_pipeline> pipeline_solve_tri_f32;1050 vk_pipeline pipeline_im2col_f32, pipeline_im2col_f32_f16;1051 vk_pipeline pipeline_im2col_3d_f32, pipeline_im2col_3d_f32_f16;1052 vk_pipeline pipeline_timestep_embedding_f32;1053 vk_pipeline pipeline_conv_transpose_1d_f32;1054 vk_pipeline pipeline_col2im_1d_f32;1055 vk_pipeline pipeline_col2im_1d_f16;1056 vk_pipeline pipeline_col2im_1d_bf16;1057 vk_pipeline pipeline_out_prod_f32;1058 vk_pipeline pipeline_snake_f32;1059 vk_pipeline pipeline_snake_f16;1060 vk_pipeline pipeline_snake_bf16;1061 vk_pipeline pipeline_pool1d_f32;1062 vk_pipeline pipeline_pool2d_f32;1063 vk_pipeline pipeline_rwkv_wkv6_f32;1064 vk_pipeline pipeline_rwkv_wkv7_f32;1065 vk_pipeline pipeline_gated_linear_attn_f32;1066 // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d1281067 vk_pipeline pipeline_gated_delta_net[4][2];1068 vk_pipeline pipeline_ssm_scan_f32_d128;1069 vk_pipeline pipeline_ssm_scan_f32_d256;1070 vk_pipeline pipeline_ssm_conv_f32;1071 vk_pipeline pipeline_ssm_conv_silu_f32;1072 vk_pipeline pipeline_ssm_conv_bias_silu_f32;1073 vk_pipeline pipeline_opt_step_adamw_f32;1074 vk_pipeline pipeline_opt_step_sgd_f32;1075 std::map<vk_conv2d_pipeline_state, vk_pipeline> pipeline_conv2d_f32[CONV_SHAPE_COUNT];1076 std::map<vk_conv2d_pipeline_state, vk_pipeline> pipeline_conv2d_f16_f32[CONV_SHAPE_COUNT];1077 std::map<vk_conv2d_pipeline_state, vk_pipeline> pipeline_conv_transpose_2d_f32[CONV_SHAPE_COUNT];1078 std::map<vk_conv2d_pipeline_state, vk_pipeline> pipeline_conv_transpose_2d_f16_f32[CONV_SHAPE_COUNT];1079 std::map<vk_conv3d_pipeline_state, vk_pipeline> pipeline_conv3d_f32[CONV_SHAPE_COUNT];1080 std::map<vk_conv3d_pipeline_state, vk_pipeline> pipeline_conv3d_f16_f32[CONV_SHAPE_COUNT];1081 vk_pipeline pipeline_conv2d_dw_whcn_f32, pipeline_conv2d_dw_whcn_f16_f32;1082 vk_pipeline pipeline_conv2d_dw_cwhn_f32, pipeline_conv2d_dw_cwhn_f16_f32;1083 1084 std::map<vk_fa_pipeline_state, vk_pipeline> pipeline_flash_attn_f32_f16;1085 1086 std::map<std::pair<uint32_t, uint32_t>, vk_pipeline> pipeline_fa_mask_opt;1087 1088 vk_pipeline pipeline_flash_attn_split_k_reduce;1089 vk_pipeline pipeline_count_experts;1090 1091 // [2] is for whether to take n_experts from spec constant (0) or push constant (1)1092 vk_pipeline pipeline_topk_moe[num_topk_moe_pipelines][2];1093 1094 std::vector<vk_pipeline_ref> all_pipelines;1095 1096 std::vector<std::tuple<void*, size_t, vk_buffer>> pinned_memory;1097 1098 vk::Fence fence;1099 vk_buffer sync_staging;1100 1101 ggml_backend_buffer_type buffer_type;1102 1103 bool disable_fusion;1104 bool disable_host_visible_vidmem;1105 bool allow_sysmem_fallback;1106 bool disable_graph_optimize;1107 1108 std::unique_ptr<vk_memory_logger> memory_logger;1109 1110 ~vk_device_struct() {1111 VK_LOG_DEBUG("destroy device " << name);1112 1113 device.destroyFence(fence);1114 1115 ggml_vk_destroy_buffer(sync_staging);1116 1117 if (compute_queue) compute_queue->cmd_pool.destroy(device);1118 if (transfer_queue) transfer_queue->cmd_pool.destroy(device);1119 1120 // Explicitly clear to ensure queues drop their shared_ptrs to handles1121 // before the Vulkan logical device instance is destroyed1122 compute_queue.reset();1123 transfer_queue.reset();1124 1125 for (auto& pipeline : all_pipelines) {1126 if (pipeline.expired()) {1127 continue;1128 }1129 1130 vk_pipeline pl = pipeline.lock();1131 ggml_vk_destroy_pipeline(device, pl);1132 }1133 all_pipelines.clear();1134 1135 device.destroyDescriptorSetLayout(dsl);1136 1137 device.destroy();1138 }1139};1140 1141void vk_command_pool::init(vk_device& device, vk_queue *q_) {1142 cmd_buffers.clear();1143 q = q_;1144 1145 vk::CommandPoolCreateInfo command_pool_create_info(1146 vk::CommandPoolCreateFlags(VK_COMMAND_POOL_CREATE_TRANSIENT_BIT | VK_COMMAND_POOL_CREATE_RESET_COMMAND_BUFFER_BIT),1147 q->queue_family_index);1148 pool = device->device.createCommandPool(command_pool_create_info);1149}1150 1151void vk_command_pool::destroy(vk::Device& device) {1152 device.destroyCommandPool(pool);1153 pool = nullptr;1154 cmd_buffers.clear();1155}1156 1157static void ggml_vk_print_device_fault_info(const vk_device& device) {1158 if (!device->device_fault || !device->pfn_vkGetDeviceFaultInfoEXT) {1159 return;1160 }1161 1162 VkDeviceFaultCountsEXT fault_counts {};1163 fault_counts.sType = VK_STRUCTURE_TYPE_DEVICE_FAULT_COUNTS_EXT;1164 VkResult res = device->pfn_vkGetDeviceFaultInfoEXT(device->device, &fault_counts, nullptr);1165 if (res != VK_SUCCESS) {1166 GGML_LOG_ERROR("ggml_vulkan: vkGetDeviceFaultInfoEXT (counts) failed: %d\n", res);1167 return;1168 }1169 1170 std::vector<VkDeviceFaultAddressInfoEXT> address_infos(fault_counts.addressInfoCount);1171 std::vector<VkDeviceFaultVendorInfoEXT> vendor_infos(fault_counts.vendorInfoCount);1172 1173 VkDeviceFaultInfoEXT fault_info {};1174 fault_info.sType = VK_STRUCTURE_TYPE_DEVICE_FAULT_INFO_EXT;1175 fault_info.pAddressInfos = address_infos.data();1176 fault_info.pVendorInfos = vendor_infos.data();1177 1178 res = device->pfn_vkGetDeviceFaultInfoEXT(device->device, &fault_counts, &fault_info);1179 if (res != VK_SUCCESS) {1180 GGML_LOG_ERROR("ggml_vulkan: vkGetDeviceFaultInfoEXT (info) failed: %d\n", res);1181 return;1182 }1183 1184 if (fault_counts.addressInfoCount == 0 && fault_counts.vendorInfoCount == 0 && fault_info.description[0] == '\0') {1185 return;1186 }1187 1188 if (fault_info.description[0] != '\0') {1189 GGML_LOG_ERROR("ggml_vulkan: device fault on %s: %s\n", device->name.c_str(), fault_info.description);1190 }1191 1192 for (uint32_t i = 0; i < fault_counts.addressInfoCount; i++) {1193 const auto& info = address_infos[i];1194 GGML_LOG_CONT(" address fault %u: type=%d address=0x%llx precision=0x%llx\n",1195 i, (int)info.addressType,1196 (unsigned long long)info.reportedAddress,1197 (unsigned long long)info.addressPrecision);1198 }1199 for (uint32_t i = 0; i < fault_counts.vendorInfoCount; i++) {1200 const auto& info = vendor_infos[i];