Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions docs/OperatorKernels.md
Original file line number Diff line number Diff line change
Expand Up @@ -292,7 +292,8 @@ The **OpSet Version** column uses the following notation:
|Mul|*in* A:**T**<br> *in* B:**T**<br> *out* C:**T**|14+|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)|
|||13|**T** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|||[7, 12]|**T** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|Multinomial|*in* input:**T1**<br> *out* output:**T2**|7+|**T1** = tensor(float)<br/> **T2** = tensor(int32), tensor(int64)|
|Multinomial|*in* input:**T1**<br> *out* output:**T2**|22+|**T1** = tensor(float)<br/> **T2** = tensor(int32), tensor(int64)|
|||[7, 21]|**T1** = tensor(float)<br/> **T2** = tensor(int32), tensor(int64)|
|Neg|*in* X:**T**<br> *out* Y:**T**|13+|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64), tensor(int8)|
|||[6, 12]|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64), tensor(int8)|
|NonZero|*in* X:**T**<br> *out* Y:**tensor(int64)**|13+|**T** = tensor(bool), tensor(float), tensor(int32), tensor(int64), tensor(uint8)|
Expand Down Expand Up @@ -337,10 +338,14 @@ The **OpSet Version** column uses the following notation:
|RNN|*in* X:**T**<br> *in* W:**T**<br> *in* R:**T**<br> *in* B:**T**<br> *in* sequence_lens:**T1**<br> *in* initial_h:**T**<br> *out* Y:**T**<br> *out* Y_h:**T**|22+|**T** = tensor(float)<br/> **T1** = tensor(int32)|
|||[14, 21]|**T** = tensor(float)<br/> **T1** = tensor(int32)|
|||[7, 13]|**T** = tensor(float)<br/> **T1** = tensor(int32)|
|RandomNormal|*out* output:**T**|1+|**T** = tensor(double), tensor(float)|
|RandomNormalLike|*in* input:**T1**<br> *out* output:**T2**|1+|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|RandomUniform|*out* output:**T**|1+|**T** = tensor(double), tensor(float)|
|RandomUniformLike|*in* input:**T1**<br> *out* output:**T2**|1+|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|RandomNormal|*out* output:**T**|22+|**T** = tensor(double), tensor(float)|
|||[1, 21]|**T** = tensor(double), tensor(float)|
|RandomNormalLike|*in* input:**T1**<br> *out* output:**T2**|22+|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|||[1, 21]|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|RandomUniform|*out* output:**T**|22+|**T** = tensor(double), tensor(float)|
|||[1, 21]|**T** = tensor(double), tensor(float)|
|RandomUniformLike|*in* input:**T1**<br> *out* output:**T2**|22+|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|||[1, 21]|**T1** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)<br/> **T2** = tensor(double), tensor(float)|
|Range|*in* start:**T**<br> *in* limit:**T**<br> *in* delta:**T**<br> *out* output:**T**|27+|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64)|
|||[11, 26]|**T** = tensor(double), tensor(float), tensor(int16), tensor(int32), tensor(int64)|
|Reciprocal|*in* X:**T**<br> *out* Y:**T**|13+|**T** = tensor(double), tensor(float)|
Expand Down
35 changes: 25 additions & 10 deletions onnxruntime/core/providers/cpu/cpu_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -87,11 +87,11 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, float, Tanh);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, double, Tanh);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 8, PRelu);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomNormal);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomUniform);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomNormalLike);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomUniformLike);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, Multinomial);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21, RandomNormal);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21, RandomUniform);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21, RandomNormalLike);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21, RandomUniformLike);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 21, Multinomial);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, float, Abs);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, double, Abs);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, int8_t, Abs);
Expand Down Expand Up @@ -1366,6 +1366,11 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, Th
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, AveragePool);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, float, LpNormalization);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, double, LpNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, Multinomial);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomNormal);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomNormalLike);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomUniform);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomUniformLike);

#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, MLFloat16, Conv);
Expand Down Expand Up @@ -1596,11 +1601,16 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12,
double, Tanh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 8, PRelu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomNormal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomUniform)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomNormalLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, RandomUniformLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, Multinomial)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21,
RandomNormal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21,
RandomUniform)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21,
RandomNormalLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 21,
RandomUniformLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 21,
Multinomial)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 12,
float, Add)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 12,
Expand Down Expand Up @@ -3495,6 +3505,11 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, LpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, MaxUnpool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, Multinomial)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomNormal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomNormalLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomUniform)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, RandomUniformLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, Softplus)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, float, Round)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 22, double, Round)>,
Expand Down
58 changes: 56 additions & 2 deletions onnxruntime/core/providers/cpu/generator/random.cc
Original file line number Diff line number Diff line change
Expand Up @@ -82,44 +82,98 @@ using EnabledRandomNormalComputeOutputTypes =
EnabledRandomNormalOutputTypes,
EnabledRandomNormalLikeOutputTypes>;

ONNX_CPU_OPERATOR_KERNEL(
// All the ops below were revised in opset 22, which widened their float type constraints to the
// full set of float types by adding bfloat16. These kernels still implement float and double only,
// so the opset 22 registrations declare the same type constraints as the earlier ones. They are
// needed because a kernel registered without an end version matches its start version exactly
// (see VerifyVersion in kernel_registry.cc), so without them these ops have no CPU kernel at all
// at opset 22+, for any dtype.
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
RandomNormal,
1,
21,
KernelDefBuilder()
.TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalOutputTypes>()),
RandomNormal);

ONNX_CPU_OPERATOR_KERNEL(
RandomNormal,
22,
KernelDefBuilder()
.TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalOutputTypes>()),
RandomNormal);

ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
RandomUniform,
1,
21,
KernelDefBuilder()
.TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformOutputTypes>()),
RandomUniform);

ONNX_CPU_OPERATOR_KERNEL(
RandomUniform,
22,
KernelDefBuilder()
.TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformOutputTypes>()),
RandomUniform);

ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
RandomNormalLike,
1,
21,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T2",
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalLikeOutputTypes>()),
RandomNormalLike);

ONNX_CPU_OPERATOR_KERNEL(
RandomNormalLike,
22,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T2",
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalLikeOutputTypes>()),
RandomNormalLike);

ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
RandomUniformLike,
1,
21,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T2",
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformLikeOutputTypes>()),
RandomUniformLike);

// https://github.com/onnx/onnx/blob/main/docs/Operators.md#multinomial
ONNX_CPU_OPERATOR_KERNEL(
RandomUniformLike,
22,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T2",
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformLikeOutputTypes>()),
RandomUniformLike);

// https://github.com/onnx/onnx/blob/main/docs/Operators.md#multinomial
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Multinomial,
7,
21,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
.TypeConstraint("T2",
BuildKernelDefConstraintsFromTypeList<EnabledMultinomialOutputTypes>()),
Multinomial);

ONNX_CPU_OPERATOR_KERNEL(
Multinomial,
22,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
.TypeConstraint("T2",
Expand Down
91 changes: 81 additions & 10 deletions onnxruntime/test/providers/cpu/generator/random_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3,15 +3,28 @@

#include "gtest/gtest.h"
#include "test/providers/provider_test_utils.h"
#include "test/util/include/default_providers.h"

#include <algorithm>
#include <random>
using namespace ONNX_NAMESPACE;
namespace onnxruntime {
namespace test {

TEST(Random, RandomNormal2DDouble) {
OpTester test("RandomNormal");
// OpTester's default Run() pre-assigns every node to an EP and, when no EP has a matching kernel,
// logs a warning and skips the model without failing. That would make a test for a kernel
// registration vacuous, so run with an explicitly specified CPU EP instead: that path leaves node
// assignment to the graph partitioner, so a missing registration surfaces as a session
// initialization failure. It also pins execution to the CPU kernel, which is what the expected
// outputs below were generated against.
static void RunOnCpuEp(OpTester& test) {
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

static void RunRandomNormal2DDouble(int opset_version) {
OpTester test("RandomNormal", opset_version);

std::vector<int64_t> dims{20, 50};

Expand All @@ -34,14 +47,28 @@ TEST(Random, RandomNormal2DDouble) {

test.AddOutput<double>("Y", dims, expected_output);

if (opset_version >= 22) {
RunOnCpuEp(test);
return;
}

// The expected_output is generated using std lib, which is used by CPU kernel only.
// So we need to exclude other EPs here. Ditto for other places.
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCudaExecutionProvider, kCudaNHWCExecutionProvider});
}

void RunRandomNormalLike3DFloat(bool infer_dtype = false) {
OpTester test("RandomNormalLike");
TEST(Random, RandomNormal2DDouble) {
RunRandomNormal2DDouble(7);
}

// The op was revised in opset 22 and needs a kernel registration of its own for that opset.
TEST(Random, RandomNormal2DDoubleOpset22) {
RunRandomNormal2DDouble(22);
}

void RunRandomNormalLike3DFloat(bool infer_dtype = false, int opset_version = 7) {
OpTester test("RandomNormalLike", opset_version);

std::vector<int64_t> dims{2, 2, 3};

Expand Down Expand Up @@ -72,6 +99,11 @@ void RunRandomNormalLike3DFloat(bool infer_dtype = false) {

test.AddOutput<float>("Y", dims, expected_output);

if (opset_version >= 22) {
RunOnCpuEp(test);
return;
}

// TensorRT does not support manual seed overrides and there will be result mismatch
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCudaExecutionProvider, kCudaNHWCExecutionProvider, kTensorrtExecutionProvider});
Expand All @@ -86,8 +118,12 @@ TEST(Random, RandomNormalLikeInferDType) {
RunRandomNormalLike3DFloat(infer_dtype);
}

TEST(Random, RandomUniform1DFloat) {
OpTester test("RandomUniform");
TEST(Random, RandomNormalLike3DFloatOpset22) {
RunRandomNormalLike3DFloat(false, 22);
}

static void RunRandomUniform1DFloat(int opset_version) {
OpTester test("RandomUniform", opset_version);

std::vector<int64_t> dims{10};

Expand All @@ -110,13 +146,26 @@ TEST(Random, RandomUniform1DFloat) {

test.AddOutput<float>("Y", dims, expected_output);

if (opset_version >= 22) {
RunOnCpuEp(test);
return;
}

// TensorRT does not support manual seed overrides and there will be result mismatch
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCudaExecutionProvider, kCudaNHWCExecutionProvider, kTensorrtExecutionProvider});
}

void RunRandomUniformLikeTest(bool infer_dtype = false) {
OpTester test("RandomUniformLike");
TEST(Random, RandomUniform1DFloat) {
RunRandomUniform1DFloat(7);
}

TEST(Random, RandomUniform1DFloatOpset22) {
RunRandomUniform1DFloat(22);
}

void RunRandomUniformLikeTest(bool infer_dtype = false, int opset_version = 7) {
OpTester test("RandomUniformLike", opset_version);

std::vector<int64_t> dims{2, 6};

Expand Down Expand Up @@ -144,6 +193,11 @@ void RunRandomUniformLikeTest(bool infer_dtype = false) {

test.AddOutput<double>("Y", dims, expected_output);

if (opset_version >= 22) {
RunOnCpuEp(test);
return;
}

// TensorRT does not support seed parameter and there will be result mismatch
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCudaExecutionProvider, kCudaNHWCExecutionProvider, kTensorrtExecutionProvider});
Expand All @@ -158,6 +212,10 @@ TEST(Random, RandomUniformLikeInferDType) {
RunRandomUniformLikeTest(infer_dtype);
}

TEST(Random, RandomUniformLike2DDoubleOpset22) {
RunRandomUniformLikeTest(false, 22);
}

TEST(Random, InvalidDType) {
constexpr float seed = 123.f;

Expand Down Expand Up @@ -236,8 +294,8 @@ test cases but they use a different RNG (Philox) and hence the test results diff
of the op is same as tensorflow, for now I've just relied on the output generated by this code as ground truth
for verification.
*/
TEST(Random, MultinomialGoodCase) {
OpTester test("Multinomial");
static void RunMultinomialGoodCase(int opset_version) {
OpTester test("Multinomial", opset_version);

constexpr int64_t num_samples = 10;
constexpr float seed = 1618.f;
Expand All @@ -263,9 +321,22 @@ TEST(Random, MultinomialGoodCase) {
#endif
test.AddOutput<int64_t>("Y", output_dims, expected_output);

if (opset_version >= 22) {
RunOnCpuEp(test);
return;
}

test.Run();
}

TEST(Random, MultinomialGoodCase) {
RunMultinomialGoodCase(7);
}

TEST(Random, MultinomialGoodCaseOpset22) {
RunMultinomialGoodCase(22);
}

TEST(Random, MultinomialDefaultDType) {
auto run_test = [](int num_run_calls, const std::vector<int32_t>& expected_output) {
OpTester test("Multinomial");
Expand Down