diff --git a/qa/L0_parameters/class_count_test.py b/qa/L0_parameters/class_count_test.py new file mode 100755 index 0000000000..9e49ae4ae2 --- /dev/null +++ b/qa/L0_parameters/class_count_test.py @@ -0,0 +1,105 @@ +#!/usr/bin/env python3 + +# Copyright 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions +# are met: +# * Redistributions of source code must retain the above copyright +# notice, this list of conditions and the following disclaimer. +# * Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# * Neither the name of NVIDIA CORPORATION nor the names of its +# contributors may be used to endorse or promote products derived +# from this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS ``AS IS'' AND ANY +# EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +# PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR +# CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, +# EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, +# PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR +# PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY +# OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +# (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import sys + +sys.path.append("../common") + +import os +import unittest + +import numpy as np +import test_util as tu +import tritonclient.grpc as grpcclient +import tritonclient.http as httpclient +from tritonclient.utils import InferenceServerException + + +class ClassificationParameterTest(tu.TestResultCollector): + def setUp(self): + self.protocol = os.environ.get("CLIENT_TYPE", "http") + if self.protocol == "http": + self.client = httpclient.InferenceServerClient("localhost:8000") + else: + self.client = grpcclient.InferenceServerClient("localhost:8001") + + def _prepare_io(self, input_data, dtype): + if self.protocol == "http": + inputs = [httpclient.InferInput("INPUT0", input_data.shape, dtype)] + outputs = [httpclient.InferRequestedOutput(name="OUTPUT0", class_count=5)] + else: + inputs = [grpcclient.InferInput("INPUT0", input_data.shape, dtype)] + outputs = [grpcclient.InferRequestedOutput(name="OUTPUT0", class_count=5)] + inputs[0].set_data_from_numpy(input_data) + return inputs, outputs + + def test_classificattion(self): + shape = (1, 8) + dtype = "FP32" + model_name = "identity_fp32" + input_data = np.ones(shape, dtype=np.float32) + + inputs, outputs = self._prepare_io(input_data, dtype) + result = self.client.infer( + model_name=model_name, inputs=inputs, outputs=outputs + ) + output = result.get_output("OUTPUT0") + if self.protocol == "http": + output_dtype = output["datatype"] + else: + output_dtype = output.datatype + + self.assertEqual(output_dtype, "BYTES") + + # Validate shape matches to the class_count + output_data = result.as_numpy("OUTPUT0") + self.assertIsNotNone(output_data) + self.assertEqual(output_data.shape, (1, 5)) + + for res_str_bytes in np.nditer(output_data, flags=["refs_ok"]): + res_str = res_str_bytes.item().decode("utf-8") + self.assertTrue(res_str.startswith("1.000000:")) + + def test_classificattion_unsupported_data_type(self): + shape = (1, 8) + model_name = "identity_bytes" + dtype = "BYTES" + input_data = np.array([["test"] * shape[1]], dtype=object) + + inputs, outputs = self._prepare_io(input_data, dtype) + with self.assertRaises(InferenceServerException) as e: + self.client.infer(model_name=model_name, inputs=inputs, outputs=outputs) + + self.assertIn( + "class result not available for output due to unsupported type 'BYTES'", + str(e.exception), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/qa/L0_parameters/test.sh b/qa/L0_parameters/test.sh index c53b02d4b7..63868237aa 100755 --- a/qa/L0_parameters/test.sh +++ b/qa/L0_parameters/test.sh @@ -1,5 +1,5 @@ #!/bin/bash -# Copyright 2023-2024, NVIDIA CORPORATION. All rights reserved. +# Copyright 2023-2025, NVIDIA CORPORATION. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -97,10 +97,65 @@ for i in "${all_tests[@]}"; do wait $SERVER_PID done + +# Test Classification Extension +PYTHON_MODELS_DIR="${PYTHON_MODELS_DIR:-/opt/tritonserver/qa/python_models}" +MODELDIR="models" +TEST_RESULT_FILE="test_results.txt" +TEST_SCRIPT_PY="./class_count_test.py" + +rm -rf $MODELDIR +mkdir -p "${MODELDIR}/identity_fp32/1" +cp ${PYTHON_MODELS_DIR}/identity_fp32/config.pbtxt "${MODELDIR}/identity_fp32/" +cp ${PYTHON_MODELS_DIR}/identity_fp32/model.py "${MODELDIR}/identity_fp32/1/" + +mkdir -p "${MODELDIR}/identity_bytes/1" +cp ${PYTHON_MODELS_DIR}/identity_fp32/config.pbtxt "${MODELDIR}/identity_bytes/" +cp ${PYTHON_MODELS_DIR}/identity_fp32/model.py "${MODELDIR}/identity_bytes/1/" +(cd "${MODELDIR}/identity_bytes" && \ + sed -i 's/identity_fp32/identity_bytes/' config.pbtxt && \ + sed -i 's/TYPE_FP32/TYPE_STRING/' config.pbtxt ) + +SERVER_ARGS="--model-repository=`pwd`/${MODELDIR} --log-verbose=1" +for client_type in http grpc; do + export CLIENT_TYPE=$client_type + SERVER_LOG="./class_count_test_${client_type}_server.log" + CLIENT_LOG="./class_count_test_${client_type}_client.log" + rm -f $SERVER_LOG $CLIENT_LOG + run_server + if [ "$SERVER_PID" == "0" ]; then + echo -e "\n***\n*** Failed to start $SERVER\n***" + cat $SERVER_LOG + exit 1 + fi + + set +e + python3 $TEST_SCRIPT_PY -v >>"$CLIENT_LOG" 2>&1 + if [ $? -ne 0 ]; then + cat $CLIENT_LOG + echo -e "\n***\n*** Test Failed - class_count_${client_type}_test_client\n***" + RET=1 + else + check_test_results $TEST_RESULT_FILE 2 + if [ $? -ne 0 ]; then + cat $TEST_RESULT_FILE + echo -e "\n***\n*** Test Result Verification Failed - class_count_${client_type}_test_client\n***" + RET=1 + fi + fi + kill $SERVER_PID + wait $SERVER_PID + + if [ $? -ne 0 ]; then + echo -e "\n***\n*** Test Server shut down non-gracefully\n***" + RET=1 + fi + set -e +done + if [ $RET -eq 0 ]; then echo -e "\n***\n*** Test Passed\n***" else - cat $CLIENT_LOG echo -e "\n***\n*** Test FAILED\n***" fi diff --git a/src/classification.cc b/src/classification.cc index 2d8cd26b9e..828edf613d 100644 --- a/src/classification.cc +++ b/src/classification.cc @@ -1,4 +1,4 @@ -// Copyright (c) 2020, NVIDIA CORPORATION. All rights reserved. +// Copyright (c) 2020-2025, NVIDIA CORPORATION. All rights reserved. // // Redistribution and use in source and binary forms, with or without // modification, are permitted provided that the following conditions @@ -77,8 +77,18 @@ TopkClassifications( const TRITONSERVER_DataType datatype, const uint32_t req_class_count, std::vector* class_strs) { - const size_t element_cnt = - byte_size / TRITONSERVER_DataTypeByteSize(datatype); + const uint32_t dtype_byte_size = TRITONSERVER_DataTypeByteSize(datatype); + if (dtype_byte_size == 0) { + return TRITONSERVER_ErrorNew( + TRITONSERVER_ERROR_INVALID_ARG, + std::string( + std::string("class result not available for output due to " + "unsupported type '") + + std::string(TRITONSERVER_DataTypeString(datatype)) + "'") + .c_str()); + } + + const size_t element_cnt = byte_size / dtype_byte_size; switch (datatype) { case TRITONSERVER_TYPE_UINT8: