Skip to main content

dispatch_gpu_sample_argmax_candidates

Function dispatch_gpu_sample_argmax_candidates 

Source
pub fn dispatch_gpu_sample_argmax_candidates(
    encoder: &mut CommandEncoder,
    registry: &mut KernelRegistry,
    device: &DeviceRef,
    logits: &MlxBuffer,
    out_top1_idx: &MlxBuffer,
    out_top1_val: &MlxBuffer,
    out_cand_count: &MlxBuffer,
    out_overflow: &MlxBuffer,
    out_cand_ids: &MlxBuffer,
    params_buf: &MlxBuffer,
    n_slots: u32,
    vocab: u32,
    cap: u32,
) -> Result<()>
Expand description

Dispatch GPU argmax+candidate-collect over [n_slots, vocab] logits.

Outputs (per slot): out_top1_idx[n], out_top1_val[n], out_cand_count[n] (atomic u32 — total count, may exceed cap), out_overflow[n] (1 if count>cap), out_cand_ids[n*cap].