diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index ff06f7fdb3..46cdf4bc1f 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -262,9 +262,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Aco class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Atanh); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scan); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scatter); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, string, TfIdfVectorizer); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, TfIdfVectorizer); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, TfIdfVectorizer); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, TfIdfVectorizer); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, bool, NonZero); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, NonZero); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, NonZero); @@ -881,12 +879,7 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { Scan)>, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo()) - .TypeConstraint("T1", DataTypeImpl::GetTensorType()), - TfIdfVectorizer); - -ONNX_CPU_OPERATOR_TYPED_KERNEL( - TfIdfVectorizer, - 9, - int32_t, - KernelDefBuilder() - .TypeConstraint("T", DataTypeImpl::GetTensorType()) - .TypeConstraint("T1", DataTypeImpl::GetTensorType()), - TfIdfVectorizer); - -ONNX_CPU_OPERATOR_TYPED_KERNEL( - TfIdfVectorizer, - 9, - int64_t, - KernelDefBuilder() - .TypeConstraint("T", DataTypeImpl::GetTensorType()) + .TypeConstraint("T", {DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType()}) .TypeConstraint("T1", DataTypeImpl::GetTensorType()), TfIdfVectorizer); @@ -405,11 +388,14 @@ void TfIdfVectorizer::OutputResult(OpKernelContext* ctx, size_t B, const std::ve std::vector output_dims; if (B == 0) { output_dims.push_back(impl.output_size_); + B = 1; // For use in the loops below } else { output_dims.push_back(B); output_dims.push_back(impl.output_size_); } + const auto row_size = impl.output_size_; + TensorShape output_shape(output_dims); assert(frequences.size() == static_cast(output_shape.Size())); @@ -424,9 +410,11 @@ void TfIdfVectorizer::OutputResult(OpKernelContext* ctx, size_t B, const std::ve } break; case kIDF: { if (!w.empty()) { - assert(frequences.size() == w.size()); - for (size_t i = 0; i < frequences.size(); ++i) { - *output_data++ = (frequences[i] > 0) ? w[i] : 0; + const auto* freqs = frequences.data(); + for (size_t batch = 0; batch < B; ++batch) { + for (size_t i = 0; i < row_size; ++i) { + *output_data++ = (*freqs++ > 0) ? w[i] : 0; + } } } else { for (auto f : frequences) { @@ -436,9 +424,11 @@ void TfIdfVectorizer::OutputResult(OpKernelContext* ctx, size_t B, const std::ve } break; case kTFIDF: { if (!w.empty()) { - assert(frequences.size() == w.size()); - for (size_t i = 0; i < frequences.size(); ++i) { - *output_data++ = frequences[i] * w[i]; + const auto* freqs = frequences.data(); + for (size_t batch = 0; batch < B; ++batch) { + for (size_t i = 0; i < row_size; ++i) { + *output_data++ = *freqs++ * w[i]; + } } } else { for (auto f : frequences) { diff --git a/onnxruntime/test/providers/cpu/nn/tfidfvectorizer_test.cc b/onnxruntime/test/providers/cpu/nn/tfidfvectorizer_test.cc index fe4396bae9..564a6fef74 100644 --- a/onnxruntime/test/providers/cpu/nn/tfidfvectorizer_test.cc +++ b/onnxruntime/test/providers/cpu/nn/tfidfvectorizer_test.cc @@ -669,5 +669,25 @@ TEST(TfIdfVectorizerTest, String_TFIDFWeights_onlyBigrams_Skip5) { test.Run(OpTester::ExpectResult::kExpectSuccess); } +TEST(TfIdfVectorizerTest, String_TFIDFWeights_onlyBigrams_Skip5_2rows) { + OpTester test("TfIdfVectorizer", opset_ver); + // s=5, Min=Max=2, weights specified, string + InitTestAttr(test, "TFIDF", 2, 2, 5, + {0, 4}, + {0, 1, 2, 3, 4, 5, 6}, //7 output indexes + {2.0f, 2.0f, 2.0f, 2.0f, 2.0f, 3.0f, 2.0f}, // weights + {}, + {"two", "three", "five", "four", //1-grams + "five", "six", "seven", "eight", "six", "seven"}); //bi-grams + + test.AddInput("T", {2, 6}, {"one", "one", "three", "three", "three", "seven", + "eight", "six", "seven", "five", "six", "eight"}); + + test.AddOutput("Y", {2, 7}, {0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, // No bi-grams in the first row + 0.f, 0.f, 0.f, 0.f, 2.f, 3.f, 2.f}); + + test.Run(OpTester::ExpectResult::kExpectSuccess); +} + } // namespace test } // namespace onnxruntime