From 50ba3bc8611c0b9ba916cb147b87425a5e386624 Mon Sep 17 00:00:00 2001 From: OtonariShunji Date: Thu, 13 Aug 2026 20:45:00 +0900 Subject: [PATCH 1/3] =?UTF-8?q?clblast=E3=81=AEwrapper=E3=82=92=E4=BD=9C?= =?UTF-8?q?=E6=88=90=E3=81=97=E3=80=81f32=E3=81=AE=E8=A1=8C=E5=88=97?= =?UTF-8?q?=E7=A9=8D=E3=82=92=E8=A1=8C=E3=81=84=E3=81=BE=E3=81=99=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitmodules | 3 + clblast-rs/.gitignore | 1 + clblast-rs/Cargo.lock | 171 ++++++++++++++++++++++++++++++++++ clblast-rs/Cargo.toml | 10 ++ clblast-rs/build.rs | 16 ++++ clblast-rs/src/clblast.rs | 187 ++++++++++++++++++++++++++++++++++++++ clblast-rs/src/lib.rs | 1 + clblast-rs/vcpkg | 1 + 8 files changed, 390 insertions(+) create mode 100644 .gitmodules create mode 100644 clblast-rs/.gitignore create mode 100644 clblast-rs/Cargo.lock create mode 100644 clblast-rs/Cargo.toml create mode 100644 clblast-rs/build.rs create mode 100644 clblast-rs/src/clblast.rs create mode 100644 clblast-rs/src/lib.rs create mode 160000 clblast-rs/vcpkg 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..de63c1a --- /dev/null +++ b/clblast-rs/build.rs @@ -0,0 +1,16 @@ +fn main() { + let mut config = vcpkg::Config::new(); + + unsafe { + if std::env::var("VCPKG_ROOT").is_err() { + panic!("VCPKG_ROOT is not set."); + } + + std::env::set_var("VCPKGRS_DYNAMIC", "1"); // DLL 読み込み許可 + }; + config.target_triplet("x64-windows"); + + 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/src/clblast.rs b/clblast-rs/src/clblast.rs new file mode 100644 index 0000000..2076d3f --- /dev/null +++ b/clblast-rs/src/clblast.rs @@ -0,0 +1,187 @@ +use opencl3::memory::{Buffer, ClMem}; + +type CLBlastStatusCode = i32; +type CLBlastLayout = u32; +type CLBlastTranspose = u32; + +type CLCommandQueue = *mut std::ffi::c_void; +type CLEvent = *mut std::ffi::c_void; +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; + +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, 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, opencl3::memory::CL_MEM_WRITE_ONLY, 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 enum Layout { + RowMajor = 101, + ColMajor = 102, +} +pub enum Transpose { + No = 111, + Yes = 112, +} + +impl CLBlast { + pub fn mat_mul(&mut self) -> Result<(), String> { + 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 = unsafe { + clblast_sgemm( + layout as CLBlastLayout, + Transpose::No as CLBlastTranspose, + Transpose::No as CLBlastTranspose, + a.rows, // m + b.cols, // n + a.cols, // k + 1.0, + a.buffer.get(), 0, a_ld, + b.buffer.get(), 0, b_ld, + 0.0, + 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)) + } + } +} \ 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..a014903 --- /dev/null +++ b/clblast-rs/src/lib.rs @@ -0,0 +1 @@ +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 From f7f75e97e714d90e88c5824f40ef049ade7767de Mon Sep 17 00:00:00 2001 From: OtonariShunji Date: Thu, 13 Aug 2026 21:38:27 +0900 Subject: [PATCH 2/3] =?UTF-8?q?build.rs=E3=81=AE=E4=BF=AE=E6=AD=A3?= =?UTF-8?q?=E3=81=8A=E3=82=88=E3=81=B3=E3=80=81=E5=86=85=E7=A9=8D=E3=82=92?= =?UTF-8?q?=E5=AE=9F=E8=A3=85=E3=81=97=E3=81=BE=E3=81=99=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- clblast-rs/build.rs | 21 ++++++--- clblast-rs/docs/new.md | 0 clblast-rs/src/clblast.rs | 97 +++++++++++++++++++++++++++++++++++++-- 3 files changed, 108 insertions(+), 10 deletions(-) create mode 100644 clblast-rs/docs/new.md diff --git a/clblast-rs/build.rs b/clblast-rs/build.rs index de63c1a..f2026dd 100644 --- a/clblast-rs/build.rs +++ b/clblast-rs/build.rs @@ -1,15 +1,24 @@ fn main() { let mut config = vcpkg::Config::new(); - unsafe { - if std::env::var("VCPKG_ROOT").is_err() { - panic!("VCPKG_ROOT is not set."); - } + 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("x64-windows"); - + + config.target_triplet(triplet); + if let Err(e) = config.probe("clblast") { // CLBlast 探索 panic!("Failed to find clblast with vcpkg. Error: {}", e); } 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 index 2076d3f..0a3dbe6 100644 --- a/clblast-rs/src/clblast.rs +++ b/clblast-rs/src/clblast.rs @@ -1,8 +1,8 @@ use opencl3::memory::{Buffer, ClMem}; type CLBlastStatusCode = i32; -type CLBlastLayout = u32; -type CLBlastTranspose = u32; +type CLBlastLayout = i32; +type CLBlastTranspose = i32; type CLCommandQueue = *mut std::ffi::c_void; type CLEvent = *mut std::ffi::c_void; @@ -99,13 +99,13 @@ impl CLBlast { } ) } - pub fn set_buffer(&mut self, target: usize, vec: &[T], row_major: bool, rows: usize, cols: usize) -> Result<(), String>{ + 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, opencl3::memory::CL_MEM_WRITE_ONLY, vec.len(), std::ptr::null_mut() + &self.context, mem_flag, vec.len(), std::ptr::null_mut() )? }; unsafe{ @@ -184,4 +184,93 @@ impl CLBlast { Err(format!("CLBlast GEMM failed with status code: {}", status)) } } + pub fn dot_mul(&mut self) -> Result<(), String> { + let x = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer X not set"))?; + let y = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer Y not set"))?; + + let n = x.rows * x.cols; // Assuming x is a vector + let mut raw_queue = self.queue.get(); + + let status = unsafe { + clblast_sdot( + n, + x.buffer.get(), 0, + x.buffer.get(), 0, 1, + y.buffer.get(), 0, 1, + &mut raw_queue, + std::ptr::null_mut(), + ) + }; + + if status == 0 { + Ok(()) + } else { + Err(format!("CLBlast DOT failed with status code: {}", status)) + } + } +} +impl CLBlast { + pub fn mat_mul(&mut self) -> Result<(), String> { + 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 = unsafe { + clblast_dgemm( + layout as CLBlastLayout, + Transpose::No as CLBlastTranspose, + Transpose::No as CLBlastTranspose, + a.rows, // m + b.cols, // n + a.cols, // k + 1.0, + a.buffer.get(), 0, a_ld, + b.buffer.get(), 0, b_ld, + 0.0, + 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> { + let x = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer X not set"))?; + let y = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer Y not set"))?; + + let n = x.rows * x.cols; // Assuming x is a vector + let mut raw_queue = self.queue.get(); + + let status = unsafe { + clblast_ddot( + n, + x.buffer.get(), 0, + x.buffer.get(), 0, 1, + y.buffer.get(), 0, 1, + &mut raw_queue, + std::ptr::null_mut(), + ) + }; + + if status == 0 { + Ok(()) + } else { + Err(format!("CLBlast DOT failed with status code: {}", status)) + } + } } \ No newline at end of file From fbde31d2e1b89036c2eae5f19d82f498ba081a8b Mon Sep 17 00:00:00 2001 From: OtonariShunji Date: Thu, 13 Aug 2026 23:31:55 +0900 Subject: [PATCH 3/3] =?UTF-8?q?trait=E3=82=92=E5=AE=9A=E7=BE=A9=E3=81=97f3?= =?UTF-8?q?2,f64=E3=81=AB=E5=AE=9F=E8=A3=85=E3=81=97=E3=81=BE=E3=81=99?= =?UTF-8?q?=E3=80=82=E3=81=93=E3=82=8C=E3=81=AB=E3=82=88=E3=82=8Af32,f64?= =?UTF-8?q?=E3=81=AE=E9=96=A2=E6=95=B0=E3=81=AE=E5=91=BC=E3=81=B3=E5=A4=89?= =?UTF-8?q?=E3=81=88=E3=82=92=E9=9A=A0=E5=8C=BF=E5=8C=96=E3=81=97=E3=81=BE?= =?UTF-8?q?=E3=81=99=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- clblast-rs/src/clblast.rs | 146 ++++++++++---------------------- clblast-rs/src/clblast_float.rs | 89 +++++++++++++++++++ clblast-rs/src/lib.rs | 1 + 3 files changed, 134 insertions(+), 102 deletions(-) create mode 100644 clblast-rs/src/clblast_float.rs diff --git a/clblast-rs/src/clblast.rs b/clblast-rs/src/clblast.rs index 0a3dbe6..9672a45 100644 --- a/clblast-rs/src/clblast.rs +++ b/clblast-rs/src/clblast.rs @@ -1,12 +1,13 @@ use opencl3::memory::{Buffer, ClMem}; +use crate::clblast_float::CLBlastFloat; -type CLBlastStatusCode = i32; -type CLBlastLayout = i32; -type CLBlastTranspose = i32; +pub type CLBlastStatusCode = i32; +pub type CLBlastLayout = i32; +pub type CLBlastTranspose = i32; -type CLCommandQueue = *mut std::ffi::c_void; -type CLEvent = *mut std::ffi::c_void; -type CLMem = *mut std::ffi::c_void; +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; @@ -15,6 +16,15 @@ 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( @@ -134,19 +144,9 @@ impl CLBlast { } Ok(()) } -} - -pub enum Layout { - RowMajor = 101, - ColMajor = 102, -} -pub enum Transpose { - No = 111, - Yes = 112, -} - -impl CLBlast { - pub fn mat_mul(&mut self) -> Result<(), String> { + 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"))?; @@ -160,23 +160,19 @@ impl CLBlast { }; let mut raw_queue = self.queue.get(); - let status = unsafe { - clblast_sgemm( + let status = T::gemm( layout as CLBlastLayout, Transpose::No as CLBlastTranspose, Transpose::No as CLBlastTranspose, a.rows, // m b.cols, // n a.cols, // k - 1.0, a.buffer.get(), 0, a_ld, b.buffer.get(), 0, b_ld, - 0.0, c.buffer.get(), 0, c_ld, &mut raw_queue, std::ptr::null_mut(), - ) - }; + ); if status == 0 { Ok(()) @@ -184,93 +180,39 @@ impl CLBlast { Err(format!("CLBlast GEMM failed with status code: {}", status)) } } - pub fn dot_mul(&mut self) -> Result<(), String> { - let x = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer X not set"))?; - let y = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer Y not set"))?; - - let n = x.rows * x.cols; // Assuming x is a vector - let mut raw_queue = self.queue.get(); - - let status = unsafe { - clblast_sdot( - n, - x.buffer.get(), 0, - x.buffer.get(), 0, 1, - y.buffer.get(), 0, 1, - &mut raw_queue, - std::ptr::null_mut(), - ) - }; - - if status == 0 { - Ok(()) - } else { - Err(format!("CLBlast DOT 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)); } - } -} -impl CLBlast { - pub fn mat_mul(&mut self) -> Result<(), String> { - 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 = unsafe { - clblast_dgemm( - layout as CLBlastLayout, - Transpose::No as CLBlastTranspose, - Transpose::No as CLBlastTranspose, - a.rows, // m - b.cols, // n - a.cols, // k - 1.0, - a.buffer.get(), 0, a_ld, - b.buffer.get(), 0, b_ld, - 0.0, - 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)) + if c.rows * c.cols < 1 { + return Err(String::from("Buffer DOT must contain at least 1 element")); } - } - pub fn dot_mul(&mut self) -> Result<(), String> { - let x = self.buffers[A_BUFFER].as_ref().ok_or_else(|| String::from("Buffer X not set"))?; - let y = self.buffers[B_BUFFER].as_ref().ok_or_else(|| String::from("Buffer Y not set"))?; - - let n = x.rows * x.cols; // Assuming x is a vector let mut raw_queue = self.queue.get(); - let status = unsafe { - clblast_ddot( - n, - x.buffer.get(), 0, - x.buffer.get(), 0, 1, - y.buffer.get(), 0, 1, + 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 DOT failed with status code: {}", status)) + Err(format!("CLBlast SDOT failed with status code: {}", status)) } } -} \ No newline at end of file +} + 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 index a014903..6ead0cb 100644 --- a/clblast-rs/src/lib.rs +++ b/clblast-rs/src/lib.rs @@ -1 +1,2 @@ +pub mod clblast_float; pub mod clblast; \ No newline at end of file