use std::hash::{DefaultHasher, Hash, Hasher};
use crate::RangaError;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum GpuError {
#[error("no suitable GPU adapter found")]
NoAdapter,
#[error("failed to request GPU device: {0}")]
DeviceRequest(String),
#[error("GPU buffer operation failed: {0}")]
BufferOp(String),
}
impl From<GpuError> for RangaError {
fn from(e: GpuError) -> Self {
RangaError::Other(e.to_string())
}
}
impl From<mabda::GpuError> for GpuError {
fn from(e: mabda::GpuError) -> Self {
match e {
mabda::GpuError::AdapterNotFound => GpuError::NoAdapter,
mabda::GpuError::DeviceRequest(inner) => GpuError::DeviceRequest(inner.to_string()),
other => GpuError::BufferOp(other.to_string()),
}
}
}
pub struct GpuContext {
pub(super) inner: mabda::GpuContext,
pub(super) cache: mabda::PipelineCache,
#[allow(dead_code)]
pub(super) shader_cache: mabda::ShaderCache,
adapter_name: String,
backend: String,
}
impl GpuContext {
pub fn new() -> Result<Self, GpuError> {
let inner = pollster::block_on(mabda::GpuContext::new())?;
let info = inner.adapter.get_info();
let adapter_name = info.name.clone();
let backend = format!("{:?}", info.backend);
Ok(Self {
inner,
cache: mabda::PipelineCache::new(),
shader_cache: mabda::ShaderCache::new(),
adapter_name,
backend,
})
}
#[cfg(feature = "hwaccel")]
pub fn new_with_hwaccel() -> Result<Self, GpuError> {
let report = crate::hwaccel::probe();
if !report.has_gpu {
return Err(GpuError::NoAdapter);
}
tracing::info!(
gpu = %report.gpu_name,
memory_mb = report.gpu_memory_mb,
"GPU detected via ai-hwaccel"
);
Self::new()
}
#[must_use]
pub fn adapter_name(&self) -> &str {
&self.adapter_name
}
#[must_use]
pub fn backend_name(&self) -> &str {
&self.backend
}
#[must_use]
#[inline]
pub fn device(&self) -> &wgpu::Device {
&self.inner.device
}
#[must_use]
#[inline]
pub fn queue(&self) -> &wgpu::Queue {
&self.inner.queue
}
pub fn get_or_create_pipeline_1buf(
&mut self,
name: &'static str,
shader_src: &str,
) -> Result<&mabda::compute::ComputePipeline, RangaError> {
let key = hash_name(name);
let device = &self.inner.device;
let pipeline = self.cache.get_or_insert_compute(key, || {
let entries: &[wgpu::BindGroupLayoutEntry] = &[
wgpu::BindGroupLayoutEntry {
binding: 0,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: false },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 1,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
];
mabda::compute::ComputePipeline::with_layout(device, shader_src, "main", entries)
});
Ok(pipeline)
}
pub fn get_or_create_pipeline_3buf(
&mut self,
name: &'static str,
shader_src: &str,
) -> Result<&mabda::compute::ComputePipeline, RangaError> {
let key = hash_name(name);
let device = &self.inner.device;
let pipeline = self.cache.get_or_insert_compute(key, || {
let entries: &[wgpu::BindGroupLayoutEntry] = &[
wgpu::BindGroupLayoutEntry {
binding: 0,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 1,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: false },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 2,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
];
mabda::compute::ComputePipeline::with_layout(device, shader_src, "main", entries)
});
Ok(pipeline)
}
pub(super) fn get_or_create_pipeline_4buf(
&mut self,
name: &'static str,
shader_src: &str,
) -> Result<&mabda::compute::ComputePipeline, RangaError> {
let key = hash_name(name);
let device = &self.inner.device;
let pipeline = self.cache.get_or_insert_compute(key, || {
let entries: &[wgpu::BindGroupLayoutEntry] = &[
wgpu::BindGroupLayoutEntry {
binding: 0,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 1,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: false },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 2,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
wgpu::BindGroupLayoutEntry {
binding: 3,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
];
mabda::compute::ComputePipeline::with_layout(device, shader_src, "main", entries)
});
Ok(pipeline)
}
}
#[inline]
fn hash_name(name: &str) -> u64 {
let mut h = DefaultHasher::new();
name.hash(&mut h);
h.finish()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gpu_error_converts_to_ranga_error() {
let gpu_err = GpuError::NoAdapter;
let ranga_err: RangaError = gpu_err.into();
assert!(ranga_err.to_string().contains("no suitable GPU"));
}
#[test]
fn gpu_error_device_request_message() {
let gpu_err = GpuError::DeviceRequest("limits exceeded".into());
assert!(gpu_err.to_string().contains("limits exceeded"));
}
#[test]
fn gpu_error_buffer_op_message() {
let gpu_err = GpuError::BufferOp("map failed".into());
assert!(gpu_err.to_string().contains("map failed"));
}
#[test]
fn mabda_error_converts_to_gpu_error() {
let mabda_err = mabda::GpuError::AdapterNotFound;
let gpu_err: GpuError = mabda_err.into();
assert!(matches!(gpu_err, GpuError::NoAdapter));
}
}