1#ifndef INFINI_OPS_BASE_RANDOM_SAMPLE_INFINILM_H_
2#define INFINI_OPS_BASE_RANDOM_SAMPLE_INFINILM_H_
17 "Migrate to an open-source-aligned operator when available.")]]
21 int64_t topk,
float temperature,
Tensor out)
22 : dtype_{logits.dtype()},
23 out_dtype_{out.dtype()},
25 logits_stride_{logits.stride(0)},
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");
44 int64_t topk,
float temperature,
49 return dtype == DataType::kFloat16 || dtype == DataType::kBFloat16 ||
50 dtype == DataType::kFloat32 || dtype == DataType::kFloat64;
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;
66 Tensor::Stride logits_stride_{1};
72 float temperature_{1.0f};
78 float ,
float topp, int64_t topk,
79 float temperature,
Tensor out)
const {
81 topk, temperature, 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 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