1#ifndef INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLER_H_
2#define INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLER_H_
19 "Migrate to an open-source-aligned operator when available.")]]
23 std::optional<Tensor> p,
Tensor out)
24 : batch_size_{logits.size(0)},
25 vocab_size_{logits.size(1)},
26 dtype_{logits.dtype()} {
27 assert(logits.ndim() == 2 &&
28 "`TopKTopPSampler` requires 2D `[batch_size, vocab_size]` logits");
29 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
30 dtype_ == DataType::kFloat32 || dtype_ == DataType::kFloat64) &&
31 "`TopKTopPSampler` requires floating-point logits");
32 assert(out.ndim() == 1 &&
33 "`TopKTopPSampler` requires 1D `[batch_size]` output");
34 assert(out.size(0) == batch_size_ &&
35 "`TopKTopPSampler` requires output batch size to match logits");
36 assert(out.dtype() == DataType::kInt32 &&
37 "`TopKTopPSampler` requires int32 output");
44 std::optional<Tensor> p,
Tensor out)
const = 0;
48 if (!k.has_value())
return;
50 assert(k->ndim() == 1 &&
51 "`TopKTopPSampler` requires `k` to be 1D when provided");
52 assert((k->size(0) == 1 || k->size(0) == batch_size_) &&
53 "`TopKTopPSampler` requires `k` shape [1] or [batch_size]");
54 assert((k->dtype() == DataType::kInt32 || k->dtype() == DataType::kInt64) &&
55 "`TopKTopPSampler` requires int32 or int64 `k`");
59 if (!p.has_value())
return;
61 assert(p->ndim() == 1 &&
62 "`TopKTopPSampler` requires `p` to be 1D when provided");
63 assert((p->size(0) == 1 || p->size(0) == batch_size_) &&
64 "`TopKTopPSampler` requires `p` shape [1] or [batch_size]");
65 assert((p->dtype() == DataType::kFloat16 ||
66 p->dtype() == DataType::kBFloat16 ||
67 p->dtype() == DataType::kFloat32 ||
68 p->dtype() == DataType::kFloat64) &&
69 "`TopKTopPSampler` requires floating-point `p`");
72 Tensor::Size batch_size_{0};
74 Tensor::Size vocab_size_{0};
Definition generated/include/operator.h:282
Definition top_k_top_p_sampler.h:20
void ValidateK(std::optional< Tensor > k) const
Definition top_k_top_p_sampler.h:47
virtual void operator()(const Tensor logits, std::optional< Tensor > k, std::optional< Tensor > p, Tensor out) const =0
void ValidateP(std::optional< Tensor > p) const
Definition top_k_top_p_sampler.h:58
TopKTopPSampler(const Tensor logits, std::optional< Tensor > k, std::optional< Tensor > p, Tensor out)
Definition top_k_top_p_sampler.h:22
DataType dtype_
Definition top_k_top_p_sampler.h:76
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8