Some changes to Sampling Op (#14218)

### Description
<!-- Describe your changes. -->
1. add an optional input to pass in seed
2. two UTs. one for top_p=0.5, another for top_p=0.01(create greedy
search result, in convert_generation.py)
3. fix a bug in cpu kernel

### Motivation and Context
<!-- - Why is this change required? What problem does it solve?
- If it fixes an open issue, please link to the issue here. -->

Co-authored-by: Ubuntu <wy@v100-2.0cdb2e52twzevn1i4fi45bylyg.jx.internal.cloudapp.net>
This commit is contained in:
Ye Wang 2023-01-12 14:15:26 -08:00 committed by GitHub
parent 3898b22a1a
commit c9a53c9255
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
14 changed files with 236 additions and 23 deletions

View file

@ -3945,7 +3945,7 @@ This version of the operator has been available since version 1 of the 'com.micr
<dd>Size of the vocabulary. If not provided, it will be inferred from the decoder subgraph's output shape</dd>
</dl>
#### Inputs (2 - 8)
#### Inputs (2 - 9)
<dl>
<dt><tt>input_ids</tt> : I</dt>
@ -3964,6 +3964,8 @@ This version of the operator has been available since version 1 of the 'com.micr
<dd>Custom attention mask. Shape is (batch_size, sequence_length)</dd>
<dt><tt>presence_mask</tt> (optional) : I</dt>
<dd>Presence penalty mask. Shape is (batch_size, vocab_size)</dd>
<dt><tt>seed</tt> (optional) : I</dt>
<dd>Seed for random number generator. Shape is (1)</dd>
</dl>
#### Outputs (1 - 2)

View file

@ -451,7 +451,7 @@ Do not modify directly.*
|QuickGelu|*in* X:**T**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|Range|*in* start:**T**<br> *in* limit:**T**<br> *in* delta:**T**<br> *out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64)|
|SampleOp|*in* X:**T**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|Sampling|*in* input_ids:**I**<br> *in* max_length:**I**<br> *in* min_length:**I**<br> *in* repetition_penalty:**T**<br> *in* vocab_mask:**I**<br> *in* prefix_vocab_mask:**I**<br> *in* attention_mask:**I**<br> *in* presence_mask:**I**<br> *out* sequences:**I**<br> *out* filtered_logits:**T**|1+|**T** = tensor(float)|
|Sampling|*in* input_ids:**I**<br> *in* max_length:**I**<br> *in* min_length:**I**<br> *in* repetition_penalty:**T**<br> *in* vocab_mask:**I**<br> *in* prefix_vocab_mask:**I**<br> *in* attention_mask:**I**<br> *in* presence_mask:**I**<br> *in* seed:**I**<br> *out* sequences:**I**<br> *out* filtered_logits:**T**|1+|**T** = tensor(float)|
|SkipLayerNormalization|*in* input:**T**<br> *in* skip:**T**<br> *in* gamma:**T**<br> *in* beta:**T**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* mean:**U**<br> *out* inv_std_var:**U**<br> *out* input_skip_bias_sum:**T**|1+|**T** = tensor(double), tensor(float)|
|SparseToDenseMatMul|*in* A:**T**<br> *in* B:**T1**<br> *out* Y:**T1**|1+|**T** = sparse_tensor(double), sparse_tensor(float), sparse_tensor(int32), sparse_tensor(int64), sparse_tensor(uint32), sparse_tensor(uint64)<br/> **T1** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|Tokenizer|*in* X:**T**<br> *out* Y:**T**|1+|**T** = tensor(string)|
@ -814,7 +814,7 @@ Do not modify directly.*
|RemovePadding|*in* input:**T**<br> *in* sequence_token_count:**M**<br> *out* output:**T**<br> *out* token_offset:**M**<br> *out* cumulated_seq_len:**M**<br> *out* max_seq_len:**M**|1+|**T** = tensor(float), tensor(float16)|
|RestorePadding|*in* input:**T**<br> *in* token_offset:**M**<br> *out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
|Rfft|*in* X:**T**<br> *out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)|
|Sampling|*in* input_ids:**I**<br> *in* max_length:**I**<br> *in* min_length:**I**<br> *in* repetition_penalty:**T**<br> *in* vocab_mask:**I**<br> *in* prefix_vocab_mask:**I**<br> *in* attention_mask:**I**<br> *in* presence_mask:**I**<br> *out* sequences:**I**<br> *out* filtered_logits:**T**|1+|**T** = tensor(float), tensor(float16)|
|Sampling|*in* input_ids:**I**<br> *in* max_length:**I**<br> *in* min_length:**I**<br> *in* repetition_penalty:**T**<br> *in* vocab_mask:**I**<br> *in* prefix_vocab_mask:**I**<br> *in* attention_mask:**I**<br> *in* presence_mask:**I**<br> *in* seed:**I**<br> *out* sequences:**I**<br> *out* filtered_logits:**T**|1+|**T** = tensor(float), tensor(float16)|
|SkipLayerNormalization|*in* input:**T**<br> *in* skip:**T**<br> *in* gamma:**T**<br> *in* beta:**T**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* mean:**U**<br> *out* inv_std_var:**U**<br> *out* input_skip_bias_sum:**T**|1+|**T** = tensor(float), tensor(float16)|
|SkipSimplifiedLayerNormalization|*in* input:**T**<br> *in* skip:**T**<br> *in* gamma:**T**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* mean:**U**<br> *out* inv_std_var:**U**<br> *out* input_skip_bias_sum:**T**|1+|**T** = tensor(float), tensor(float16)|
|TransposeMatMul|*in* A:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)|

View file

@ -105,6 +105,12 @@ void CpuTensorConsoleDumper::Print(const char* name, const MLFloat16* tensor, in
DumpCpuTensor<MLFloat16>(name, tensor, dim0, dim1);
}
void CpuTensorConsoleDumper::Print(const char* name, const size_t* tensor, int dim0, int dim1) const {
if (!is_enabled_)
return;
DumpCpuTensor<size_t>(name, tensor, dim0, dim1);
}
void CpuTensorConsoleDumper::Print(const char* name, const int64_t* tensor, int dim0, int dim1) const {
if (!is_enabled_)
return;
@ -180,6 +186,9 @@ void CpuTensorConsoleDumper::Print(const char*, const float*, int, int) const {
void CpuTensorConsoleDumper::Print(const char*, const MLFloat16*, int, int) const {
}
void CpuTensorConsoleDumper::Print(const char*, const size_t*, int, int) const {
}
void CpuTensorConsoleDumper::Print(const char*, const int64_t*, int, int) const {
}

View file

@ -16,6 +16,7 @@ class CpuTensorConsoleDumper : public IConsoleDumper {
virtual ~CpuTensorConsoleDumper() {}
void Print(const char* name, const float* tensor, int dim0, int dim1) const override;
void Print(const char* name, const MLFloat16* tensor, int dim0, int dim1) const override;
void Print(const char* name, const size_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const int64_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const int32_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const float* tensor, int dim0, int dim1, int dim2) const override;

View file

@ -172,6 +172,7 @@ class IConsoleDumper {
bool IsEnabled() const { return is_enabled_; }
virtual void Print(const char* name, const float* tensor, int dim0, int dim1) const = 0;
virtual void Print(const char* name, const MLFloat16* tensor, int dim0, int dim1) const = 0;
virtual void Print(const char* name, const size_t* tensor, int dim0, int dim1) const = 0;
virtual void Print(const char* name, const int64_t* tensor, int dim0, int dim1) const = 0;
virtual void Print(const char* name, const int32_t* tensor, int dim0, int dim1) const = 0;
virtual void Print(const char* name, const float* tensor, int dim0, int dim1, int dim2) const = 0;

View file

@ -10,9 +10,10 @@ template <typename T>
void filter_scores(std::vector<size_t>& sorted_indice,
gsl::span<T>& next_token_score,
const transformers::IGenerationParameters* parameters,
size_t index) {
size_t real_index = sorted_indice[index];
next_token_score[real_index] = (T)parameters->filter_value;
size_t chunk_offset,
size_t offset) {
size_t real_index = sorted_indice[chunk_offset + offset];
next_token_score[chunk_offset + real_index] = (T)parameters->filter_value;
}
template <typename T>
@ -23,12 +24,12 @@ void cumulate_and_filter_custom(gsl::span<T>& next_token_scores,
for (size_t i = 0; i < static_cast<size_t>(parameters->batch_size); i++) {
size_t offset = i * parameters->vocab_size;
if (cumulative_probs[offset] > parameters->top_p) {
filter_scores(sorted_indices, next_token_scores, parameters, 1 + offset);
filter_scores(sorted_indices, next_token_scores, parameters, offset, 1);
}
for (size_t j = 1; j < static_cast<size_t>(parameters->vocab_size) - 1; j++) {
cumulative_probs[j + offset] += cumulative_probs[j + offset - 1];
if (cumulative_probs[j + offset] > parameters->top_p) {
filter_scores(sorted_indices, next_token_scores, parameters, j + offset + 1);
filter_scores(sorted_indices, next_token_scores, parameters, offset, j + 1);
}
}
}
@ -42,12 +43,12 @@ void cumulate_and_filter(gsl::span<T>& next_token_scores,
for (size_t i = 0; i < static_cast<size_t>(parameters->batch_size); i++) {
size_t offset = i * parameters->vocab_size;
if (cumulative_probs[offset] <= 1 - parameters->top_p) {
filter_scores(sorted_indices, next_token_scores, parameters, offset);
filter_scores(sorted_indices, next_token_scores, parameters, offset, 0);
}
for (size_t j = 1; j < static_cast<size_t>(parameters->vocab_size) - static_cast<size_t>(parameters->min_tokens_to_keep); j++) {
cumulative_probs[j + offset] += cumulative_probs[j + offset - 1];
if (cumulative_probs[j + offset] <= 1 - parameters->top_p) {
filter_scores(sorted_indices, next_token_scores, parameters, j + offset);
filter_scores(sorted_indices, next_token_scores, parameters, offset, j);
}
}
}
@ -78,10 +79,11 @@ Status Sample(AllocatorPtr& allocator,
for (size_t i = 0; i < static_cast<size_t>(parameters->batch_size); i++) {
auto indices_begin = sorted_indices.begin() + i * parameters->vocab_size;
auto indices_end = sorted_indices.begin() + (i + 1) * parameters->vocab_size;
gsl::span<T> next_token_score = next_token_scores.subspan(i * parameters->vocab_size, parameters->vocab_size);
std::iota(indices_begin, indices_end, 0);
std::sort(indices_begin, indices_end,
[&next_token_scores, &predicator](size_t i1, size_t i2) {
return !predicator(next_token_scores[i1], next_token_scores[i2]);
[&next_token_score, &predicator](size_t i1, size_t i2) {
return predicator(next_token_score[i1], next_token_score[i2]);
});
std::sort(sorted_scores.begin() + i * parameters->vocab_size,
@ -89,6 +91,11 @@ Status Sample(AllocatorPtr& allocator,
predicator);
}
#ifdef DEBUG_GENERATION
dumper->Print("sorted_scores", sorted_scores.data(), parameters->batch_size, parameters->vocab_size);
dumper->Print("sorted_indices", sorted_indices.data(), parameters->batch_size, parameters->vocab_size);
#endif
gsl::span<T>& cumulative_probs = sampling_state->cumulative_probs;
ORT_RETURN_IF_ERROR(SoftmaxCPU<T>(parameters->batch_size,
@ -104,13 +111,10 @@ Status Sample(AllocatorPtr& allocator,
cumulate_and_filter(next_token_scores, cumulative_probs, parameters, sorted_indices);
}
gsl::span<T>& next_token_probs = sampling_state->h_softmaxed_score;
ORT_RETURN_IF_ERROR(SoftmaxCPU<T>(parameters->batch_size,
parameters->vocab_size,
next_token_scores.data(),
next_token_probs.data(),
false,
thread_pool));
#ifdef DEBUG_GENERATION
dumper->Print("cumulative_probs after filtering", cumulative_probs.data(), parameters->batch_size, parameters->vocab_size);
dumper->Print("next_token_scores after filtering", next_token_scores.data(), parameters->batch_size, parameters->vocab_size);
#endif
// torch.multinomial()
int64_t next_token_probs_dims[] = {static_cast<int64_t>(parameters->batch_size), parameters->vocab_size};
@ -119,7 +123,7 @@ Status Sample(AllocatorPtr& allocator,
OrtValue next_token_probs_value;
Tensor::InitOrtValue(element_type,
next_token_probs_shape,
next_token_probs.data(),
next_token_scores.data(),
allocator->Info(),
next_token_probs_value);
const Tensor& input = next_token_probs_value.Get<Tensor>();

View file

@ -21,6 +21,14 @@ void SamplingParameters::ParseFromAttributes(const OpKernelInfo& info) {
vocab_size = static_cast<int>(info.GetAttrOrDefault<int64_t>("vocab_size", -1));
}
void SamplingParameters::ParseFromInputs(OpKernelContext* context) {
this->GreedySearchParameters::ParseFromInputs(context);
auto* seed_tensor = context->Input<Tensor>(8);
seed = seed_tensor ? static_cast<int>(*seed_tensor->Data<int32_t>()) : 0;
ORT_ENFORCE(seed >= 0, "Seed must be >= 0");
}
} // namespace transformers
} // namespace contrib
} // namespace onnxruntime

View file

@ -12,6 +12,8 @@ namespace transformers {
struct SamplingParameters : public GreedySearchParameters {
void ParseFromAttributes(const OpKernelInfo& info);
void ParseFromInputs(OpKernelContext* context);
};
} // namespace transformers

View file

@ -145,6 +145,11 @@ void CudaTensorConsoleDumper::Print(const char* name, const MLFloat16* tensor, i
DumpGpuTensor<MLFloat16>(name, tensor, dim0, dim1, true);
}
void CudaTensorConsoleDumper::Print(const char* name, const size_t* tensor, int dim0, int dim1) const {
if (is_enabled_)
DumpGpuTensor<size_t>(name, tensor, dim0, dim1, true);
}
void CudaTensorConsoleDumper::Print(const char* name, const int64_t* tensor, int dim0, int dim1) const {
if (is_enabled_)
DumpGpuTensor<int64_t>(name, tensor, dim0, dim1, true);
@ -212,6 +217,9 @@ void CudaTensorConsoleDumper::Print(const char*, const float*, int, int) const {
void CudaTensorConsoleDumper::Print(const char*, const MLFloat16*, int, int) const {
}
void CudaTensorConsoleDumper::Print(const char*, const size_t*, int, int) const {
}
void CudaTensorConsoleDumper::Print(const char*, const int64_t*, int, int) const {
}

View file

@ -19,6 +19,7 @@ class CudaTensorConsoleDumper : public onnxruntime::contrib::transformers::ICons
virtual ~CudaTensorConsoleDumper() {}
void Print(const char* name, const float* tensor, int dim0, int dim1) const override;
void Print(const char* name, const MLFloat16* tensor, int dim0, int dim1) const override;
void Print(const char* name, const size_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const int64_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const int32_t* tensor, int dim0, int dim1) const override;
void Print(const char* name, const float* tensor, int dim0, int dim1, int dim2) const override;

View file

@ -1155,6 +1155,7 @@ ONNX_MS_OPERATOR_SET_SCHEMA(Sampling, 1,
.Input(5, "prefix_vocab_mask", "Mask of vocabulary for first step. Words that masked with 0 are not allowed to be generated, and 1 is allowed. Shape is (batch_size, vocab_size)", "I", OpSchema::Optional)
.Input(6, "attention_mask", "Custom attention mask. Shape is (batch_size, sequence_length)", "I", OpSchema::Optional)
.Input(7, "presence_mask", "Presence penalty mask. Shape is (batch_size, vocab_size)", "I", OpSchema::Optional)
.Input(8, "seed", "Seed for random number generator. Shape is (1)", "I", OpSchema::Optional)
.Output(0, "sequences", "Word IDs of generated sequences. Shape is (batch_size, max_sequence_length)", "I")
.Output(1, "filtered_logits", "Filtered logits as input to the mutinomial function for debug purpose. Shape is (batch_size, vocab_size)", "T", OpSchema::Optional)
.TypeConstraint("T", {"tensor(float)"}, "Constrain input and output types to float tensors.")

View file

@ -273,6 +273,14 @@ def parse_arguments(argv: Optional[List[str]] = None) -> argparse.Namespace:
)
model_group.set_defaults(presence_mask=False)
model_group.add_argument(
"--seed",
required=False,
action="store_true",
help="Random seed for sampling op",
)
model_group.set_defaults(seed=False)
beam_parameters_group = parser.add_argument_group(
"Beam search parameters not stored in the output model, for testing parity and performance"
)
@ -1531,6 +1539,11 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
if is_sampling and args.custom and args.presence_mask:
inputs.append("presence_mask")
else:
inputs.append("")
if is_sampling and args.seed:
inputs.append("seed")
outputs = ["sequences"]
if args.output_sequences_scores:
@ -1709,6 +1722,10 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
)
graph_inputs.append(presence_mask)
if is_sampling and args.seed:
seed = onnx.helper.make_tensor_value_info("seed", TensorProto.INT32, [1])
graph_inputs.append(seed)
# graph outputs
sequences = None
if is_beamsearch:
@ -2278,9 +2295,13 @@ def main(argv: Optional[List[str]] = None, sentences: Optional[List[str]] = None
if args.model_type == "gpt2" and is_greedy:
if args.top_p > 0.0 and args.top_p < 1.0:
convert_generation_model(args, GenerationType.SAMPLING)
logger.info("The test for gpt2_sampling onnx model is not implemented yet")
return
convert_generation_model(args, GenerationType.GREEDYSEARCH)
logger.info(
"The test for gpt2_sampling onnx model is limited to non-custom model with small top_p(e.g <=0.01) value. The result should be the same as gpt2 greedy search."
)
if args.top_p > 0.01 or args.custom or args.seed:
return
else:
convert_generation_model(args, GenerationType.GREEDYSEARCH)
else:
convert_generation_model(args)

View file

@ -0,0 +1,155 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include <memory>
#include <vector>
#include "gtest/gtest.h"
#include "core/common/gsl.h"
#include "core/session/onnxruntime_cxx_api.h"
#include "test/common/cuda_op_test_utils.h"
extern std::unique_ptr<Ort::Env> ort_env;
namespace onnxruntime {
namespace test {
#if defined(__linux__) && !defined(__ANDROID__)
#ifdef USE_CUDA
TEST(SamplingTest, Gpt2Sampling_CUDA) {
std::vector<int32_t> input_ids{
0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620,
41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572,
0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328};
std::vector<int32_t> max_length{15};
std::vector<int32_t> min_length{1};
std::vector<float> repetition_penalty{1.0f};
std::vector<int32_t> expected_output{
0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620, 125, 543, 668,
41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572, 776, 213, 697,
0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328, 450};
const int64_t batch_size = 3;
const int64_t sequence_length = 12;
std::vector<int64_t> input_ids_shape{batch_size, sequence_length};
std::vector<int64_t> parameter_shape{1};
std::vector<int64_t> expected_output_shape{input_ids_shape[0], max_length[0]};
Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
auto input_ids_tensor = Ort::Value::CreateTensor(
info, input_ids.data(), input_ids.size(), input_ids_shape.data(), input_ids_shape.size());
auto max_length_tensor = Ort::Value::CreateTensor(
info, max_length.data(), max_length.size(), parameter_shape.data(), parameter_shape.size());
auto min_length_tensor = Ort::Value::CreateTensor(
info, min_length.data(), min_length.size(), parameter_shape.data(), parameter_shape.size());
auto repetition_penalty_tensor = Ort::Value::CreateTensor(
info, repetition_penalty.data(), repetition_penalty.size(), parameter_shape.data(), parameter_shape.size());
std::vector<Ort::Value> ort_inputs;
ort_inputs.push_back(std::move(input_ids_tensor));
ort_inputs.push_back(std::move(max_length_tensor));
ort_inputs.push_back(std::move(min_length_tensor));
ort_inputs.push_back(std::move(repetition_penalty_tensor));
const char* input_names[] = {"input_ids", "max_length", "min_length", "repetition_penalty"};
const char* const output_names[] = {"sequences"};
Ort::SessionOptions session_options;
constexpr int min_cuda_architecture = 530;
if (HasCudaEnvironment(min_cuda_architecture)) {
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0));
Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_sampling.onnx"), session_options);
auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(),
output_names, 1);
ASSERT_EQ(ort_outputs.size(), 1U);
const auto& sequences = ort_outputs[0];
ASSERT_TRUE(sequences.IsTensor());
auto result_ts = sequences.GetTensorTypeAndShapeInfo();
ASSERT_EQ(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, result_ts.GetElementType());
ASSERT_EQ(expected_output_shape, result_ts.GetShape());
const auto* result_vals = sequences.GetTensorData<int32_t>();
auto result_span = gsl::make_span(result_vals, expected_output.size());
ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end()));
}
}
#endif
TEST(SamplingTest, Gpt2Sampling_CPU) {
std::vector<int32_t> input_ids{
0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620,
41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572,
0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328};
std::vector<int32_t> max_length{15};
std::vector<int32_t> min_length{1};
std::vector<float> repetition_penalty{1.0f};
std::vector<int32_t> expected_output{
0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620, 125, 669, 28,
41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572, 475, 944, 527,
0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328, 210};
const int64_t batch_size = 3;
const int64_t sequence_length = 12;
std::vector<int64_t> input_ids_shape{batch_size, sequence_length};
std::vector<int64_t> parameter_shape{1};
std::vector<int64_t> expected_output_shape{input_ids_shape[0], max_length[0]};
Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
auto input_ids_tensor = Ort::Value::CreateTensor(
info, input_ids.data(), input_ids.size(), input_ids_shape.data(), input_ids_shape.size());
auto max_length_tensor = Ort::Value::CreateTensor(
info, max_length.data(), max_length.size(), parameter_shape.data(), parameter_shape.size());
auto min_length_tensor = Ort::Value::CreateTensor(
info, min_length.data(), min_length.size(), parameter_shape.data(), parameter_shape.size());
auto repetition_penalty_tensor = Ort::Value::CreateTensor(
info, repetition_penalty.data(), repetition_penalty.size(), parameter_shape.data(), parameter_shape.size());
std::vector<Ort::Value> ort_inputs;
ort_inputs.push_back(std::move(input_ids_tensor));
ort_inputs.push_back(std::move(max_length_tensor));
ort_inputs.push_back(std::move(min_length_tensor));
ort_inputs.push_back(std::move(repetition_penalty_tensor));
const char* input_names[] = {"input_ids", "max_length", "min_length", "repetition_penalty"};
const char* const output_names[] = {"sequences"};
Ort::SessionOptions session_options;
Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_sampling.onnx"), session_options);
auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(),
output_names, 1);
ASSERT_EQ(ort_outputs.size(), 1U);
const auto& sequences = ort_outputs[0];
ASSERT_TRUE(sequences.IsTensor());
auto result_ts = sequences.GetTensorTypeAndShapeInfo();
ASSERT_EQ(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, result_ts.GetElementType());
ASSERT_EQ(expected_output_shape, result_ts.GetShape());
const auto* result_vals = sequences.GetTensorData<int32_t>();
auto result_span = gsl::make_span(result_vals, expected_output.size());
ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end()));
}
#endif
} // namespace test
} // namespace onnxruntime

Binary file not shown.