logits-processor / torch-ext /torch_binding.h
CharlesCNorton
Add logits-processor source, tests, and model card
9aa33e7
Raw
History Blame
924 Bytes
#pragma once
#include <torch/torch.h>
// Set logits of disallowed tokens to -inf, in place. bitmask is packed int32,
// bit == 1 means the token is allowed (XGrammar convention); one row of
// ceil(V/32) words per batch element. logits: [B, V], bitmask: [B, ceil(V/32)].
void apply_token_bitmask_inplace(at::Tensor logits, at::Tensor bitmask);
// Fused logits -> token sampler. Applies per-row temperature, then the
// intersection of top-k / top-p (nucleus) / min-p filtering, renormalizes the
// surviving set, and draws one token per row with a counter-based RNG keyed by
// (seed, row, offset). temperature <= 0 is greedy (argmax). logits: [B, V]
// (float/half/bf16); temperature/top_p/min_p: [B] float; top_k: [B] int32.
// Returns tokens: [B] int64.
at::Tensor sample(at::Tensor logits, at::Tensor temperature, at::Tensor top_p,
at::Tensor top_k, at::Tensor min_p, int64_t seed, int64_t offset);