use crate::op_registry::{MetalKernel, register_metal_kernel};
use rlx_ir::Shape;
use std::sync::Arc;
pub const OP_NAME: &str = "llada2.group_limited_gate";
#[derive(Debug)]
struct Llada2GateMetal;
impl MetalKernel for Llada2GateMetal {
fn name(&self) -> &str {
OP_NAME
}
fn execute(
&self,
inputs: &[(&[u8], &Shape)],
output: (&mut [u8], &Shape),
attrs: &[u8],
) -> Result<(), String> {
let sig_bytes = inputs[0].0;
let route_bytes = inputs[1].0;
let out_bytes = output.0;
if !sig_bytes.len().is_multiple_of(4)
|| !route_bytes.len().is_multiple_of(4)
|| !out_bytes.len().is_multiple_of(4)
{
return Err("gate: non-f32-aligned buffers".into());
}
let sig = bytemuck::cast_slice::<u8, f32>(sig_bytes);
let route = bytemuck::cast_slice::<u8, f32>(route_bytes);
let out = bytemuck::cast_slice_mut::<u8, f32>(out_bytes);
rlx_cpu::llada2_gate::execute_gate_f32(sig, route, out, attrs)
}
}
pub fn register() {
register_metal_kernel(Arc::new(Llada2GateMetal));
}