InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
random_sample_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_RANDOM_SAMPLE_INFINILM_H_
2#define INFINI_OPS_BASE_RANDOM_SAMPLE_INFINILM_H_
3
4#include <cassert>
5#include <cstdint>
6#include <limits>
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.")]]
18RandomSampleInfinilm : public Operator<RandomSampleInfinilm> {
19 public:
20 RandomSampleInfinilm(const Tensor logits, float random_val, float topp,
21 int64_t topk, float temperature, Tensor out)
22 : dtype_{logits.dtype()},
23 out_dtype_{out.dtype()},
24 n_{logits.size(0)},
25 logits_stride_{logits.stride(0)},
26 topp_{topp},
27 topk_{topk},
28 temperature_{temperature} {
29 assert(logits.ndim() == 1 && "`RandomSampleInfinilm` requires 1D logits");
30 assert(n_ > 0 && "`RandomSampleInfinilm` requires non-empty logits");
31 assert(logits.stride(0) == 1 &&
32 "`RandomSampleInfinilm` requires contiguous logits");
33 assert(out.numel() == 1 && "`RandomSampleInfinilm` requires scalar output");
34 assert(IsFloatDtype(dtype_) &&
35 "`RandomSampleInfinilm` requires floating-point logits");
36 assert(IsIntDtype(out_dtype_) &&
37 "`RandomSampleInfinilm` requires integer output");
38 assert(topk > 0 && "`RandomSampleInfinilm` requires `topk > 0`");
39 assert(topk <= std::numeric_limits<int>::max() &&
40 "`RandomSampleInfinilm` requires `topk` to fit in int");
41 }
42
43 virtual void operator()(const Tensor logits, float random_val, float topp,
44 int64_t topk, float temperature,
45 Tensor out) const = 0;
46
47 protected:
48 static bool IsFloatDtype(DataType dtype) {
49 return dtype == DataType::kFloat16 || dtype == DataType::kBFloat16 ||
50 dtype == DataType::kFloat32 || dtype == DataType::kFloat64;
51 }
52
53 static bool IsIntDtype(DataType dtype) {
54 return dtype == DataType::kInt8 || dtype == DataType::kInt16 ||
55 dtype == DataType::kInt32 || dtype == DataType::kInt64 ||
56 dtype == DataType::kUInt8 || dtype == DataType::kUInt16 ||
57 dtype == DataType::kUInt32 || dtype == DataType::kUInt64;
58 }
59
60 DataType dtype_;
61
62 DataType out_dtype_;
63
64 Tensor::Size n_{0};
65
66 Tensor::Stride logits_stride_{1};
67
68 float topp_{0.0f};
69
70 int64_t topk_{1};
71
72 float temperature_{1.0f};
73};
74
75template <>
77 detail::CacheKey operator()(const Config& config, const Tensor logits,
78 float /*random_val*/, float topp, int64_t topk,
79 float temperature, Tensor out) const {
80 return detail::CacheKey::Build(config.implementation_index(), logits, topp,
81 topk, temperature, out);
82 }
83};
84
85} // namespace infini::ops
86
87#endif
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 random_sample_infinilm.h:18
DataType out_dtype_
Definition random_sample_infinilm.h:62
virtual void operator()(const Tensor logits, float random_val, float topp, int64_t topk, float temperature, Tensor out) const =0
RandomSampleInfinilm(const Tensor logits, float random_val, float topp, int64_t topk, float temperature, Tensor out)
Definition random_sample_infinilm.h:20
static bool IsFloatDtype(DataType dtype)
Definition random_sample_infinilm.h:48
DataType dtype_
Definition random_sample_infinilm.h:60
static bool IsIntDtype(DataType dtype)
Definition random_sample_infinilm.h:53
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8
detail::CacheKey operator()(const Config &config, const Tensor logits, float, float topp, int64_t topk, float temperature, Tensor out) const
Definition random_sample_infinilm.h:77
Definition generated/include/operator.h:225