1use std::{ffi::c_void, sync::Arc};
4
5use libloading::Library;
6
7use crate::{RocmLibrarySet, RocmLoadError};
8
9const HIP_MEMCPY_HOST_TO_DEVICE: i32 = 1;
10const HIP_MEMCPY_DEVICE_TO_HOST: i32 = 2;
11const ROCBLAS_OPERATION_NONE: i32 = 111;
12
13type HipMalloc = unsafe extern "C" fn(*mut *mut c_void, usize) -> i32;
14type HipFree = unsafe extern "C" fn(*mut c_void) -> i32;
15type HipMemcpy = unsafe extern "C" fn(*mut c_void, *const c_void, usize, i32) -> i32;
16type HipDeviceSynchronize = unsafe extern "C" fn() -> i32;
17type RocblasCreate = unsafe extern "C" fn(*mut *mut c_void) -> i32;
18type RocblasDestroy = unsafe extern "C" fn(*mut c_void) -> i32;
19type RocblasSgemm = unsafe extern "C" fn(
20 *mut c_void,
21 i32,
22 i32,
23 i32,
24 i32,
25 i32,
26 *const f32,
27 *const f32,
28 i32,
29 *const f32,
30 i32,
31 *const f32,
32 *mut f32,
33 i32,
34) -> i32;
35
36pub(crate) struct RocmDeviceBuffer {
37 runtime: Arc<RocmLibrarySet>,
38 address: usize,
39 len: usize,
40}
41
42unsafe impl Send for RocmDeviceBuffer {}
45unsafe impl Sync for RocmDeviceBuffer {}
47
48impl RocmDeviceBuffer {
49 pub(crate) fn len(&self) -> usize {
50 self.len
51 }
52
53 pub(crate) fn runtime(&self) -> &Arc<RocmLibrarySet> {
54 &self.runtime
55 }
56
57 fn pointer(&self) -> *mut c_void {
58 self.address as *mut c_void
59 }
60
61 pub(crate) fn read(&self) -> Result<Vec<f32>, RocmLoadError> {
62 self.runtime.download(self)
63 }
64}
65
66impl Drop for RocmDeviceBuffer {
67 fn drop(&mut self) {
68 let _ = self.runtime.free(self.pointer());
69 }
70}
71
72impl RocmLibrarySet {
73 pub(crate) fn upload(
74 self: &Arc<Self>,
75 values: &[f32],
76 ) -> Result<Arc<RocmDeviceBuffer>, RocmLoadError> {
77 let bytes = byte_count(values.len())?;
78 let (hip, _) = self.execution_handles();
79 let malloc = symbol::<HipMalloc>(hip, b"hipMalloc\0")?;
80 let copy = symbol::<HipMemcpy>(hip, b"hipMemcpy\0")?;
81 let mut pointer = std::ptr::null_mut();
82 check_hip(unsafe { malloc(&mut pointer, bytes) }, "hipMalloc")?;
84 let status = unsafe {
87 copy(
88 pointer,
89 values.as_ptr().cast::<c_void>(),
90 bytes,
91 HIP_MEMCPY_HOST_TO_DEVICE,
92 )
93 };
94 if let Err(error) = check_hip(status, "hipMemcpy host-to-device") {
95 let _ = self.free(pointer);
96 return Err(error);
97 }
98 Ok(Arc::new(RocmDeviceBuffer {
99 runtime: Arc::clone(self),
100 address: pointer as usize,
101 len: values.len(),
102 }))
103 }
104
105 pub(crate) fn matmul(
106 self: &Arc<Self>,
107 left: &RocmDeviceBuffer,
108 right: &RocmDeviceBuffer,
109 rows: usize,
110 inner: usize,
111 cols: usize,
112 ) -> Result<Arc<RocmDeviceBuffer>, RocmLoadError> {
113 validate_matrix_lengths(left.len(), right.len(), rows, inner, cols)?;
114 if !Arc::ptr_eq(self, left.runtime()) || !Arc::ptr_eq(self, right.runtime()) {
115 return Err(error("ROCm inputs belong to another runtime"));
116 }
117 let output_len = rows
118 .checked_mul(cols)
119 .ok_or_else(|| error("ROCm output length overflowed"))?;
120 let output = self.upload(&vec![0.0_f32; output_len])?;
121 let (_, rocblas) = self.execution_handles();
122 let create = symbol::<RocblasCreate>(rocblas, b"rocblas_create_handle\0")?;
123 let destroy = symbol::<RocblasDestroy>(rocblas, b"rocblas_destroy_handle\0")?;
124 let sgemm = symbol::<RocblasSgemm>(rocblas, b"rocblas_sgemm\0")?;
125 let mut handle = std::ptr::null_mut();
126 check_rocblas(unsafe { create(&mut handle) }, "rocblas_create_handle")?;
128 let dimensions = matrix_dimensions(rows, inner, cols)?;
129 let alpha = 1.0_f32;
130 let beta = 0.0_f32;
131 let status = unsafe {
134 sgemm(
135 handle,
136 ROCBLAS_OPERATION_NONE,
137 ROCBLAS_OPERATION_NONE,
138 dimensions.cols,
139 dimensions.rows,
140 dimensions.inner,
141 &alpha,
142 right.pointer().cast::<f32>(),
143 dimensions.cols,
144 left.pointer().cast::<f32>(),
145 dimensions.inner,
146 &beta,
147 output.pointer().cast::<f32>(),
148 dimensions.cols,
149 )
150 };
151 let gemm_result = check_rocblas(status, "rocblas_sgemm");
152 let destroy_result = check_rocblas(unsafe { destroy(handle) }, "rocblas_destroy_handle");
154 gemm_result?;
155 destroy_result?;
156 Ok(output)
157 }
158
159 pub fn matmul_f32(
162 self: &Arc<Self>,
163 left: &[f32],
164 right: &[f32],
165 rows: usize,
166 inner: usize,
167 cols: usize,
168 ) -> Result<Vec<f32>, RocmLoadError> {
169 let left = self.upload(left)?;
170 let right = self.upload(right)?;
171 self.matmul(&left, &right, rows, inner, cols)?.read()
172 }
173
174 fn download(&self, buffer: &RocmDeviceBuffer) -> Result<Vec<f32>, RocmLoadError> {
175 let bytes = byte_count(buffer.len())?;
176 let (hip, _) = self.execution_handles();
177 let copy = symbol::<HipMemcpy>(hip, b"hipMemcpy\0")?;
178 let synchronize = symbol::<HipDeviceSynchronize>(hip, b"hipDeviceSynchronize\0")?;
179 let mut values = vec![0.0_f32; buffer.len()];
180 check_hip(
182 unsafe {
183 copy(
184 values.as_mut_ptr().cast::<c_void>(),
185 buffer.pointer(),
186 bytes,
187 HIP_MEMCPY_DEVICE_TO_HOST,
188 )
189 },
190 "hipMemcpy device-to-host",
191 )?;
192 check_hip(unsafe { synchronize() }, "hipDeviceSynchronize")?;
194 Ok(values)
195 }
196
197 fn free(&self, pointer: *mut c_void) -> Result<(), RocmLoadError> {
198 if pointer.is_null() {
199 return Ok(());
200 }
201 let (hip, _) = self.execution_handles();
202 let free = symbol::<HipFree>(hip, b"hipFree\0")?;
203 check_hip(unsafe { free(pointer) }, "hipFree")
205 }
206}
207
208struct MatrixDimensions {
209 rows: i32,
210 inner: i32,
211 cols: i32,
212}
213
214fn matrix_dimensions(
215 rows: usize,
216 inner: usize,
217 cols: usize,
218) -> Result<MatrixDimensions, RocmLoadError> {
219 Ok(MatrixDimensions {
220 rows: i32::try_from(rows).map_err(|_| error("ROCm row count exceeds i32"))?,
221 inner: i32::try_from(inner).map_err(|_| error("ROCm inner count exceeds i32"))?,
222 cols: i32::try_from(cols).map_err(|_| error("ROCm column count exceeds i32"))?,
223 })
224}
225
226fn validate_matrix_lengths(
227 left: usize,
228 right: usize,
229 rows: usize,
230 inner: usize,
231 cols: usize,
232) -> Result<(), RocmLoadError> {
233 if rows.checked_mul(inner) != Some(left) || inner.checked_mul(cols) != Some(right) {
234 return Err(error("ROCm matmul shape does not match input lengths"));
235 }
236 Ok(())
237}
238
239fn byte_count(len: usize) -> Result<usize, RocmLoadError> {
240 len.checked_mul(std::mem::size_of::<f32>())
241 .ok_or_else(|| error("ROCm byte count overflowed"))
242}
243
244fn symbol<'library, T>(
245 library: &'library Library,
246 name: &[u8],
247) -> Result<libloading::Symbol<'library, T>, RocmLoadError> {
248 unsafe { library.get(name) }.map_err(|load| error(load.to_string()))
250}
251
252fn check_hip(status: i32, operation: &str) -> Result<(), RocmLoadError> {
253 (status == 0)
254 .then_some(())
255 .ok_or_else(|| error(format!("{operation} failed with HIP status {status}")))
256}
257
258fn check_rocblas(status: i32, operation: &str) -> Result<(), RocmLoadError> {
259 (status == 0)
260 .then_some(())
261 .ok_or_else(|| error(format!("{operation} failed with rocBLAS status {status}")))
262}
263
264fn error(message: impl Into<String>) -> RocmLoadError {
265 RocmLoadError {
266 message: message.into(),
267 }
268}