//! # RustyTorch++ ROCm Backend //! //! AMD GPU backend implementation using HIP runtime. //! //! ## Features //! //! - **HIP Runtime**: AMD's CUDA-compatible API for GPU computing //! - **rocBLAS**: Optimized matrix operations //! - **MIOpen**: Deep learning primitives //! - **Source Compatibility**: HIP kernels are ~95% CUDA-compatible //! //! ## Architecture //! //! ```text //! RocmBackend //! ├── RocmTensorPrimitive - HIP device memory //! ├── RocmDevice - Device context and streams //! └── Ops //! ├── Basic - Add, mul, etc. (HIP kernels) //! ├── GEMM - Matrix multiply (rocBLAS) //! └── Attention - Flash Attention (ported from CUDA) //! ``` //! //! ## Example //! //! ```rust,ignore //! use rtx_backend_rocm::{RocmBackend, RocmDevice}; //! use rtx_backend::Backend; //! //! let device = RocmDevice::new(0)?; //! let a = RocmBackend::zeros([1024, 1024], &device); //! let b = RocmBackend::randn([1024, 1024], &device); //! let c = RocmBackend::matmul(&a, &b); //! ``` #![warn(missing_docs)] mod device; mod error; /// HIP FFI bindings for low-level GPU access. pub mod hip_ffi; /// HIP kernel implementations for GPU operations. #[cfg(feature = "hip-runtime")] pub mod hip_kernels; /// Operations module - re-exported for tests and direct access. pub mod ops; mod tensor; pub use device::{RocmDevice, RocmDeviceInfo}; pub use error::{RocmBackendError, RocmBackendResult}; pub use tensor::{RocmBuffer, RocmTensorPrimitive}; use rtx_backend::{Backend, BoolU8}; /// ROCm backend for RustyTorch++. /// /// This backend uses AMD GPUs via HIP runtime and provides: /// - rocBLAS for matrix operations /// - HIP kernels (CUDA source-compatible) /// - Multi-GPU support #[derive(Clone, Debug, Default)] pub struct RocmBackend; impl Backend for RocmBackend { type TensorPrimitive = RocmTensorPrimitive; type Device = RocmDevice; type FloatElem = f32; type IntElem = i32; type BoolElem = BoolU8; fn name() -> &'static str { "rocm" } fn seed(seed: u64) { ops::seed_rng(seed); } // ==================== Tensor Creation ==================== fn zeros(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::zeros(shape, device) } fn ones(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::ones(shape, device) } fn full( shape: [usize; D], fill_value: Self::FloatElem, device: &Self::Device, ) -> Self::TensorPrimitive { ops::creation::full(shape, fill_value, device) } fn rand(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::rand(shape, device) } fn randn(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::randn(shape, device) } fn from_data( data: &[Self::FloatElem], shape: [usize; D], device: &Self::Device, ) -> Self::TensorPrimitive { ops::creation::from_data(data, shape, device) } // ==================== Basic Operations ==================== fn add( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::add(&lhs, &rhs) } fn sub( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::sub(&lhs, &rhs) } fn mul( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::mul(&lhs, &rhs) } fn div( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::div(&lhs, &rhs) } fn matmul( lhs: Self::TensorPrimitive<2>, rhs: Self::TensorPrimitive<2>, ) -> Self::TensorPrimitive<2> { ops::gemm::matmul(&lhs, &rhs) } fn bmm( lhs: Self::TensorPrimitive<3>, rhs: Self::TensorPrimitive<3>, ) -> Self::TensorPrimitive<3> { ops::gemm::bmm(&lhs, &rhs) } // ==================== Unary Operations ==================== fn neg(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::neg(&tensor) } fn exp(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::exp(&tensor) } fn log(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::log(&tensor) } fn sqrt(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::sqrt(&tensor) } fn abs(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::abs(&tensor) } fn sin(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::sin(&tensor) } fn cos(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::cos(&tensor) } fn pow( tensor: Self::TensorPrimitive, exp: Self::FloatElem, ) -> Self::TensorPrimitive { ops::unary::pow(&tensor, exp) } fn clamp( tensor: Self::TensorPrimitive, min: Self::FloatElem, max: Self::FloatElem, ) -> Self::TensorPrimitive { ops::comparison::clamp(&tensor, Some(min), Some(max)) } // ==================== Activation Functions ==================== fn relu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::relu(&tensor) } fn sigmoid(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::sigmoid(&tensor) } fn tanh(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::tanh(&tensor) } // ==================== Reduction Operations ==================== fn sum(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::sum(&tensor) } fn sum_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::sum_dim(&tensor, dim) } fn mean(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::mean(&tensor) } fn mean_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::mean_dim(&tensor, dim) } fn var(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::var(&tensor) } fn var_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::var_dim(&tensor, dim) } fn max(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::max(&tensor) } fn min(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::min(&tensor) } // ==================== Shape Operations ==================== fn shape(tensor: &Self::TensorPrimitive) -> [usize; D] { tensor.shape } fn reshape( tensor: Self::TensorPrimitive, shape: [usize; D2], ) -> Self::TensorPrimitive { ops::shape::reshape(tensor, shape) } fn transpose(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::shape::transpose(&tensor) } fn swap_dims( tensor: Self::TensorPrimitive, dim1: usize, dim2: usize, ) -> Self::TensorPrimitive { ops::shape::swap_dims(&tensor, dim1, dim2) } // ==================== LLM-Specific Operations ==================== fn flash_attention( query: Self::TensorPrimitive<4>, key: Self::TensorPrimitive<4>, value: Self::TensorPrimitive<4>, mask: Option<&Self::TensorPrimitive<4>>, scale: Self::FloatElem, causal: bool, ) -> Self::TensorPrimitive<4> { ops::attention::flash_attention(&query, &key, &value, mask, scale, causal) } fn softmax( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::activation::softmax(&tensor, dim) } fn layer_norm( tensor: Self::TensorPrimitive, weight: &Self::TensorPrimitive<1>, bias: Option<&Self::TensorPrimitive<1>>, eps: Self::FloatElem, ) -> Self::TensorPrimitive { ops::normalization::layer_norm(&tensor, weight, bias, eps) } fn rms_norm( tensor: Self::TensorPrimitive, weight: &Self::TensorPrimitive<1>, eps: Self::FloatElem, ) -> Self::TensorPrimitive { ops::normalization::rms_norm(&tensor, weight, eps) } fn rope( tensor: Self::TensorPrimitive, cos: &Self::TensorPrimitive<2>, sin: &Self::TensorPrimitive<2>, ) -> Self::TensorPrimitive { ops::attention::rope(&tensor, cos, sin) } fn gelu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::gelu(&tensor) } fn silu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::silu(&tensor) } fn leaky_relu( tensor: Self::TensorPrimitive, negative_slope: Self::FloatElem, ) -> Self::TensorPrimitive { ops::activation::leaky_relu(&tensor, negative_slope) } fn elu( tensor: Self::TensorPrimitive, alpha: Self::FloatElem, ) -> Self::TensorPrimitive { ops::activation::elu(&tensor, alpha) } // ==================== Comparison Operations ==================== fn gt_scalar( tensor: Self::TensorPrimitive, value: Self::FloatElem, ) -> Self::TensorPrimitive { ops::comparison::gt_scalar(&tensor, value) } // ==================== Convolution Operations ==================== fn conv2d( input: Self::TensorPrimitive<4>, weight: &Self::TensorPrimitive<4>, bias: Option<&Self::TensorPrimitive<1>>, stride: [usize; 2], padding: [usize; 2], dilation: [usize; 2], groups: usize, ) -> Self::TensorPrimitive<4> { let config = ops::convolution::Conv2dConfig { kernel_size: (weight.shape[2], weight.shape[3]), stride: (stride[0], stride[1]), padding: (padding[0], padding[1]), dilation: (dilation[0], dilation[1]), groups, }; ops::convolution::conv2d(&input, weight, bias, &config) } // ==================== Pooling Operations ==================== fn max_pool2d( input: Self::TensorPrimitive<4>, kernel_size: [usize; 2], stride: [usize; 2], padding: [usize; 2], ) -> Self::TensorPrimitive<4> { let config = ops::pooling::Pool2dConfig { kernel_size: (kernel_size[0], kernel_size[1]), stride: (stride[0], stride[1]), padding: (padding[0], padding[1]), dilation: (1, 1), ceil_mode: false, }; ops::pooling::max_pool2d(&input, &config) } fn avg_pool2d( input: Self::TensorPrimitive<4>, kernel_size: [usize; 2], stride: [usize; 2], padding: [usize; 2], count_include_pad: bool, ) -> Self::TensorPrimitive<4> { let config = ops::pooling::Pool2dConfig { kernel_size: (kernel_size[0], kernel_size[1]), stride: (stride[0], stride[1]), padding: (padding[0], padding[1]), dilation: (1, 1), ceil_mode: false, }; ops::pooling::avg_pool2d(&input, &config, count_include_pad) } // ==================== Device Management ==================== fn device(tensor: &Self::TensorPrimitive) -> Self::Device { tensor.device.clone() } fn to_device( tensor: Self::TensorPrimitive, device: &Self::Device, ) -> Self::TensorPrimitive { if tensor.device == *device { tensor } else { ops::device::copy_to_device(tensor, device) } } fn to_data(tensor: &Self::TensorPrimitive) -> Vec { ops::device::copy_to_host(tensor) } fn sync(device: &Self::Device) { let _ = device.synchronize(); } } /// Check if ROCm is available on this system. pub fn is_available() -> bool { #[cfg(feature = "hip-runtime")] { hip_ffi::HipRuntime::is_available() } #[cfg(not(feature = "hip-runtime"))] { device::is_available() } } /// Get the number of available AMD GPUs. pub fn device_count() -> usize { hip_ffi::HipRuntime::device_count() } /// Check if HIP runtime feature is enabled. pub fn has_hip_runtime() -> bool { cfg!(feature = "hip-runtime") } /// Check if rocBLAS feature is enabled. pub fn has_rocblas() -> bool { cfg!(feature = "rocblas") } /// Check if MIOpen feature is enabled. pub fn has_miopen() -> bool { cfg!(feature = "miopen") } /// Type alias for training with ROCm + autodiff. pub type RocmTraining = RocmBackend; /// Type alias for inference with ROCm (no autodiff overhead). pub type RocmInference = RocmBackend;