File size: 3,436 Bytes
a8baeed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
use super::{GpuContext, GpuError, Result, UnaryOp};
use ocl::Buffer;
use std::fmt;

pub(super) const KERNEL_SRC: &str = r#"
__kernel void negate_u8(__global const uchar* input,
                        __global uchar* output,
                        const ulong element_count) {
    const ulong i = (ulong)get_global_id(0);
    if (i >= element_count) {
        return;
    }
    output[i] = input[i] ^ (uchar)1;
}
"#;

/// Error type for Negate validation and execution failures.
#[derive(Debug, Clone)]
pub enum NegateError {
    /// Non-canonical Boolean value found at the given index.
    InvalidBoolean { index: usize, value: u8 },
    /// OpenCL error during device operations.
    OpenCl(String),
}

impl fmt::Display for NegateError {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            NegateError::InvalidBoolean { index, value } => {
                write!(
                    f,
                    "invalid boolean at index {}: expected 0x00 or 0x01, got 0x{:02x}",
                    index, value
                )
            }
            NegateError::OpenCl(msg) => write!(f, "OpenCL error: {}", msg),
        }
    }
}

impl std::error::Error for NegateError {}

impl From<ocl::Error> for NegateError {
    fn from(e: ocl::Error) -> Self {
        NegateError::OpenCl(e.to_string())
    }
}

/// Logical negation `¬x`: `output[i] = !input[i]`, executed on the OpenCL device.
#[derive(Clone, Copy, Debug, Default)]
pub struct Negate;

impl Negate {
    /// Reference semantics of the functor on a single Boolean.
    pub fn denote(x: bool) -> bool {
        !x
    }

    /// Validate that all elements in the input buffer are canonical (0x00 or 0x01).
    /// Returns the first invalid (index, value) pair found, or Ok(()) if all canonical.
    fn validate_input(input_data: &[u8]) -> std::result::Result<(), NegateError> {
        for (i, &byte) in input_data.iter().enumerate() {
            if byte != 0x00 && byte != 0x01 {
                return Err(NegateError::InvalidBoolean {
                    index: i,
                    value: byte,
                });
            }
        }
        Ok(())
    }
}

impl UnaryOp for Negate {
    fn launch(
        &self,
        ctx: &GpuContext,
        input: &Buffer<u8>,
        output: &Buffer<u8>,
        count: usize,
    ) -> Result<()> {
        let (global, local) =
            ctx.work_sizes(count, &[("input", input.len()), ("output", output.len())])?;
        if count == 0 {
            return Ok(());
        }

        // Download and validate input data before kernel submission
        let input_data = ctx.download_raw(input)?;
        let input_slice = input_data.get(..count).ok_or(GpuError::BufferTooShort {
            which: "input",
            len: input_data.len(),
            count,
        })?;

        Self::validate_input(input_slice)
            .map_err(|e| GpuError::Ocl(ocl::Error::from(e.to_string())))?;

        let kernel = ctx
            .proque()
            .kernel_builder("negate_u8")
            .global_work_size(global)
            .local_work_size(local)
            .arg(input)
            .arg(output)
            .arg(count as u64)
            .build()?;
        // SAFETY: arguments match the kernel signature; the kernel guards gid < count and
        // count <= both buffer lengths was validated above.
        unsafe { kernel.enq()? };
        Ok(())
    }
}