use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
use super::{Backend, DeviceError, DeviceInfo, Handle, Plan};
fn instance() -> Option<&'static wgpu::Instance> {
static INSTANCE: OnceLock<Option<wgpu::Instance>> = OnceLock::new();
INSTANCE
.get_or_init(|| {
if wgpu::Instance::enabled_backend_features().is_empty() {
return None;
}
Some(wgpu::Instance::new(
wgpu::InstanceDescriptor::new_without_display_handle_from_env(),
))
})
.as_ref()
}
fn kind_name(t: wgpu::DeviceType) -> &'static str {
match t {
wgpu::DeviceType::DiscreteGpu => "discrete GPU",
wgpu::DeviceType::IntegratedGpu => "integrated GPU",
wgpu::DeviceType::VirtualGpu => "virtual GPU",
wgpu::DeviceType::Cpu => "CPU",
wgpu::DeviceType::Other => "other",
}
}
fn describe(a: &wgpu::Adapter) -> DeviceInfo {
let info = a.get_info();
DeviceInfo {
name: info.name,
backend: info.backend.to_string(),
kind: kind_name(info.device_type).to_string(),
f64: a.features().contains(wgpu::Features::SHADER_F64),
}
}
pub(super) fn enumerate() -> Vec<DeviceInfo> {
let Some(inst) = instance() else { return Vec::new() };
pollster::block_on(inst.enumerate_adapters(wgpu::Backends::all()))
.iter()
.map(describe)
.collect()
}
pub(super) fn shared() -> Option<Arc<dyn Backend>> {
static GPU: OnceLock<Option<Arc<Gpu>>> = OnceLock::new();
let gpu = GPU.get_or_init(open).clone()?;
Some(gpu)
}
fn open() -> Option<Arc<Gpu>> {
let inst = instance()?;
let adapter = pollster::block_on(inst.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
force_fallback_adapter: false,
compatible_surface: None,
apply_limit_buckets: false,
}))
.ok()?;
let info = describe(&adapter);
let mut features = wgpu::Features::empty();
if info.f64 {
features |= wgpu::Features::SHADER_F64;
}
let (device, queue) = pollster::block_on(adapter.request_device(&wgpu::DeviceDescriptor {
label: Some("libjay"),
required_features: features,
required_limits: adapter.limits(),
..Default::default()
}))
.ok()?;
device.on_uncaptured_error(Arc::new(|e| {
eprintln!("libjay: the device reported an error: {e}");
}));
Some(Arc::new(Gpu {
info,
device,
queue,
shaders: Mutex::new(HashMap::new()),
pipelines: Mutex::new(HashMap::new()),
}))
}
struct Gpu {
info: DeviceInfo,
device: wgpu::Device,
queue: wgpu::Queue,
shaders: Mutex<HashMap<String, wgpu::ShaderModule>>,
pipelines: Mutex<HashMap<(String, String), wgpu::ComputePipeline>>,
}
impl Gpu {
fn scoped<T>(&self, f: impl FnOnce() -> T) -> Result<T, DeviceError> {
let scope = self.device.push_error_scope(wgpu::ErrorFilter::Validation);
let v = f();
match pollster::block_on(scope.pop()) {
None => Ok(v),
Some(e) => Err(DeviceError(e.to_string())),
}
}
fn pipeline(&self, source: &str, entry: &str) -> Result<wgpu::ComputePipeline, DeviceError> {
let key = (source.to_string(), entry.to_string());
if let Some(p) = self.pipelines.lock().expect("pipeline cache").get(&key) {
return Ok(p.clone());
}
let module = {
let mut shaders = self.shaders.lock().expect("shader cache");
match shaders.get(source) {
Some(m) => m.clone(),
None => {
let m = self.scoped(|| {
self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("libjay kernel"),
source: wgpu::ShaderSource::Wgsl(source.into()),
})
})?;
shaders.insert(source.to_string(), m.clone());
m
}
}
};
let pipeline = self.scoped(|| {
self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("libjay kernel"),
layout: None,
module: &module,
entry_point: Some(entry),
compilation_options: Default::default(),
cache: None,
})
})?;
self.pipelines.lock().expect("pipeline cache").insert(key, pipeline.clone());
Ok(pipeline)
}
}
const MAP_ALIGN: usize = 8;
impl Backend for Gpu {
fn info(&self) -> &DeviceInfo {
&self.info
}
fn upload(&self, values: &[f64], p: super::Precision) -> Result<Handle, DeviceError> {
let bytes = values.len() * p.size();
let size = bytes.max(MAP_ALIGN).next_multiple_of(MAP_ALIGN) as u64;
let buffer = self.scoped(|| {
self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("libjay input"),
size,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: true,
})
})?;
{
let mut view = buffer
.get_mapped_range_mut(..)
.map_err(|e| DeviceError(format!("mapping an input buffer: {e}")))?;
view.slice(..bytes).write_iter(super::codegen::byte_iter(values, p));
}
buffer.unmap();
Ok(Handle(Arc::new(buffer)))
}
fn dispatch(&self, plan: &Plan<'_>) -> Result<Vec<u8>, DeviceError> {
let pipeline = self.pipeline(plan.source, plan.entry)?;
let bytes = (plan.out_elems * plan.elem_size).next_multiple_of(MAP_ALIGN) as u64;
self.scoped(|| self.run(plan, &pipeline, bytes))?
.map(|mut v| {
v.truncate(plan.out_elems * plan.elem_size);
v
})
}
}
impl Gpu {
fn run(
&self,
plan: &Plan<'_>,
pipeline: &wgpu::ComputePipeline,
bytes: u64,
) -> Result<Vec<u8>, DeviceError> {
let meta = [plan.n, plan.stride, 0u32, 0u32];
let meta_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("libjay meta"),
size: 16,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
self.queue.write_buffer(&meta_buf, 0, bytemuck_u32(&meta));
let out_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("libjay out"),
size: bytes,
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("libjay readback"),
size: bytes,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut entries = vec![
wgpu::BindGroupEntry { binding: 0, resource: meta_buf.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: out_buf.as_entire_binding() },
];
let inputs: Vec<&wgpu::Buffer> = plan
.inputs
.iter()
.map(|h| {
h.0.downcast_ref::<wgpu::Buffer>()
.expect("a handle this backend made")
})
.collect();
for (i, b) in inputs.iter().enumerate() {
entries.push(wgpu::BindGroupEntry {
binding: i as u32 + 2,
resource: b.as_entire_binding(),
});
}
let bind = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("libjay bindings"),
layout: &pipeline.get_bind_group_layout(0),
entries: &entries,
});
let mut enc = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("libjay") });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("libjay kernel"),
timestamp_writes: None,
});
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bind, &[]);
pass.dispatch_workgroups(plan.groups, 1, 1);
}
enc.copy_buffer_to_buffer(&out_buf, 0, &staging, 0, bytes);
self.queue.submit(Some(enc.finish()));
let done = Arc::new(Mutex::new(None));
let flag = done.clone();
staging.map_async(wgpu::MapMode::Read, .., move |r| {
*flag.lock().expect("map result") = Some(r);
});
self.device
.poll(wgpu::PollType::wait_indefinitely())
.map_err(|e| DeviceError(format!("waiting for the device: {e}")))?;
match done.lock().expect("map result").take() {
Some(Ok(())) => {}
Some(Err(e)) => return Err(DeviceError(format!("reading back: {e}"))),
None => return Err(DeviceError("the device never finished".into())),
}
let out = staging
.get_mapped_range(..)
.map_err(|e| DeviceError(format!("mapping the readback buffer: {e}")))?
.to_vec();
staging.unmap();
Ok(out)
}
}
fn bytemuck_u32(v: &[u32; 4]) -> &[u8] {
unsafe { std::slice::from_raw_parts(v.as_ptr().cast::<u8>(), std::mem::size_of_val(v)) }
}