sigmoid_gate_mul_batched

Function sigmoid_gate_mul_batched 

Source
pub fn sigmoid_gate_mul_batched(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    input: DevicePtr,
    gate: DevicePtr,
    output: DevicePtr,
    dim: u32,
    gate_stride: u32,
    num_tokens: u32,
    stream: u64,
) -> Result<()>
Expand description

Batched sigmoid gate multiply across multiple tokens.

Replaces per-token sigmoid_gate_mul launches with a single kernel. gate is strided (gate_stride elements between tokens in gate buffer).