use ocl::{Queue, Buffer, Kernel, Context, Program, builders::DeviceSpecifier, error::Result};
#[derive(PartialEq, Eq)]
pub enum Op {
Add,
Min,
Mul,
Div,
Mod,
None
}
pub struct MapProgram(Program);
pub struct MapKernel(Kernel, usize);
impl MapProgram {
pub fn from<D: Into<DeviceSpecifier>>(devices: D, op: Op, context: &Context) -> Result<Self> {
let src = if op == Op::None {
String::from("__kernel void __main__(__global float* buffer, float scalar) {}")
} else {
format!(r#"
__kernel void __main__(__global float* buffer, float scalar) {{
buffer[get_global_id(0)] {}= scalar;
}}
"#, match op {
Op::Add => "+",
Op::Min => "-",
Op::Mul => "*",
Op::Div => "/",
Op::Mod => "%",
Op::None => panic!("creating program failed")
})
};
Program::builder()
.devices(devices)
.src(src)
.build(&context)
.map(|program| Self(program))
}
}
impl MapKernel {
pub fn from(program: &MapProgram, queue: Queue, buffer: &Buffer<f32>, val: &f32) -> Result<Self> {
let buffer_len = buffer.len();
Kernel::builder()
.program(&program.0)
.name("__main__")
.queue(queue.clone())
.global_work_size(buffer_len)
.arg(buffer)
.arg(val)
.build()
.map(|kernel| Self(kernel, buffer_len))
}
pub fn cmd_enq(&self, queue: &Queue) {
unsafe {
self.0.cmd()
.queue(&queue)
.global_work_offset(self.0.default_global_work_offset())
.global_work_size(self.1)
.local_work_size(self.0.default_local_work_size())
.enq().unwrap();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use ocl::{flags, Platform, Device, Context, Queue, Program, Buffer, Kernel};
#[test]
fn test_add_unsafe() {
let src = r#"
__kernel void add(__global float* buffer, float scalar) {
buffer[get_global_id(0)] += scalar;
}
"#;
let platform = Platform::default();
let device = Device::first(platform).unwrap();
let context = Context::builder()
.platform(platform)
.devices(device.clone())
.build().unwrap();
let program = Program::builder()
.devices(device)
.src(src)
.build(&context).unwrap();
let queue = Queue::new(&context, device, None).unwrap();
let dims = 1 << 20;
let buffer = Buffer::<f32>::builder()
.queue(queue.clone())
.flags(flags::MEM_READ_WRITE)
.len(dims)
.fill_val(0f32)
.build().unwrap();
let kernel = Kernel::builder()
.program(&program)
.name("add")
.queue(queue.clone())
.global_work_size(dims)
.arg(&buffer)
.arg(&10.0f32)
.build().unwrap();
unsafe {
kernel.cmd()
.queue(&queue)
.global_work_offset(kernel.default_global_work_offset())
.global_work_size(dims)
.local_work_size(kernel.default_local_work_size())
.enq().unwrap();
}
let mut vec = vec![0.0f32; dims];
buffer.cmd()
.queue(&queue)
.offset(0)
.read(&mut vec)
.enq().unwrap();
assert_eq!(vec, vec![10.0f32; dims]);
}
#[test]
fn test_add() {
let platform = Platform::default();
let device = Device::first(platform).unwrap();
let context = Context::builder()
.platform(platform)
.devices(device.clone())
.build().unwrap();
let program = MapProgram::from(device, Op::Add, &context).unwrap();
let queue = Queue::new(&context, device, None).unwrap();
let dims = 1 << 20;
let buffer = Buffer::<f32>::builder()
.queue(queue.clone())
.flags(flags::MEM_READ_WRITE) .len(dims)
.fill_val(0f32)
.build().unwrap();
let kernel = MapKernel::from(&program, queue.clone(), &buffer, &10.0f32).unwrap();
kernel.cmd_enq(&queue);
let mut vec = vec![0.0f32; dims];
buffer.cmd()
.queue(&queue)
.offset(0)
.read(&mut vec)
.enq().unwrap();
assert_eq!(vec, vec![10.0f32; dims]);
}
}