1#ifndef INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLING_FROM_LOGITS_H_
2#define INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLING_FROM_LOGITS_H_
20 const std::optional<Tensor> indices,
21 const std::string filter_apply_order,
22 const bool deterministic,
const bool check_nan,
23 const std::optional<int64_t> seed,
24 const std::optional<int64_t> offset,
Tensor out)
28 assert(logits.ndim() == 2 &&
29 "`TopKTopPSamplingFromLogits` requires 2D "
30 "`[batch_size, vocab_size]` logits.");
32 "`TopKTopPSamplingFromLogits` requires floating-point logits.");
33 assert(top_k.ndim() == 1 && top_k.size(0) ==
batch_size_ &&
35 "`TopKTopPSamplingFromLogits` requires integer `top_k` with shape "
37 assert(top_p.ndim() == 1 && top_p.size(0) ==
batch_size_ &&
39 "`TopKTopPSamplingFromLogits` requires floating-point `top_p` with "
40 "shape `[batch_size]`.");
41 assert(out.ndim() == 1 &&
42 "`TopKTopPSamplingFromLogits` requires 1D output.");
43 assert((filter_apply_order ==
"top_k_first" ||
44 filter_apply_order ==
"joint") &&
45 "`TopKTopPSamplingFromLogits` requires `filter_apply_order` to be "
46 "`top_k_first` or `joint`.");
47 assert((!offset.has_value() || *offset >= 0) &&
48 "`TopKTopPSamplingFromLogits` requires a nonnegative `offset`.");
50 if (indices.has_value()) {
51 assert(indices->ndim() == 1 && indices->size(0) ==
batch_size_ &&
53 "`TopKTopPSamplingFromLogits` requires integer `indices` with "
54 "shape `[batch_size]`.");
55 assert(out.dtype() == indices->dtype() &&
56 "`TopKTopPSamplingFromLogits` requires output and `indices` to "
57 "have the same dtype.");
60 "`TopKTopPSamplingFromLogits` requires output batch size to "
61 "match logits when `indices` is absent.");
62 assert(out.dtype() == DataType::kInt32 &&
63 "`TopKTopPSamplingFromLogits` requires int32 output when "
64 "`indices` is absent.");
74 const std::optional<Tensor> indices,
75 const std::string filter_apply_order,
76 const bool deterministic,
const bool check_nan,
77 const std::optional<int64_t> seed,
78 const std::optional<int64_t> offset,
83 return dtype == DataType::kFloat16 || dtype == DataType::kBFloat16 ||
84 dtype == DataType::kFloat32 || dtype == DataType::kFloat64;
88 return dtype == DataType::kInt32 || dtype == DataType::kInt64;
102 const std::optional<Tensor> indices,
103 const std::string filter_apply_order,
104 const bool deterministic,
const bool check_nan,
105 const std::optional<int64_t> ,
106 const std::optional<int64_t> ,
109 top_p, indices, filter_apply_order,
110 deterministic, check_nan, out);
Definition generated/include/config.h:12
std::size_t implementation_index() const
Definition generated/include/config.h:20
Definition generated/include/operator.h:282
Definition top_k_top_p_sampling_from_logits.h:16
Tensor::Size vocab_size_
Definition top_k_top_p_sampling_from_logits.h:93
static bool IsFloatDtype(DataType dtype)
Definition top_k_top_p_sampling_from_logits.h:82
DataType dtype_
Definition top_k_top_p_sampling_from_logits.h:95
virtual void operator()(const Tensor logits, const Tensor top_k, const Tensor top_p, const std::optional< Tensor > indices, const std::string filter_apply_order, const bool deterministic, const bool check_nan, const std::optional< int64_t > seed, const std::optional< int64_t > offset, Tensor out) const =0
TopKTopPSamplingFromLogits(const Tensor logits, const Tensor top_k, const Tensor top_p, const std::optional< Tensor > indices, const std::string filter_apply_order, const bool deterministic, const bool check_nan, const std::optional< int64_t > seed, const std::optional< int64_t > offset, Tensor out)
Definition top_k_top_p_sampling_from_logits.h:18
Tensor::Size batch_size_
Definition top_k_top_p_sampling_from_logits.h:91
static bool IsIntegerDtype(DataType dtype)
Definition top_k_top_p_sampling_from_logits.h:87
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8
detail::CacheKey operator()(const Config &config, const Tensor logits, const Tensor top_k, const Tensor top_p, const std::optional< Tensor > indices, const std::string filter_apply_order, const bool deterministic, const bool check_nan, const std::optional< int64_t >, const std::optional< int64_t >, Tensor out) const
Definition top_k_top_p_sampling_from_logits.h:100
Definition generated/include/operator.h:225