Skip to main content

sim_lib_compute_rocm/
runtime.rs

1//! Narrow runtime-loaded HIP/rocBLAS execution boundary.
2
3use 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
42// SAFETY: HIP device allocations are opaque integer handles. HIP and rocBLAS
43// calls are thread-safe, and the allocation is freed once after its last Arc.
44unsafe impl Send for RocmDeviceBuffer {}
45// SAFETY: See `Send`; device addresses are never dereferenced by host code.
46unsafe 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        // SAFETY: Function pointers use the documented HIP runtime ABI.
83        check_hip(unsafe { malloc(&mut pointer, bytes) }, "hipMalloc")?;
84        // SAFETY: HIP owns `pointer` for `bytes`, and the host slice is valid
85        // for the same byte count.
86        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        // SAFETY: Function pointers use the documented rocBLAS ABI.
127        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        // Row-major C=A*B is column-major C^T=B^T*A^T.
132        // SAFETY: Device buffers cover the validated dimensions.
133        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        // SAFETY: `handle` came from the successful create call.
153        let destroy_result = check_rocblas(unsafe { destroy(handle) }, "rocblas_destroy_handle");
154        gemm_result?;
155        destroy_result?;
156        Ok(output)
157    }
158
159    /// Runs one row-major dense f32 matrix multiplication through rocBLAS and
160    /// returns a synchronized host result.
161    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        // SAFETY: Host and device allocations both cover `bytes`.
181        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        // SAFETY: The function has no arguments and uses the documented ABI.
193        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        // SAFETY: The pointer came from `hipMalloc` and is freed once.
204        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    // SAFETY: Callers supply the official HIP/rocBLAS signature for each name.
249    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}