diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000..e11a51e --- /dev/null +++ b/.gitmodules @@ -0,0 +1,3 @@ +[submodule "clblast-rs/vcpkg"] + path = clblast-rs/vcpkg + url = https://github.com/microsoft/vcpkg diff --git a/clblast-rs/.gitignore b/clblast-rs/.gitignore new file mode 100644 index 0000000..9f97022 --- /dev/null +++ b/clblast-rs/.gitignore @@ -0,0 +1 @@ +target/ \ No newline at end of file diff --git a/clblast-rs/Cargo.lock b/clblast-rs/Cargo.lock new file mode 100644 index 0000000..9561aa1 --- /dev/null +++ b/clblast-rs/Cargo.lock @@ -0,0 +1,171 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "cl3" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ad5f4170f520224684c5c6e86e8d5474c2a10696e553b1db8389970cd82ece7" +dependencies = [ + "dlopen2", + "libc", + "opencl-sys", + "thiserror", +] + +[[package]] +name = "clblast-rs" +version = "0.1.0" +dependencies = [ + "opencl3", + "vcpkg", +] + +[[package]] +name = "dlopen2" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e2c5bd4158e66d1e215c49b837e11d62f3267b30c92f1d171c4d3105e3dc4d4" +dependencies = [ + "dlopen2_derive", + "libc", + "once_cell", + "winapi", +] + +[[package]] +name = "dlopen2_derive" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fbbb781877580993a8707ec48672673ec7b81eeba04cfd2310bd28c08e47c8f" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "libc" +version = "0.2.189" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "opencl-sys" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9bd005f352b2f05acd01d04122448e84f8bc66bdae49045927841bc63de62123" +dependencies = [ + "libc", +] + +[[package]] +name = "opencl3" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecb3473c4f01afd0eea3ee5b31649fec47dbc23353a63e1b916aa67cd810fb73" +dependencies = [ + "cl3", + "libc", +] + +[[package]] +name = "proc-macro2" +version = "1.0.107" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "syn" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "thiserror" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec86235f5fcc2a73650310756d2ac5b138a5780bbbdfae3eeccec992c435ba4f" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc04cd3e1236dd4a98afca4569f2deb3f120e5422a4023be2cb683f8486292af" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] +name = "winapi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" +dependencies = [ + "winapi-i686-pc-windows-gnu", + "winapi-x86_64-pc-windows-gnu", +] + +[[package]] +name = "winapi-i686-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" + +[[package]] +name = "winapi-x86_64-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" diff --git a/clblast-rs/Cargo.toml b/clblast-rs/Cargo.toml new file mode 100644 index 0000000..1a05994 --- /dev/null +++ b/clblast-rs/Cargo.toml @@ -0,0 +1,10 @@ +[package] +name = "clblast-rs" +version = "0.1.0" +edition = "2024" + +[build-dependencies] +vcpkg = "0.2.15" + +[dependencies] +opencl3 = "0.12.3" diff --git a/clblast-rs/build.rs b/clblast-rs/build.rs new file mode 100644 index 0000000..f2026dd --- /dev/null +++ b/clblast-rs/build.rs @@ -0,0 +1,25 @@ +fn main() { + let mut config = vcpkg::Config::new(); + + if std::env::var("VCPKG_ROOT").is_err() { + println!("cargo:warning=VCPKG_ROOT is not set. Please set it to the path of your vcpkg installation."); + } + + unsafe { + std::env::set_var("VCPKGRS_DYNAMIC", "1"); // DLL 読み込み許可 + } + + let target = std::env::var("TARGET").unwrap_or_default(); + let triplet = match target.as_str() { + "x86_64-pc-windows-msvc" => "x64-windows", + "i686-pc-windows-msvc" => "x86-windows", + "aarch64-pc-windows-msvc" => "arm64-windows", + _ => "x64-windows", + }; + + config.target_triplet(triplet); + + if let Err(e) = config.probe("clblast") { // CLBlast 探索 + panic!("Failed to find clblast with vcpkg. Error: {}", e); + } +} \ No newline at end of file diff --git a/clblast-rs/docs/new.md b/clblast-rs/docs/new.md new file mode 100644 index 0000000..e69de29 diff --git a/clblast-rs/src/clblast.rs b/clblast-rs/src/clblast.rs new file mode 100644 index 0000000..9672a45 --- /dev/null +++ b/clblast-rs/src/clblast.rs @@ -0,0 +1,218 @@ +use opencl3::memory::{Buffer, ClMem}; +use crate::clblast_float::CLBlastFloat; + +pub type CLBlastStatusCode = i32; +pub type CLBlastLayout = i32; +pub type CLBlastTranspose = i32; + +pub type CLCommandQueue = *mut std::ffi::c_void; +pub type CLEvent = *mut std::ffi::c_void; +pub type CLMem = *mut std::ffi::c_void; + +const DEFAULT_PLATFORM_ID: usize = 0; +const DEFAULT_DEVICE_ID: usize = 0; + +pub const A_BUFFER: usize = 0; +pub const B_BUFFER: usize = 1; +pub const C_BUFFER: usize = 2; + +pub enum Layout { + RowMajor = 101, + ColMajor = 102, +} +pub enum Transpose { + No = 111, + Yes = 112, +} + +unsafe extern "C" { + #[link_name = "CLBlastSgemm"] + pub unsafe fn clblast_sgemm( + layout: CLBlastLayout, a_transpose: CLBlastTranspose, + b_transpose: CLBlastTranspose, m: usize, n: usize, + k: usize, alpha: f32, a_buffer: CLMem, + a_offset: usize, a_ld: usize, b_buffer: CLMem, + b_offset: usize, b_ld: usize, beta: f32, c_buffer: CLMem, + c_offset: usize, c_ld: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode; + #[link_name = "CLBlastDgemm"] + pub unsafe fn clblast_dgemm( + layout: CLBlastLayout, a_transpose: CLBlastTranspose, + b_transpose: CLBlastTranspose, m: usize, n: usize, + k: usize, alpha: f64, a_buffer: CLMem, + a_offset: usize, a_ld: usize, b_buffer: CLMem, + b_offset: usize, b_ld: usize, beta: f64, c_buffer: CLMem, + c_offset: usize, c_ld: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode; + + #[link_name = "CLBlastSaxpy"] + pub unsafe fn clblast_saxpy( + n: usize, alpha: f32, x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode; + #[link_name = "CLBlastDaxpy"] + pub unsafe fn clblast_daxpy( + n: usize, alpha: f64, x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode; + + #[link_name = "CLBlastSdot"] + pub unsafe fn clblast_sdot( + n: usize, dot_buffer: CLMem, dot_offset: usize, + x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, + queue: *mut CLCommandQueue, event: *mut CLEvent + ) -> CLBlastStatusCode; + #[link_name = "CLBlastDdot"] + pub unsafe fn clblast_ddot( + n: usize, dot_buffer: CLMem, dot_offset: usize, + x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, + queue: *mut CLCommandQueue, event: *mut CLEvent + ) -> CLBlastStatusCode; +} + +struct CLBlastMatrix { + buffer: Buffer, + row_major: bool, + rows: usize, + cols: usize, +} +pub struct CLBlast { + queue: opencl3::command_queue::CommandQueue, + context: opencl3::context::Context, + buffers: [Option>; 3], + ty: std::marker::PhantomData, +} +impl CLBlast { + pub fn new(platform_id: usize, device_id: usize) -> Result, String> { + let platforms = opencl3::platform::get_platforms()?; + if platforms.len() <= platform_id { return Err(String::from("Platform ID is out of range")); } + let platform = platforms[platform_id]; + + let devices = platform.get_devices(opencl3::device::CL_DEVICE_TYPE_GPU)?; + if devices.len() <= device_id { return Err(String::from("Device ID is out of range")); } + let device = opencl3::device::Device::new(devices[device_id]); + let context = opencl3::context::Context::from_device(&device)?; + let queue = unsafe{ opencl3::command_queue::CommandQueue::create(&context, device.id(), opencl3::command_queue::CL_QUEUE_PROFILING_ENABLE)? }; + + Ok( + Self { + queue: queue, + context: context, + buffers: [None, None, None], + ty: std::marker::PhantomData + } + ) + } + pub fn set_buffer(&mut self, target: usize, mem_flag: opencl3::memory::cl_mem_flags, vec: &[T], row_major: bool, rows: usize, cols: usize) -> Result<(), String>{ + if target >= self.buffers.len() { + return Err(String::from("Target index is out of range")); + } + let mut buffer = unsafe { + opencl3::memory::Buffer::::create( + &self.context, mem_flag, vec.len(), std::ptr::null_mut() + )? + }; + unsafe{ + self.queue.enqueue_write_buffer(&mut buffer, opencl3::types::CL_TRUE, 0, &vec[..], &[])?; + } + + self.buffers[target] = Some(CLBlastMatrix { buffer: buffer, row_major: row_major, rows: rows, cols: cols }); + Ok(()) + } + pub fn swap_buffer(&mut self, target1: usize, target2: usize) -> Result<(), String> { + if target1 >= self.buffers.len() || target2 >= self.buffers.len() { + return Err(String::from("Target1 index is out of range")); + } + if target1 == target2 { + return Ok(()); + } + + self.buffers.swap(target1, target2); + Ok(()) + } + pub fn read_buffer(&mut self, target: usize, vec: &mut [T]) -> Result<(), String> { + let buf = self.buffers[target].as_ref().ok_or_else(|| String::from("Buffer not set"))?; + + unsafe{ + self.queue.enqueue_read_buffer(&buf.buffer, opencl3::types::CL_TRUE, 0, vec, &[])?; + } + Ok(()) + } + pub fn mat_mul(&mut self) -> Result<(), String> + where T: CLBlastFloat + { + let a = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer A not set"))?; + let b = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer B not set"))?; + let c = self.buffers[C_BUFFER].as_ref().ok_or_else(|| String::from("Buffer C not set"))?; + + let layout = if a.row_major { Layout::RowMajor } else { Layout::ColMajor }; + + let (a_ld, b_ld, c_ld) = if a.row_major { + (a.cols, b.cols, c.cols) + } else { + (a.rows, b.rows, c.rows) + }; + let mut raw_queue = self.queue.get(); + + let status = T::gemm( + layout as CLBlastLayout, + Transpose::No as CLBlastTranspose, + Transpose::No as CLBlastTranspose, + a.rows, // m + b.cols, // n + a.cols, // k + a.buffer.get(), 0, a_ld, + b.buffer.get(), 0, b_ld, + c.buffer.get(), 0, c_ld, + &mut raw_queue, + std::ptr::null_mut(), + ); + + if status == 0 { + Ok(()) + } else { + Err(format!("CLBlast GEMM failed with status code: {}", status)) + } + } + pub fn dot_mul(&mut self) -> Result<(), String> + where T: CLBlastFloat + { + let a = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer X not set"))?; + let b = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer Y not set"))?; + let c = self.buffers[C_BUFFER].as_ref().ok_or_else(|| String::from("Buffer DOT not set"))?; + + let n_a = a.rows * a.cols; + let n_b = b.rows * b.cols; + + if n_a != n_b { + return Err(format!("Vector dimension mismatch: X size is {}, but Y size is {}", n_a, n_b)); + } + if c.rows * c.cols < 1 { + return Err(String::from("Buffer DOT must contain at least 1 element")); + } + + let mut raw_queue = self.queue.get(); + + let status = T::dot( + n_a, + c.buffer.get(), 0, + a.buffer.get(), 0, 1, + b.buffer.get(), 0, 1, + &mut raw_queue, + std::ptr::null_mut(), + ); + + if status == 0 { + Ok(()) + } else { + Err(format!("CLBlast SDOT failed with status code: {}", status)) + } + } +} + diff --git a/clblast-rs/src/clblast_float.rs b/clblast-rs/src/clblast_float.rs new file mode 100644 index 0000000..deffee4 --- /dev/null +++ b/clblast-rs/src/clblast_float.rs @@ -0,0 +1,89 @@ +use crate::clblast::{CLBlastLayout, CLBlastTranspose, CLCommandQueue, CLEvent, CLMem, CLBlastStatusCode}; + +pub trait CLBlastFloat: Sized { + fn gemm( + layout: CLBlastLayout, a_transpose: CLBlastTranspose, + b_transpose: CLBlastTranspose, m: usize, n: usize, + k: usize, a_buffer: CLMem, + a_offset: usize, a_ld: usize, b_buffer: CLMem, + b_offset: usize, b_ld: usize, c_buffer: CLMem, + c_offset: usize, c_ld: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode; + fn dot( + n: usize, dot_buffer: CLMem, dot_offset: usize, + x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, + queue: *mut CLCommandQueue, event: *mut CLEvent + ) -> CLBlastStatusCode; +} +impl CLBlastFloat for f32 { + fn gemm( + layout: CLBlastLayout, a_transpose: CLBlastTranspose, + b_transpose: CLBlastTranspose, m: usize, n: usize, + k: usize, a_buffer: CLMem, + a_offset: usize, a_ld: usize, b_buffer: CLMem, + b_offset: usize, b_ld: usize, c_buffer: CLMem, + c_offset: usize, c_ld: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode { + unsafe { + crate::clblast::clblast_sgemm( + layout, a_transpose, b_transpose, m, n, + k, 1.0 as f32, a_buffer, a_offset, a_ld, b_buffer, + b_offset, b_ld, 0.0 as f32, c_buffer, c_offset, c_ld, + queue, event + ) + } + } + fn dot( + n: usize, dot_buffer: CLMem, dot_offset: usize, + x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, + queue: *mut CLCommandQueue, event: *mut CLEvent + ) -> CLBlastStatusCode { + unsafe { + crate::clblast::clblast_sdot( + n, dot_buffer, dot_offset, + x_buffer, x_offset, x_inc, + y_buffer, y_offset, y_inc, + queue, event + ) + } + } +} +impl CLBlastFloat for f64 { + fn gemm( + layout: CLBlastLayout, a_transpose: CLBlastTranspose, + b_transpose: CLBlastTranspose, m: usize, n: usize, + k: usize, a_buffer: CLMem, + a_offset: usize, a_ld: usize, b_buffer: CLMem, + b_offset: usize, b_ld: usize, c_buffer: CLMem, + c_offset: usize, c_ld: usize, queue: *mut CLCommandQueue, + event: *mut CLEvent + ) -> CLBlastStatusCode { + unsafe { + crate::clblast::clblast_dgemm( + layout, a_transpose, b_transpose, m, n, + k, 1.0 as f64, a_buffer, a_offset, a_ld, b_buffer, + b_offset, b_ld, 0.0 as f64, c_buffer, c_offset, c_ld, + queue, event + ) + } + } + fn dot( + n: usize, dot_buffer: CLMem, dot_offset: usize, + x_buffer: CLMem, x_offset: usize, x_inc: usize, + y_buffer: CLMem, y_offset: usize, y_inc: usize, + queue: *mut CLCommandQueue, event: *mut CLEvent + ) -> CLBlastStatusCode { + unsafe { + crate::clblast::clblast_ddot( + n, dot_buffer, dot_offset, + x_buffer, x_offset, x_inc, + y_buffer, y_offset, y_inc, + queue, event + ) + } + } +} \ No newline at end of file diff --git a/clblast-rs/src/lib.rs b/clblast-rs/src/lib.rs new file mode 100644 index 0000000..6ead0cb --- /dev/null +++ b/clblast-rs/src/lib.rs @@ -0,0 +1,2 @@ +pub mod clblast_float; +pub mod clblast; \ No newline at end of file diff --git a/clblast-rs/vcpkg b/clblast-rs/vcpkg new file mode 160000 index 0000000..40f3c70 --- /dev/null +++ b/clblast-rs/vcpkg @@ -0,0 +1 @@ +Subproject commit 40f3c709db80acf154ac4b17a1f83c564ebd022e