InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
top_k_top_p_sampling_from_logits.h
Go to the documentation of this file.
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_
3
4#include <cassert>
5#include <cstdint>
6#include <optional>
7#include <string>
8
9#include "data_type.h"
10#include "operator.h"
11#include "tensor.h"
12
13namespace infini::ops {
14
15// Targets the public FlashInfer `top_k_top_p_sampling_from_logits` interface.
16class TopKTopPSamplingFromLogits : public Operator<TopKTopPSamplingFromLogits> {
17 public:
18 TopKTopPSamplingFromLogits(const Tensor logits, const Tensor top_k,
19 const Tensor top_p,
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)
25 : batch_size_{out.size(0)},
26 vocab_size_{logits.size(1)},
27 dtype_{logits.dtype()} {
28 assert(logits.ndim() == 2 &&
29 "`TopKTopPSamplingFromLogits` requires 2D "
30 "`[batch_size, vocab_size]` logits.");
31 assert(IsFloatDtype(dtype_) &&
32 "`TopKTopPSamplingFromLogits` requires floating-point logits.");
33 assert(top_k.ndim() == 1 && top_k.size(0) == batch_size_ &&
34 IsIntegerDtype(top_k.dtype()) &&
35 "`TopKTopPSamplingFromLogits` requires integer `top_k` with shape "
36 "`[batch_size]`.");
37 assert(top_p.ndim() == 1 && top_p.size(0) == batch_size_ &&
38 IsFloatDtype(top_p.dtype()) &&
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`.");
49
50 if (indices.has_value()) {
51 assert(indices->ndim() == 1 && indices->size(0) == batch_size_ &&
52 IsIntegerDtype(indices->dtype()) &&
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.");
58 } else {
59 assert(logits.size(0) == batch_size_ &&
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.");
65 }
66
67 (void)deterministic;
68 (void)check_nan;
69 (void)seed;
70 }
71
72 virtual void operator()(const Tensor logits, const Tensor top_k,
73 const Tensor top_p,
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,
79 Tensor out) const = 0;
80
81 protected:
82 static bool IsFloatDtype(DataType dtype) {
83 return dtype == DataType::kFloat16 || dtype == DataType::kBFloat16 ||
84 dtype == DataType::kFloat32 || dtype == DataType::kFloat64;
85 }
86
87 static bool IsIntegerDtype(DataType dtype) {
88 return dtype == DataType::kInt32 || dtype == DataType::kInt64;
89 }
90
91 Tensor::Size batch_size_{0};
92
93 Tensor::Size vocab_size_{0};
94
95 DataType dtype_;
96};
97
98template <>
100 detail::CacheKey operator()(const Config& config, const Tensor logits,
101 const Tensor top_k, const Tensor top_p,
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> /*seed*/,
106 const std::optional<int64_t> /*offset*/,
107 Tensor out) const {
108 return detail::CacheKey::Build(config.implementation_index(), logits, top_k,
109 top_p, indices, filter_apply_order,
110 deterministic, check_nan, out);
111 }
112};
113
114} // namespace infini::ops
115
116#endif // INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLING_FROM_LOGITS_H_
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