use crate::{Error, Result};
mod shaders {
pub const REDUCE_MIN: &str = r#"
@group(0) @binding(0) var<storage, read> input: array<f32>;
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
var<workgroup> shared_data: array<f32, 256>;
@compute @workgroup_size(256)
fn main(
@builtin(global_invocation_id) global_id: vec3<u32>,
@builtin(local_invocation_id) local_id: vec3<u32>,
@builtin(workgroup_id) workgroup_id: vec3<u32>,
) {
let tid = local_id.x;
let gid = global_id.x;
let input_size = arrayLength(&input);
// Load data into shared memory
if (gid < input_size) {
shared_data[tid] = input[gid];
} else {
shared_data[tid] = 3.4028235e+38; // f32::MAX
}
workgroupBarrier();
// Parallel reduction in shared memory
var stride: u32 = 128u;
while (stride > 0u) {
if (tid < stride && tid + stride < 256u) {
shared_data[tid] = min(shared_data[tid], shared_data[tid + stride]);
}
workgroupBarrier();
stride = stride >> 1u;
}
// Write result from first thread of workgroup
if (tid == 0u) {
output[workgroup_id.x] = shared_data[0];
}
}
"#;
pub const REDUCE_SUM: &str = r#"
@group(0) @binding(0) var<storage, read> input: array<f32>;
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
var<workgroup> shared_data: array<f32, 256>;
@compute @workgroup_size(256)
fn main(
@builtin(global_invocation_id) global_id: vec3<u32>,
@builtin(local_invocation_id) local_id: vec3<u32>,
@builtin(workgroup_id) workgroup_id: vec3<u32>,
) {
let tid = local_id.x;
let gid = global_id.x;
let input_size = arrayLength(&input);
// Load data into shared memory
if (gid < input_size) {
shared_data[tid] = input[gid];
} else {
shared_data[tid] = 0.0;
}
workgroupBarrier();
// Parallel reduction in shared memory
var stride: u32 = 128u;
while (stride > 0u) {
if (tid < stride && tid + stride < 256u) {
shared_data[tid] = shared_data[tid] + shared_data[tid + stride];
}
workgroupBarrier();
stride = stride >> 1u;
}
// Write result from first thread of workgroup
if (tid == 0u) {
output[workgroup_id.x] = shared_data[0];
}
}
"#;
pub const ELEMENTWISE_MIN: &str = r#"
@group(0) @binding(0) var<storage, read> a: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let idx = global_id.x;
let size = arrayLength(&a);
if (idx < size) {
result[idx] = min(a[idx], b[idx]);
}
}
"#;
pub const ELEMENTWISE_ADD: &str = r#"
@group(0) @binding(0) var<storage, read> a: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let idx = global_id.x;
let size = arrayLength(&a);
if (idx < size) {
result[idx] = a[idx] + b[idx];
}
}
"#;
pub const ELEMENTWISE_MUL: &str = r#"
@group(0) @binding(0) var<storage, read> a: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let idx = global_id.x;
let size = arrayLength(&a);
if (idx < size) {
result[idx] = a[idx] * b[idx];
}
}
"#;
pub const LOG_SUM_EXP: &str = r#"
@group(0) @binding(0) var<storage, read> input: array<f32>;
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
var<workgroup> shared_max: array<f32, 256>;
var<workgroup> shared_sum: array<f32, 256>;
@compute @workgroup_size(256)
fn main(
@builtin(global_invocation_id) global_id: vec3<u32>,
@builtin(local_invocation_id) local_id: vec3<u32>,
@builtin(workgroup_id) workgroup_id: vec3<u32>,
) {
let tid = local_id.x;
let gid = global_id.x;
let input_size = arrayLength(&input);
// Load data and find max
var val: f32;
if (gid < input_size) {
val = -input[gid]; // Negate for log semiring
shared_max[tid] = val;
} else {
val = -3.4028235e+38;
shared_max[tid] = -3.4028235e+38;
}
workgroupBarrier();
// First pass: find maximum for numerical stability
var stride: u32 = 128u;
while (stride > 0u) {
if (tid < stride && tid + stride < 256u) {
shared_max[tid] = max(shared_max[tid], shared_max[tid + stride]);
}
workgroupBarrier();
stride = stride >> 1u;
}
let max_val = shared_max[0];
workgroupBarrier();
// Second pass: compute sum of exp(x - max)
if (gid < input_size) {
shared_sum[tid] = exp(val - max_val);
} else {
shared_sum[tid] = 0.0;
}
workgroupBarrier();
stride = 128u;
while (stride > 0u) {
if (tid < stride && tid + stride < 256u) {
shared_sum[tid] = shared_sum[tid] + shared_sum[tid + stride];
}
workgroupBarrier();
stride = stride >> 1u;
}
// Final result: -log(sum) + max (negated back to log semiring convention)
if (tid == 0u) {
output[workgroup_id.x] = -(log(shared_sum[0]) + max_val);
}
}
"#;
pub const BATCH_WEIGHT_SCALE: &str = r#"
@group(0) @binding(0) var<storage, read_write> weights: array<f32>;
@group(0) @binding(1) var<uniform> scale: f32;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let idx = global_id.x;
let size = arrayLength(&weights);
if (idx < size) {
weights[idx] = weights[idx] + scale; // For tropical semiring, scaling is addition
}
}
"#;
#[allow(dead_code)]
pub const PREFIX_SUM: &str = r#"
@group(0) @binding(0) var<storage, read_write> data: array<f32>;
var<workgroup> temp: array<f32, 512>;
@compute @workgroup_size(256)
fn main(
@builtin(local_invocation_id) local_id: vec3<u32>,
@builtin(workgroup_id) workgroup_id: vec3<u32>,
) {
let tid = local_id.x;
let n = 512u;
let offset = workgroup_id.x * n;
// Load data into shared memory
temp[2u * tid] = data[offset + 2u * tid];
temp[2u * tid + 1u] = data[offset + 2u * tid + 1u];
// Build sum in place up the tree
var d: u32 = n >> 1u;
var stride: u32 = 1u;
while (d > 0u) {
workgroupBarrier();
if (tid < d) {
let ai = stride * (2u * tid + 1u) - 1u;
let bi = stride * (2u * tid + 2u) - 1u;
temp[bi] = min(temp[ai], temp[bi]);
}
stride = stride << 1u;
d = d >> 1u;
}
// Clear the last element
if (tid == 0u) {
temp[n - 1u] = 3.4028235e+38;
}
// Traverse down tree and build scan
d = 1u;
while (d < n) {
stride = stride >> 1u;
workgroupBarrier();
if (tid < d) {
let ai = stride * (2u * tid + 1u) - 1u;
let bi = stride * (2u * tid + 2u) - 1u;
let t = temp[ai];
temp[ai] = temp[bi];
temp[bi] = min(t, temp[bi]);
}
d = d << 1u;
}
workgroupBarrier();
// Write results
data[offset + 2u * tid] = temp[2u * tid];
data[offset + 2u * tid + 1u] = temp[2u * tid + 1u];
}
"#;
}
pub struct GpuContext {
device: wgpu::Device,
queue: wgpu::Queue,
adapter_info: wgpu::AdapterInfo,
reduce_min_pipeline: wgpu::ComputePipeline,
reduce_sum_pipeline: wgpu::ComputePipeline,
elementwise_min_pipeline: wgpu::ComputePipeline,
elementwise_add_pipeline: wgpu::ComputePipeline,
elementwise_mul_pipeline: wgpu::ComputePipeline,
log_sum_exp_pipeline: wgpu::ComputePipeline,
#[allow(dead_code)]
batch_weight_scale_pipeline: wgpu::ComputePipeline,
}
impl GpuContext {
pub async fn new() -> Result<Self> {
let instance = wgpu::Instance::new(&wgpu::InstanceDescriptor {
backends: wgpu::Backends::all(),
..Default::default()
});
let adapter = instance
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
compatible_surface: None,
force_fallback_adapter: false,
})
.await
.ok_or_else(|| Error::Algorithm("No GPU adapter found".to_string()))?;
let adapter_info = adapter.get_info();
let (device, queue) = adapter
.request_device(
&wgpu::DeviceDescriptor {
label: Some("ArcWeight GPU"),
required_features: wgpu::Features::empty(),
required_limits: wgpu::Limits::default(),
memory_hints: Default::default(),
},
None,
)
.await
.map_err(|e| Error::Algorithm(format!("Failed to create GPU device: {}", e)))?;
let reduce_min_pipeline = Self::create_pipeline(&device, shaders::REDUCE_MIN, "reduce_min");
let reduce_sum_pipeline = Self::create_pipeline(&device, shaders::REDUCE_SUM, "reduce_sum");
let elementwise_min_pipeline =
Self::create_pipeline(&device, shaders::ELEMENTWISE_MIN, "elementwise_min");
let elementwise_add_pipeline =
Self::create_pipeline(&device, shaders::ELEMENTWISE_ADD, "elementwise_add");
let elementwise_mul_pipeline =
Self::create_pipeline(&device, shaders::ELEMENTWISE_MUL, "elementwise_mul");
let log_sum_exp_pipeline =
Self::create_pipeline(&device, shaders::LOG_SUM_EXP, "log_sum_exp");
let batch_weight_scale_pipeline =
Self::create_pipeline(&device, shaders::BATCH_WEIGHT_SCALE, "batch_weight_scale");
Ok(Self {
device,
queue,
adapter_info,
reduce_min_pipeline,
reduce_sum_pipeline,
elementwise_min_pipeline,
elementwise_add_pipeline,
elementwise_mul_pipeline,
log_sum_exp_pipeline,
batch_weight_scale_pipeline,
})
}
fn create_pipeline(
device: &wgpu::Device,
shader_source: &str,
label: &str,
) -> wgpu::ComputePipeline {
let shader_module = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(shader_source.into()),
});
device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(label),
layout: None, module: &shader_module,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
})
}
pub fn adapter_info(&self) -> &wgpu::AdapterInfo {
&self.adapter_info
}
pub fn device(&self) -> &wgpu::Device {
&self.device
}
pub fn queue(&self) -> &wgpu::Queue {
&self.queue
}
pub fn create_weight_buffer(&self, weights: &[f32]) -> Result<GpuWeightBuffer> {
use wgpu::util::DeviceExt;
let buffer = self
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Weight Buffer"),
contents: bytemuck::cast_slice(weights),
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
});
Ok(GpuWeightBuffer {
buffer,
len: weights.len(),
})
}
fn create_output_buffer(&self, size: usize) -> wgpu::Buffer {
self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Output Buffer"),
size: (size * std::mem::size_of::<f32>()) as u64,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
pub fn reduce_min_gpu(&self, weights: &GpuWeightBuffer) -> Result<f32> {
if weights.is_empty() {
return Ok(f32::INFINITY);
}
let workgroup_size = 256usize;
let mut current_size = weights.len;
let mut input_buffer = &weights.buffer;
let mut temp_buffers: Vec<wgpu::Buffer> = Vec::new();
while current_size > 1 {
let num_workgroups = current_size.div_ceil(workgroup_size);
let output_buffer = self.create_output_buffer(num_workgroups);
let bind_group_layout = self.reduce_min_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Reduce Min Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input_buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: output_buffer.as_entire_binding(),
},
],
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Reduce Min Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Reduce Min Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.reduce_min_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
temp_buffers.push(output_buffer);
input_buffer = temp_buffers.last().unwrap();
current_size = num_workgroups;
}
let staging_buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Staging Buffer"),
size: std::mem::size_of::<f32>() as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Copy Encoder"),
});
encoder.copy_buffer_to_buffer(
input_buffer,
0,
&staging_buffer,
0,
std::mem::size_of::<f32>() as u64,
);
self.queue.submit(std::iter::once(encoder.finish()));
let buffer_slice = staging_buffer.slice(..);
let (sender, receiver) = std::sync::mpsc::channel();
buffer_slice.map_async(wgpu::MapMode::Read, move |result| {
let _ = sender.send(result);
});
self.device.poll(wgpu::Maintain::Wait);
receiver
.recv()
.map_err(|_| Error::Algorithm("Failed to receive buffer mapping result".to_string()))?
.map_err(|e| Error::Algorithm(format!("Buffer mapping failed: {:?}", e)))?;
let data = buffer_slice.get_mapped_range();
let result: f32 = bytemuck::cast_slice(&data)[0];
drop(data);
staging_buffer.unmap();
Ok(result)
}
pub fn reduce_sum_gpu(&self, weights: &GpuWeightBuffer) -> Result<f32> {
if weights.is_empty() {
return Ok(0.0);
}
let workgroup_size = 256usize;
let mut current_size = weights.len;
let mut input_buffer = &weights.buffer;
let mut temp_buffers: Vec<wgpu::Buffer> = Vec::new();
while current_size > 1 {
let num_workgroups = current_size.div_ceil(workgroup_size);
let output_buffer = self.create_output_buffer(num_workgroups);
let bind_group_layout = self.reduce_sum_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Reduce Sum Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input_buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: output_buffer.as_entire_binding(),
},
],
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Reduce Sum Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Reduce Sum Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.reduce_sum_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
temp_buffers.push(output_buffer);
input_buffer = temp_buffers.last().unwrap();
current_size = num_workgroups;
}
let staging_buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Staging Buffer"),
size: std::mem::size_of::<f32>() as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Copy Encoder"),
});
encoder.copy_buffer_to_buffer(
input_buffer,
0,
&staging_buffer,
0,
std::mem::size_of::<f32>() as u64,
);
self.queue.submit(std::iter::once(encoder.finish()));
let buffer_slice = staging_buffer.slice(..);
let (sender, receiver) = std::sync::mpsc::channel();
buffer_slice.map_async(wgpu::MapMode::Read, move |result| {
let _ = sender.send(result);
});
self.device.poll(wgpu::Maintain::Wait);
receiver
.recv()
.map_err(|_| Error::Algorithm("Failed to receive buffer mapping result".to_string()))?
.map_err(|e| Error::Algorithm(format!("Buffer mapping failed: {:?}", e)))?;
let data = buffer_slice.get_mapped_range();
let result: f32 = bytemuck::cast_slice(&data)[0];
drop(data);
staging_buffer.unmap();
Ok(result)
}
pub fn reduce_min(&self, weights: &GpuWeightBuffer) -> Result<f32> {
self.reduce_min_gpu(weights)
}
pub fn reduce_sum(&self, weights: &GpuWeightBuffer) -> Result<f32> {
self.reduce_sum_gpu(weights)
}
pub fn read_buffer(&self, buffer: &GpuWeightBuffer) -> Result<Vec<f32>> {
let staging_buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Staging Buffer"),
size: (buffer.len * std::mem::size_of::<f32>()) as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Read Encoder"),
});
encoder.copy_buffer_to_buffer(
&buffer.buffer,
0,
&staging_buffer,
0,
(buffer.len * std::mem::size_of::<f32>()) as u64,
);
self.queue.submit(std::iter::once(encoder.finish()));
let buffer_slice = staging_buffer.slice(..);
let (sender, receiver) = std::sync::mpsc::channel();
buffer_slice.map_async(wgpu::MapMode::Read, move |result| {
let _ = sender.send(result);
});
self.device.poll(wgpu::Maintain::Wait);
receiver
.recv()
.map_err(|_| Error::Algorithm("Failed to receive buffer mapping result".to_string()))?
.map_err(|e| Error::Algorithm(format!("Buffer mapping failed: {:?}", e)))?;
let data = buffer_slice.get_mapped_range();
let result: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
staging_buffer.unmap();
Ok(result)
}
pub fn elementwise_min_gpu(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
if a.len != b.len {
return Err(Error::InvalidOperation(
"Buffer lengths must match".to_string(),
));
}
let result_buffer = self.create_output_buffer(a.len);
let bind_group_layout = self.elementwise_min_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Elementwise Min Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: b.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: result_buffer.as_entire_binding(),
},
],
});
let num_workgroups = a.len.div_ceil(256);
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Elementwise Min Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Elementwise Min Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.elementwise_min_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
Ok(GpuWeightBuffer {
buffer: result_buffer,
len: a.len,
})
}
pub fn elementwise_min(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
self.elementwise_min_gpu(a, b)
}
pub fn elementwise_add_gpu(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
if a.len != b.len {
return Err(Error::InvalidOperation(
"Buffer lengths must match".to_string(),
));
}
let result_buffer = self.create_output_buffer(a.len);
let bind_group_layout = self.elementwise_add_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Elementwise Add Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: b.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: result_buffer.as_entire_binding(),
},
],
});
let num_workgroups = a.len.div_ceil(256);
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Elementwise Add Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Elementwise Add Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.elementwise_add_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
Ok(GpuWeightBuffer {
buffer: result_buffer,
len: a.len,
})
}
pub fn elementwise_add(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
self.elementwise_add_gpu(a, b)
}
pub fn elementwise_mul_gpu(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
if a.len != b.len {
return Err(Error::InvalidOperation(
"Buffer lengths must match".to_string(),
));
}
let result_buffer = self.create_output_buffer(a.len);
let bind_group_layout = self.elementwise_mul_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Elementwise Mul Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: b.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: result_buffer.as_entire_binding(),
},
],
});
let num_workgroups = a.len.div_ceil(256);
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Elementwise Mul Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Elementwise Mul Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.elementwise_mul_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
Ok(GpuWeightBuffer {
buffer: result_buffer,
len: a.len,
})
}
pub fn elementwise_mul(
&self,
a: &GpuWeightBuffer,
b: &GpuWeightBuffer,
) -> Result<GpuWeightBuffer> {
self.elementwise_mul_gpu(a, b)
}
pub fn log_sum_exp_gpu(&self, weights: &GpuWeightBuffer) -> Result<f32> {
if weights.is_empty() {
return Ok(f32::INFINITY); }
let workgroup_size = 256usize;
let mut current_size = weights.len;
let mut input_buffer = &weights.buffer;
let mut temp_buffers: Vec<wgpu::Buffer> = Vec::new();
while current_size > 1 {
let num_workgroups = current_size.div_ceil(workgroup_size);
let output_buffer = self.create_output_buffer(num_workgroups);
let bind_group_layout = self.log_sum_exp_pipeline.get_bind_group_layout(0);
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("Log Sum Exp Bind Group"),
layout: &bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input_buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: output_buffer.as_entire_binding(),
},
],
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Log Sum Exp Encoder"),
});
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("Log Sum Exp Pass"),
timestamp_writes: None,
});
compute_pass.set_pipeline(&self.log_sum_exp_pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
compute_pass.dispatch_workgroups(num_workgroups as u32, 1, 1);
}
self.queue.submit(std::iter::once(encoder.finish()));
temp_buffers.push(output_buffer);
input_buffer = temp_buffers.last().unwrap();
current_size = num_workgroups;
}
let staging_buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Staging Buffer"),
size: std::mem::size_of::<f32>() as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut encoder = self
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("Copy Encoder"),
});
encoder.copy_buffer_to_buffer(
input_buffer,
0,
&staging_buffer,
0,
std::mem::size_of::<f32>() as u64,
);
self.queue.submit(std::iter::once(encoder.finish()));
let buffer_slice = staging_buffer.slice(..);
let (sender, receiver) = std::sync::mpsc::channel();
buffer_slice.map_async(wgpu::MapMode::Read, move |result| {
let _ = sender.send(result);
});
self.device.poll(wgpu::Maintain::Wait);
receiver
.recv()
.map_err(|_| Error::Algorithm("Failed to receive buffer mapping result".to_string()))?
.map_err(|e| Error::Algorithm(format!("Buffer mapping failed: {:?}", e)))?;
let data = buffer_slice.get_mapped_range();
let result: f32 = bytemuck::cast_slice(&data)[0];
drop(data);
staging_buffer.unmap();
Ok(result)
}
pub fn batch_reduce_min(&self, weights_list: &[&[f32]]) -> Result<Vec<f32>> {
let mut results = Vec::with_capacity(weights_list.len());
for weights in weights_list {
let buffer = self.create_weight_buffer(weights)?;
let min = self.reduce_min_gpu(&buffer)?;
results.push(min);
}
Ok(results)
}
}
impl std::fmt::Debug for GpuContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GpuContext")
.field("adapter", &self.adapter_info.name)
.field("backend", &self.adapter_info.backend)
.finish()
}
}
pub struct GpuWeightBuffer {
buffer: wgpu::Buffer,
len: usize,
}
impl GpuWeightBuffer {
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
}
impl std::fmt::Debug for GpuWeightBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GpuWeightBuffer")
.field("len", &self.len)
.finish()
}
}
pub fn is_gpu_available() -> bool {
let instance = wgpu::Instance::new(&wgpu::InstanceDescriptor {
backends: wgpu::Backends::all(),
..Default::default()
});
pollster::block_on(async {
instance
.request_adapter(&wgpu::RequestAdapterOptions::default())
.await
.is_some()
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gpu_availability_check() {
let _available = is_gpu_available();
}
#[test]
fn test_gpu_context_creation() {
if !is_gpu_available() {
eprintln!("Skipping GPU test - no GPU available");
return;
}
let ctx = pollster::block_on(GpuContext::new());
assert!(ctx.is_ok(), "Failed to create GPU context: {:?}", ctx.err());
let ctx = ctx.unwrap();
println!(
"GPU: {} ({:?})",
ctx.adapter_info().name,
ctx.adapter_info().backend
);
}
#[test]
fn test_weight_buffer_operations() {
if !is_gpu_available() {
eprintln!("Skipping GPU test - no GPU available");
return;
}
let ctx = pollster::block_on(GpuContext::new()).unwrap();
let weights = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
let buffer = ctx.create_weight_buffer(&weights).unwrap();
assert_eq!(buffer.len(), 5);
let min = ctx.reduce_min_gpu(&buffer).unwrap();
assert_eq!(min, 1.0);
let sum = ctx.reduce_sum_gpu(&buffer).unwrap();
assert_eq!(sum, 15.0);
}
#[test]
fn test_elementwise_operations() {
if !is_gpu_available() {
eprintln!("Skipping GPU test - no GPU available");
return;
}
let ctx = pollster::block_on(GpuContext::new()).unwrap();
let a = ctx.create_weight_buffer(&[1.0, 3.0, 5.0, 7.0]).unwrap();
let b = ctx.create_weight_buffer(&[2.0, 2.0, 6.0, 4.0]).unwrap();
let min_result = ctx.elementwise_min_gpu(&a, &b).unwrap();
let min_data = ctx.read_buffer(&min_result).unwrap();
assert_eq!(min_data, vec![1.0, 2.0, 5.0, 4.0]);
let add_result = ctx.elementwise_add_gpu(&a, &b).unwrap();
let add_data = ctx.read_buffer(&add_result).unwrap();
assert_eq!(add_data, vec![3.0, 5.0, 11.0, 11.0]);
let mul_result = ctx.elementwise_mul_gpu(&a, &b).unwrap();
let mul_data = ctx.read_buffer(&mul_result).unwrap();
assert_eq!(mul_data, vec![2.0, 6.0, 30.0, 28.0]);
}
#[test]
fn test_large_reduction() {
if !is_gpu_available() {
eprintln!("Skipping GPU test - no GPU available");
return;
}
let ctx = pollster::block_on(GpuContext::new()).unwrap();
let weights: Vec<f32> = (0..10000).map(|i| i as f32).collect();
let buffer = ctx.create_weight_buffer(&weights).unwrap();
let min = ctx.reduce_min_gpu(&buffer).unwrap();
assert_eq!(min, 0.0);
let sum = ctx.reduce_sum_gpu(&buffer).unwrap();
let expected_sum: f32 = 49995000.0;
let relative_error = (sum - expected_sum).abs() / expected_sum;
assert!(
relative_error < 0.001,
"sum={}, expected={}, relative_error={}",
sum,
expected_sum,
relative_error
);
}
#[test]
fn test_batch_reduce() {
if !is_gpu_available() {
eprintln!("Skipping GPU test - no GPU available");
return;
}
let ctx = pollster::block_on(GpuContext::new()).unwrap();
let arrays: Vec<Vec<f32>> = vec![
vec![5.0, 3.0, 7.0, 1.0],
vec![10.0, 20.0, 5.0, 15.0],
vec![100.0, 50.0, 75.0, 25.0],
];
let array_refs: Vec<&[f32]> = arrays.iter().map(|v| v.as_slice()).collect();
let mins = ctx.batch_reduce_min(&array_refs).unwrap();
assert_eq!(mins, vec![1.0, 5.0, 25.0]);
}
}