InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
top_k_top_p_sample_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLE_INFINILM_H_
2#define INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLE_INFINILM_H_
3
4#include <cassert>
5#include <cstdint>
6#include <optional>
7
8#include "data_type.h"
9#include "operator.h"
10#include "tensor.h"
11
12namespace infini::ops {
13
16class [[deprecated(
17 "Migrate to an open-source-aligned operator when available.")]]
18TopKTopPSampleInfinilm : public Operator<TopKTopPSampleInfinilm> {
19 public:
20 TopKTopPSampleInfinilm(const Tensor logits, std::optional<Tensor> k,
21 std::optional<Tensor> p, uint64_t seed,
22 uint64_t offset, Tensor out)
23 : batch_size_{logits.size(0)},
24 vocab_size_{logits.size(1)},
25 dtype_{logits.dtype()} {
26 (void)seed;
27 (void)offset;
28 assert(logits.ndim() == 2 &&
29 "`TopKTopPSampleInfinilm` requires 2D `[batch_size, vocab_size]` "
30 "logits");
31 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
32 dtype_ == DataType::kFloat32 || dtype_ == DataType::kFloat64) &&
33 "`TopKTopPSampleInfinilm` requires floating-point logits");
34 assert(out.ndim() == 1 &&
35 "`TopKTopPSampleInfinilm` requires 1D `[batch_size]` output");
36 assert(out.size(0) == batch_size_ &&
37 "`TopKTopPSampleInfinilm` requires output batch size to match "
38 "logits");
39 assert(out.dtype() == DataType::kInt32 &&
40 "`TopKTopPSampleInfinilm` requires int32 output");
41
42 ValidateK(k);
43 ValidateP(p);
44 }
45
46 virtual void operator()(const Tensor logits, std::optional<Tensor> k,
47 std::optional<Tensor> p, uint64_t seed,
48 uint64_t offset, Tensor out) const = 0;
49
50 protected:
51 void ValidateK(std::optional<Tensor> k) const {
52 if (!k.has_value()) return;
53
54 assert(k->ndim() == 1 &&
55 "`TopKTopPSampleInfinilm` requires `k` to be 1D when provided");
56 assert((k->size(0) == 1 || k->size(0) == batch_size_) &&
57 "`TopKTopPSampleInfinilm` requires `k` shape [1] or [batch_size]");
58 assert((k->dtype() == DataType::kInt32 || k->dtype() == DataType::kInt64) &&
59 "`TopKTopPSampleInfinilm` requires int32 or int64 `k`");
60 }
61
62 void ValidateP(std::optional<Tensor> p) const {
63 if (!p.has_value()) return;
64
65 assert(p->ndim() == 1 &&
66 "`TopKTopPSampleInfinilm` requires `p` to be 1D when provided");
67 assert((p->size(0) == 1 || p->size(0) == batch_size_) &&
68 "`TopKTopPSampleInfinilm` requires `p` shape [1] or [batch_size]");
69 assert((p->dtype() == DataType::kFloat16 ||
70 p->dtype() == DataType::kBFloat16 ||
71 p->dtype() == DataType::kFloat32 ||
72 p->dtype() == DataType::kFloat64) &&
73 "`TopKTopPSampleInfinilm` requires floating-point `p`");
74 }
75
76 Tensor::Size batch_size_{0};
77
78 Tensor::Size vocab_size_{0};
79
80 DataType dtype_;
81};
82
83} // namespace infini::ops
84
85#endif // INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLE_INFINILM_H_
Definition generated/include/operator.h:282
Definition top_k_top_p_sample_infinilm.h:18
TopKTopPSampleInfinilm(const Tensor logits, std::optional< Tensor > k, std::optional< Tensor > p, uint64_t seed, uint64_t offset, Tensor out)
Definition top_k_top_p_sample_infinilm.h:20
virtual void operator()(const Tensor logits, std::optional< Tensor > k, std::optional< Tensor > p, uint64_t seed, uint64_t offset, Tensor out) const =0
void ValidateP(std::optional< Tensor > p) const
Definition top_k_top_p_sample_infinilm.h:62
void ValidateK(std::optional< Tensor > k) const
Definition top_k_top_p_sample_infinilm.h:51
DataType dtype_
Definition top_k_top_p_sample_infinilm.h:80
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8