1#ifndef INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLE_INFINILM_H_
2#define INFINI_OPS_BASE_TOP_K_TOP_P_SAMPLE_INFINILM_H_
17 "Migrate to an open-source-aligned operator when available.")]]
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()} {
28 assert(logits.ndim() == 2 &&
29 "`TopKTopPSampleInfinilm` requires 2D `[batch_size, vocab_size]` "
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 "
39 assert(out.dtype() == DataType::kInt32 &&
40 "`TopKTopPSampleInfinilm` requires int32 output");
47 std::optional<Tensor> p, uint64_t seed,
48 uint64_t offset,
Tensor out)
const = 0;
52 if (!k.has_value())
return;
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`");
63 if (!p.has_value())
return;
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`");
76 Tensor::Size batch_size_{0};
78 Tensor::Size vocab_size_{0};
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