InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
top_k_top_p_sampler.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLER_H_
2#define INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLER_H_
3
4#include <cassert>
5#include <optional>
6
7#include "data_type.h"
8#include "operator.h"
9#include "tensor.h"
10
11namespace infini::ops {
12
13// Legacy sampler for 2D `logits` after optional rank and nucleus filtering.
14// Temperature scaling is intentionally handled by callers.
15// The optional `k` and `p` tensors may be shaped as `[1]` or `[batch_size]`.
18class [[deprecated(
19 "Migrate to an open-source-aligned operator when available.")]]
20TopKTopPSampler : public Operator<TopKTopPSampler> {
21 public:
22 TopKTopPSampler(const Tensor logits, std::optional<Tensor> k,
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");
38
39 ValidateK(k);
40 ValidateP(p);
41 }
42
43 virtual void operator()(const Tensor logits, std::optional<Tensor> k,
44 std::optional<Tensor> p, Tensor out) const = 0;
45
46 protected:
47 void ValidateK(std::optional<Tensor> k) const {
48 if (!k.has_value()) return;
49
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`");
56 }
57
58 void ValidateP(std::optional<Tensor> p) const {
59 if (!p.has_value()) return;
60
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`");
70 }
71
72 Tensor::Size batch_size_{0};
73
74 Tensor::Size vocab_size_{0};
75
76 DataType dtype_;
77};
78
79} // namespace infini::ops
80
81#endif // INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLER_H_
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