use crate::HwdnaError;
use ash::{vk, Device, Entry, Instance};
use std::ffi::CString;
const DOT_F64_GLSL: &str = r"
#version 460
layout(local_size_x = 256) in;
layout(std430, binding = 0) buffer A { double a[]; };
layout(std430, binding = 1) buffer B { double b[]; };
layout(std430, binding = 2) buffer Result { double result[]; };
layout(std430, binding = 3) buffer Dim { uint n; };
void main() {
uint id = gl_GlobalInvocationID.x;
uint total = gl_NumWorkGroups.x * 256u;
double sum = 0.0;
for (uint i = id; i < n; i += total) {
sum += a[i] * b[i];
}
result[id] = sum;
}
";
const MATMUL_F64_GLSL: &str = r"
#version 460
layout(local_size_x = 16, local_size_y = 16) in;
layout(std430, binding = 0) buffer A { double a[]; };
layout(std430, binding = 1) buffer B { double b[]; };
layout(std430, binding = 2) buffer C { double c[]; };
layout(std430, binding = 3) buffer Dim { uint n; };
void main() {
uint row = gl_GlobalInvocationID.y;
uint col = gl_GlobalInvocationID.x;
uint nv = n;
if (row >= nv || col >= nv) return;
double sum = 0.0;
for (uint k = 0; k < nv; ++k) {
sum += a[row * nv + k] * b[k * nv + col];
}
c[row * nv + col] = sum;
}
";
const REDUCE_SUM_F64_GLSL: &str = r"
#version 460
layout(local_size_x = 256) in;
layout(std430, binding = 0) buffer Data { double data[]; };
layout(std430, binding = 1) buffer Result { double result[]; };
layout(std430, binding = 2) buffer Dim { uint n; };
void main() {
uint id = gl_GlobalInvocationID.x;
uint total = gl_NumWorkGroups.x * 256u;
double sum = 0.0;
for (uint i = id; i < n; i += total) {
sum += data[i];
}
result[id] = sum;
}
";
const DOT_F32_GLSL: &str = r"
#version 460
layout(local_size_x = 256) in;
layout(std430, binding = 0) buffer A { float a[]; };
layout(std430, binding = 1) buffer B { float b[]; };
layout(std430, binding = 2) buffer Result { float result[]; };
layout(std430, binding = 3) buffer Dim { uint n; };
void main() {
uint id = gl_GlobalInvocationID.x;
uint total = gl_NumWorkGroups.x * 256u;
float sum = 0.0;
for (uint i = id; i < n; i += total) {
sum += a[i] * b[i];
}
result[id] = sum;
}
";
const MATMUL_F32_GLSL: &str = r"
#version 460
layout(local_size_x = 16, local_size_y = 16) in;
layout(std430, binding = 0) buffer A { float a[]; };
layout(std430, binding = 1) buffer B { float b[]; };
layout(std430, binding = 2) buffer C { float c[]; };
layout(std430, binding = 3) buffer Dim { uint n; };
void main() {
uint row = gl_GlobalInvocationID.y;
uint col = gl_GlobalInvocationID.x;
uint nv = n;
if (row >= nv || col >= nv) return;
float sum = 0.0;
for (uint k = 0; k < nv; ++k) {
sum += a[row * nv + k] * b[k * nv + col];
}
c[row * nv + col] = sum;
}
";
const REDUCE_SUM_F32_GLSL: &str = r"
#version 460
layout(local_size_x = 256) in;
layout(std430, binding = 0) buffer Data { float data[]; };
layout(std430, binding = 1) buffer Result { float result[]; };
layout(std430, binding = 2) buffer Dim { uint n; };
void main() {
uint id = gl_GlobalInvocationID.x;
uint total = gl_NumWorkGroups.x * 256u;
float sum = 0.0;
for (uint i = id; i < n; i += total) {
sum += data[i];
}
result[id] = sum;
}
";
struct ComputeShader {
pipeline: vk::Pipeline,
layout: vk::PipelineLayout,
shader_module: vk::ShaderModule,
}
pub struct VulkanRuntime {
initialized: bool,
device_name: String,
entry: Option<Entry>,
instance: Option<Instance>,
physical_device: Option<vk::PhysicalDevice>,
device: Option<Device>,
queue: Option<vk::Queue>,
command_pool: Option<vk::CommandPool>,
descriptor_set_layout: Option<vk::DescriptorSetLayout>,
descriptor_pool: Option<vk::DescriptorPool>,
dot_f64_shader: Option<ComputeShader>,
matmul_f64_shader: Option<ComputeShader>,
reduce_f64_shader: Option<ComputeShader>,
dot_f32_shader: Option<ComputeShader>,
matmul_f32_shader: Option<ComputeShader>,
reduce_f32_shader: Option<ComputeShader>,
}
unsafe impl Send for VulkanRuntime {}
unsafe impl Sync for VulkanRuntime {}
impl VulkanRuntime {
pub fn new() -> Self {
Self {
initialized: false,
device_name: String::new(),
entry: None,
instance: None,
physical_device: None,
device: None,
queue: None,
command_pool: None,
descriptor_set_layout: None,
descriptor_pool: None,
dot_f64_shader: None,
matmul_f64_shader: None,
reduce_f64_shader: None,
dot_f32_shader: None,
matmul_f32_shader: None,
reduce_f32_shader: None,
}
}
pub fn is_available() -> bool {
detect_vulkan()
}
pub fn init(&mut self) -> Result<(), HwdnaError> {
if self.initialized {
return Ok(());
}
if !Self::is_available() {
return Err(HwdnaError::NoSupportedKernels);
}
self.init_vulkan()?;
self.compile_shaders()?;
self.initialized = true;
Ok(())
}
fn init_vulkan(&mut self) -> Result<(), HwdnaError> {
let entry = unsafe { Entry::load() }.map_err(|_| HwdnaError::NoSupportedKernels)?;
let app_name = CString::new("Himada").unwrap();
let engine_name = CString::new("Himada").unwrap();
let app_info = vk::ApplicationInfo::default()
.application_name(&app_name)
.application_version(vk::make_api_version(0, 0, 1, 0))
.engine_name(&engine_name)
.engine_version(vk::make_api_version(0, 0, 1, 0))
.api_version(vk::make_api_version(0, 1, 3, 0));
let instance_create_info = vk::InstanceCreateInfo::default().application_info(&app_info);
let instance = unsafe { entry.create_instance(&instance_create_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let physical_devices = unsafe { instance.enumerate_physical_devices() }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let (phys_device, props) = physical_devices
.iter()
.filter_map(|&pd| {
let props = unsafe { instance.get_physical_device_properties(pd) };
let family_props =
unsafe { instance.get_physical_device_queue_family_properties(pd) };
let has_compute = family_props.iter().any(|qf| {
qf.queue_flags.contains(vk::QueueFlags::COMPUTE)
});
has_compute.then_some((pd, props))
})
.next()
.ok_or(HwdnaError::NoSupportedKernels)?;
self.device_name = {
let name = props.device_name;
let len = name.iter().position(|&c| c == 0).unwrap_or(name.len());
let bytes: &[u8] = unsafe { std::slice::from_raw_parts(name.as_ptr() as *const u8, len) };
String::from_utf8_lossy(bytes).to_string()
};
let family_props = unsafe { instance.get_physical_device_queue_family_properties(phys_device) };
let compute_queue_family = family_props
.iter()
.position(|qf| qf.queue_flags.contains(vk::QueueFlags::COMPUTE))
.ok_or(HwdnaError::NoSupportedKernels)? as u32;
let queue_priority = [1.0f32];
let device_queue_info = vk::DeviceQueueCreateInfo::default()
.queue_family_index(compute_queue_family)
.queue_priorities(&queue_priority);
let queue_create_infos = [device_queue_info];
let device_create_info = vk::DeviceCreateInfo::default()
.queue_create_infos(&queue_create_infos);
let device = unsafe { instance.create_device(phys_device, &device_create_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let queue = unsafe { device.get_device_queue(compute_queue_family, 0) };
let pool_info = vk::CommandPoolCreateInfo::default()
.queue_family_index(compute_queue_family)
.flags(vk::CommandPoolCreateFlags::RESET_COMMAND_BUFFER);
let command_pool = unsafe { device.create_command_pool(&pool_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let bindings = [
vk::DescriptorSetLayoutBinding::default()
.binding(0)
.descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
.descriptor_count(1)
.stage_flags(vk::ShaderStageFlags::COMPUTE),
vk::DescriptorSetLayoutBinding::default()
.binding(1)
.descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
.descriptor_count(1)
.stage_flags(vk::ShaderStageFlags::COMPUTE),
vk::DescriptorSetLayoutBinding::default()
.binding(2)
.descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
.descriptor_count(1)
.stage_flags(vk::ShaderStageFlags::COMPUTE),
vk::DescriptorSetLayoutBinding::default()
.binding(3)
.descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
.descriptor_count(1)
.stage_flags(vk::ShaderStageFlags::COMPUTE),
];
let dsl_info = vk::DescriptorSetLayoutCreateInfo::default().bindings(&bindings);
let descriptor_set_layout =
unsafe { device.create_descriptor_set_layout(&dsl_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let pool_sizes = [vk::DescriptorPoolSize {
ty: vk::DescriptorType::STORAGE_BUFFER,
descriptor_count: 16,
}];
let pool_create_info = vk::DescriptorPoolCreateInfo::default()
.pool_sizes(&pool_sizes)
.max_sets(4)
.flags(vk::DescriptorPoolCreateFlags::FREE_DESCRIPTOR_SET);
let descriptor_pool = unsafe { device.create_descriptor_pool(&pool_create_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
self.entry = Some(entry);
self.instance = Some(instance);
self.physical_device = Some(phys_device);
self.device = Some(device);
self.queue = Some(queue);
self.command_pool = Some(command_pool);
self.descriptor_set_layout = Some(descriptor_set_layout);
self.descriptor_pool = Some(descriptor_pool);
Ok(())
}
fn compile_glsl_to_spirv(&self, glsl: &str) -> Result<Vec<u32>, HwdnaError> {
use naga::front::glsl::{Frontend, Options};
let mut frontend = Frontend::default();
let options = Options::from(naga::ShaderStage::Compute);
let module = frontend.parse(&options, glsl)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let mut validator = naga::valid::Validator::new(
naga::valid::ValidationFlags::all(),
naga::valid::Capabilities::all(),
);
let module_info = validator
.validate(&module)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let mut writer = naga::back::spv::Writer::new(&naga::back::spv::Options::default())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let mut words = Vec::new();
writer.write(&module, &module_info, None, &None, &mut words)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
Ok(words)
}
fn create_compute_shader(&mut self, glsl: &str) -> Result<ComputeShader, HwdnaError> {
let device = self.device.as_ref().ok_or(HwdnaError::NoSupportedKernels)?;
let spirv = self.compile_glsl_to_spirv(glsl)?;
let shader_info = vk::ShaderModuleCreateInfo::default().code(&spirv);
let shader_module = unsafe { device.create_shader_module(&shader_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let entry_name = CString::new("main").unwrap();
let dsl = self.descriptor_set_layout.ok_or(HwdnaError::NoSupportedKernels)?;
let layouts = [dsl];
let layout_info = vk::PipelineLayoutCreateInfo::default().set_layouts(&layouts);
let layout = unsafe { device.create_pipeline_layout(&layout_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let stage = vk::PipelineShaderStageCreateInfo::default()
.stage(vk::ShaderStageFlags::COMPUTE)
.module(shader_module)
.name(&entry_name);
let compute_info = vk::ComputePipelineCreateInfo::default()
.stage(stage)
.layout(layout);
let pipelines = unsafe {
device.create_compute_pipelines(
vk::PipelineCache::null(),
&[compute_info],
None,
)
}
.map_err(|_| HwdnaError::NoSupportedKernels)?;
Ok(ComputeShader { pipeline: pipelines[0], layout, shader_module })
}
fn compile_shaders(&mut self) -> Result<(), HwdnaError> {
self.dot_f64_shader = Some(self.create_compute_shader(DOT_F64_GLSL)?);
self.matmul_f64_shader = Some(self.create_compute_shader(MATMUL_F64_GLSL)?);
self.reduce_f64_shader = Some(self.create_compute_shader(REDUCE_SUM_F64_GLSL)?);
self.dot_f32_shader = Some(self.create_compute_shader(DOT_F32_GLSL)?);
self.matmul_f32_shader = Some(self.create_compute_shader(MATMUL_F32_GLSL)?);
self.reduce_f32_shader = Some(self.create_compute_shader(REDUCE_SUM_F32_GLSL)?);
Ok(())
}
fn create_buffer(
&self,
size: vk::DeviceSize,
usage: vk::BufferUsageFlags,
data: Option<&[u8]>,
) -> Result<(vk::Buffer, vk::DeviceMemory), HwdnaError> {
let device = self.device.as_ref().ok_or(HwdnaError::NoSupportedKernels)?;
let phys_device = self.physical_device.ok_or(HwdnaError::NoSupportedKernels)?;
let instance = self.instance.as_ref().ok_or(HwdnaError::NoSupportedKernels)?;
let buffer_info = vk::BufferCreateInfo::default()
.size(size)
.usage(usage)
.sharing_mode(vk::SharingMode::EXCLUSIVE);
let buffer = unsafe { device.create_buffer(&buffer_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let mem_reqs = unsafe { device.get_buffer_memory_requirements(buffer) };
let mem_props = unsafe { instance.get_physical_device_memory_properties(phys_device) };
let mem_type_index = mem_props
.memory_types
.iter()
.enumerate()
.find(|(i, mt)| {
(mem_reqs.memory_type_bits & (1 << i)) != 0
&& mt.property_flags.contains(vk::MemoryPropertyFlags::HOST_VISIBLE)
&& mt.property_flags.contains(vk::MemoryPropertyFlags::HOST_COHERENT)
})
.map(|(i, _)| i)
.ok_or(HwdnaError::NoSupportedKernels)? as u32;
let alloc_info = vk::MemoryAllocateInfo::default()
.allocation_size(mem_reqs.size)
.memory_type_index(mem_type_index);
let memory = unsafe { device.allocate_memory(&alloc_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
unsafe { device.bind_buffer_memory(buffer, memory, 0) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
if let Some(bytes) = data {
unsafe {
let ptr = device
.map_memory(memory, 0, size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
std::ptr::copy_nonoverlapping(bytes.as_ptr(), ptr as *mut u8, bytes.len());
device.unmap_memory(memory);
}
}
Ok((buffer, memory))
}
fn destroy_buffer(&self, buffer: vk::Buffer, memory: vk::DeviceMemory) {
if let Some(ref device) = self.device {
unsafe {
device.destroy_buffer(buffer, None);
device.free_memory(memory, None);
}
}
}
fn allocate_descriptor_set(
&self,
buffers: &[vk::Buffer],
) -> Result<vk::DescriptorSet, HwdnaError> {
let device = self.device.as_ref().ok_or(HwdnaError::NoSupportedKernels)?;
let pool = self.descriptor_pool.ok_or(HwdnaError::NoSupportedKernels)?;
let layout = self.descriptor_set_layout.ok_or(HwdnaError::NoSupportedKernels)?;
let layouts = [layout];
let alloc_info = vk::DescriptorSetAllocateInfo::default()
.descriptor_pool(pool)
.set_layouts(&layouts);
let descriptor_sets = unsafe { device.allocate_descriptor_sets(&alloc_info) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let descriptor_set = descriptor_sets[0];
let buf_infos: Vec<vk::DescriptorBufferInfo> = buffers
.iter()
.map(|&buf| {
vk::DescriptorBufferInfo::default()
.buffer(buf)
.offset(0)
.range(vk::WHOLE_SIZE)
})
.collect();
let writes: Vec<vk::WriteDescriptorSet> = buf_infos
.iter()
.enumerate()
.map(|(i, buf_info)| {
vk::WriteDescriptorSet::default()
.dst_set(descriptor_set)
.dst_binding(i as u32)
.descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
.buffer_info(std::slice::from_ref(buf_info))
})
.collect();
unsafe { device.update_descriptor_sets(&writes, &[]) };
Ok(descriptor_set)
}
fn dispatch_one_shot(
&self,
shader: &ComputeShader,
descriptor_set: vk::DescriptorSet,
group_count_x: u32,
group_count_y: u32,
group_count_z: u32,
) -> Result<(), HwdnaError> {
let device = self.device.as_ref().ok_or(HwdnaError::NoSupportedKernels)?;
let queue = self.queue.ok_or(HwdnaError::NoSupportedKernels)?;
let cmd_pool = self.command_pool.ok_or(HwdnaError::NoSupportedKernels)?;
let alloc_info = vk::CommandBufferAllocateInfo::default()
.command_pool(cmd_pool)
.level(vk::CommandBufferLevel::PRIMARY)
.command_buffer_count(1);
let cmd_buffers = unsafe { device.allocate_command_buffers(&alloc_info) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let cmd_buffer = cmd_buffers[0];
let begin_info = vk::CommandBufferBeginInfo::default()
.flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT);
unsafe {
device
.begin_command_buffer(cmd_buffer, &begin_info)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
device.cmd_bind_pipeline(cmd_buffer, vk::PipelineBindPoint::COMPUTE, shader.pipeline);
device.cmd_bind_descriptor_sets(
cmd_buffer,
vk::PipelineBindPoint::COMPUTE,
shader.layout,
0,
&[descriptor_set],
&[],
);
device.cmd_dispatch(cmd_buffer, group_count_x, group_count_y, group_count_z);
device
.end_command_buffer(cmd_buffer)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
}
let fence_info = vk::FenceCreateInfo::default();
let fence = unsafe { device.create_fence(&fence_info, None) }
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let cmd_buffers = [cmd_buffer];
let submit_info = vk::SubmitInfo::default().command_buffers(&cmd_buffers);
unsafe {
device
.queue_submit(queue, &[submit_info], fence)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
device
.wait_for_fences(&[fence], true, u64::MAX)
.map_err(|_| HwdnaError::NoSupportedKernels)?;
device.destroy_fence(fence, None);
device.free_command_buffers(cmd_pool, &[cmd_buffer]);
}
Ok(())
}
pub fn device_name(&self) -> &str {
&self.device_name
}
pub fn is_initialized(&self) -> bool {
self.initialized
}
pub fn dot_f64(&self, a: &[f64], b: &[f64]) -> Result<f64, HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
let device = match self.device.as_ref() {
Some(d) => d,
None => return Ok(a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()),
};
let shader = match self.dot_f64_shader.as_ref() {
Some(s) => s,
None => return Ok(a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()),
};
let n = a.len().min(b.len()) as u64;
if n == 0 {
return Ok(0.0);
}
let num_groups = ((n + 255) / 256) as u32;
let result_count = (num_groups as u64) * 256;
let buf_size = n * size_of::<f64>() as u64;
let result_size = result_count * size_of::<f64>() as u64;
let (a_buf, a_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(a)),
)?;
let (b_buf, b_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(b)),
)?;
let (result_buf, result_mem) =
self.create_buffer(result_size, vk::BufferUsageFlags::STORAGE_BUFFER, None)?;
let n_val: u32 = n as u32;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set =
self.allocate_descriptor_set(&[a_buf, b_buf, result_buf, n_buf])?;
self.dispatch_one_shot(shader, descriptor_set, num_groups, 1, 1)?;
let sum = unsafe {
let ptr = device
.map_memory(result_mem, 0, result_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f64, result_count as usize);
let total: f64 = slice.iter().sum();
device.unmap_memory(result_mem);
total
};
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(a_buf, a_mem);
self.destroy_buffer(b_buf, b_mem);
self.destroy_buffer(result_buf, result_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(sum)
}
pub fn matmul_f64(
&self,
a: &[f64],
b: &[f64],
c: &mut [f64],
n: usize,
) -> Result<(), HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
if n == 0 {
return Ok(());
}
if self.device.is_none() || self.matmul_f64_shader.is_none() {
for i in 0..n {
for j in 0..n {
let mut sum = 0.0;
for k in 0..n {
sum += a[i * n + k] * b[k * n + j];
}
c[i * n + j] = sum;
}
}
return Ok(());
}
let device = self.device.as_ref().unwrap();
let shader = self.matmul_f64_shader.as_ref().unwrap();
let mat_size = (n * n) as u64 * size_of::<f64>() as u64;
let n_val = n as u32;
let (a_buf, a_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(a)),
)?;
let (b_buf, b_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(b)),
)?;
let (c_buf, c_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
None,
)?;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set =
self.allocate_descriptor_set(&[a_buf, b_buf, c_buf, n_buf])?;
let gx = ((n as u32) + 15) / 16;
let gy = ((n as u32) + 15) / 16;
self.dispatch_one_shot(shader, descriptor_set, gx, gy, 1)?;
unsafe {
let ptr = device
.map_memory(c_mem, 0, mat_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f64, n * n);
c.copy_from_slice(slice);
device.unmap_memory(c_mem);
}
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(a_buf, a_mem);
self.destroy_buffer(b_buf, b_mem);
self.destroy_buffer(c_buf, c_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(())
}
pub fn reduce_sum_f64(&self, data: &[f64]) -> Result<f64, HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
if data.is_empty() {
return Ok(0.0);
}
let device = match self.device.as_ref() {
Some(d) => d,
None => return Ok(data.iter().sum()),
};
let shader = match self.reduce_f64_shader.as_ref() {
Some(s) => s,
None => return Ok(data.iter().sum()),
};
let n = data.len() as u64;
let num_groups = ((n + 255) / 256) as u32;
let result_count = (num_groups as u64) * 256;
let buf_size = n * size_of::<f64>() as u64;
let result_size = result_count * size_of::<f64>() as u64;
let (data_buf, data_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(data)),
)?;
let (result_buf, result_mem) =
self.create_buffer(result_size, vk::BufferUsageFlags::STORAGE_BUFFER, None)?;
let n_val: u32 = n as u32;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set = self.allocate_descriptor_set(&[data_buf, result_buf, n_buf])?;
self.dispatch_one_shot(shader, descriptor_set, num_groups, 1, 1)?;
let sum = unsafe {
let ptr = device
.map_memory(result_mem, 0, result_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f64, result_count as usize);
let total: f64 = slice.iter().sum();
device.unmap_memory(result_mem);
total
};
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(data_buf, data_mem);
self.destroy_buffer(result_buf, result_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(sum)
}
pub fn dot_f32(&self, a: &[f32], b: &[f32]) -> Result<f32, HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
let device = match self.device.as_ref() {
Some(d) => d,
None => return Ok(a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()),
};
let shader = match self.dot_f32_shader.as_ref() {
Some(s) => s,
None => return Ok(a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()),
};
let n = a.len().min(b.len()) as u64;
if n == 0 {
return Ok(0.0);
}
let num_groups = ((n + 255) / 256) as u32;
let result_count = (num_groups as u64) * 256;
let buf_size = n * size_of::<f32>() as u64;
let result_size = result_count * size_of::<f32>() as u64;
let (a_buf, a_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(a)),
)?;
let (b_buf, b_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(b)),
)?;
let (result_buf, result_mem) =
self.create_buffer(result_size, vk::BufferUsageFlags::STORAGE_BUFFER, None)?;
let n_val: u32 = n as u32;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set =
self.allocate_descriptor_set(&[a_buf, b_buf, result_buf, n_buf])?;
self.dispatch_one_shot(shader, descriptor_set, num_groups, 1, 1)?;
let sum = unsafe {
let ptr = device
.map_memory(result_mem, 0, result_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f32, result_count as usize);
let total: f32 = slice.iter().sum();
device.unmap_memory(result_mem);
total
};
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(a_buf, a_mem);
self.destroy_buffer(b_buf, b_mem);
self.destroy_buffer(result_buf, result_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(sum)
}
pub fn matmul_f32(
&self,
a: &[f32],
b: &[f32],
c: &mut [f32],
n: usize,
) -> Result<(), HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
if n == 0 {
return Ok(());
}
if self.device.is_none() || self.matmul_f32_shader.is_none() {
for i in 0..n {
for j in 0..n {
let mut sum = 0.0f32;
for k in 0..n {
sum += a[i * n + k] * b[k * n + j];
}
c[i * n + j] = sum;
}
}
return Ok(());
}
let device = self.device.as_ref().unwrap();
let shader = self.matmul_f32_shader.as_ref().unwrap();
let mat_size = (n * n) as u64 * size_of::<f32>() as u64;
let n_val = n as u32;
let (a_buf, a_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(a)),
)?;
let (b_buf, b_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(b)),
)?;
let (c_buf, c_mem) = self.create_buffer(
mat_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
None,
)?;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set =
self.allocate_descriptor_set(&[a_buf, b_buf, c_buf, n_buf])?;
let gx = ((n as u32) + 15) / 16;
let gy = ((n as u32) + 15) / 16;
self.dispatch_one_shot(shader, descriptor_set, gx, gy, 1)?;
unsafe {
let ptr = device
.map_memory(c_mem, 0, mat_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f32, n * n);
c.copy_from_slice(slice);
device.unmap_memory(c_mem);
}
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(a_buf, a_mem);
self.destroy_buffer(b_buf, b_mem);
self.destroy_buffer(c_buf, c_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(())
}
pub fn reduce_sum_f32(&self, data: &[f32]) -> Result<f32, HwdnaError> {
if !self.initialized {
return Err(HwdnaError::NoSupportedKernels);
}
if data.is_empty() {
return Ok(0.0);
}
let device = match self.device.as_ref() {
Some(d) => d,
None => return Ok(data.iter().sum()),
};
let shader = match self.reduce_f32_shader.as_ref() {
Some(s) => s,
None => return Ok(data.iter().sum()),
};
let n = data.len() as u64;
let num_groups = ((n + 255) / 256) as u32;
let result_count = (num_groups as u64) * 256;
let buf_size = n * size_of::<f32>() as u64;
let result_size = result_count * size_of::<f32>() as u64;
let (data_buf, data_mem) = self.create_buffer(
buf_size,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::cast_slice(data)),
)?;
let (result_buf, result_mem) =
self.create_buffer(result_size, vk::BufferUsageFlags::STORAGE_BUFFER, None)?;
let n_val: u32 = n as u32;
let (n_buf, n_mem) = self.create_buffer(
4,
vk::BufferUsageFlags::STORAGE_BUFFER,
Some(bytemuck::bytes_of(&n_val)),
)?;
let descriptor_set = self.allocate_descriptor_set(&[data_buf, result_buf, n_buf])?;
self.dispatch_one_shot(shader, descriptor_set, num_groups, 1, 1)?;
let sum = unsafe {
let ptr = device
.map_memory(result_mem, 0, result_size, vk::MemoryMapFlags::empty())
.map_err(|_| HwdnaError::NoSupportedKernels)?;
let slice = std::slice::from_raw_parts(ptr as *const f32, result_count as usize);
let total: f32 = slice.iter().sum();
device.unmap_memory(result_mem);
total
};
unsafe {
let _ = device.free_descriptor_sets(
self.descriptor_pool.unwrap(),
&[descriptor_set],
);
}
self.destroy_buffer(data_buf, data_mem);
self.destroy_buffer(result_buf, result_mem);
self.destroy_buffer(n_buf, n_mem);
Ok(sum)
}
}
impl Drop for VulkanRuntime {
fn drop(&mut self) {
if let Some(ref device) = self.device {
unsafe {
let destroy_shader = |s: &ComputeShader| {
device.destroy_shader_module(s.shader_module, None);
device.destroy_pipeline_layout(s.layout, None);
device.destroy_pipeline(s.pipeline, None);
};
if let Some(ref s) = self.dot_f64_shader { destroy_shader(s); }
if let Some(ref s) = self.matmul_f64_shader { destroy_shader(s); }
if let Some(ref s) = self.reduce_f64_shader { destroy_shader(s); }
if let Some(ref s) = self.dot_f32_shader { destroy_shader(s); }
if let Some(ref s) = self.matmul_f32_shader { destroy_shader(s); }
if let Some(ref s) = self.reduce_f32_shader { destroy_shader(s); }
if let Some(pool) = self.descriptor_pool {
device.destroy_descriptor_pool(pool, None);
}
if let Some(layout) = self.descriptor_set_layout {
device.destroy_descriptor_set_layout(layout, None);
}
if let Some(pool) = self.command_pool {
device.destroy_command_pool(pool, None);
}
device.destroy_device(None);
}
}
if let Some(ref instance) = self.instance {
unsafe { instance.destroy_instance(None); }
}
}
}
fn detect_vulkan() -> bool {
#[cfg(target_os = "linux")]
{
unsafe {
let lib = libc::dlopen(
b"libvulkan.so.1\0".as_ptr() as *const i8,
libc::RTLD_LAZY,
);
if !lib.is_null() {
libc::dlclose(lib);
return true;
}
}
false
}
#[cfg(target_os = "windows")]
{
true
}
#[cfg(target_os = "macos")]
{
unsafe {
let lib = libc::dlopen(
b"libMoltenVK.dylib\0".as_ptr() as *const i8,
libc::RTLD_LAZY,
);
if !lib.is_null() {
libc::dlclose(lib);
return true;
}
}
false
}
#[cfg(not(any(target_os = "linux", target_os = "windows", target_os = "macos")))]
{
false
}
}