2756 const std::string descriptorName{
"QuantizedLstmQueueDescriptor"};
2759 ValidateNumInputs(workloadInfo, descriptorName, 3);
2760 ValidateNumOutputs(workloadInfo, descriptorName, 2);
2770 std::vector<DataType> inputOutputSupportedTypes =
2775 std::vector<DataType> cellStateSupportedTypes =
2780 std::vector<DataType> weightsSupportedTypes =
2785 std::vector<DataType> biasSupportedTypes =
2791 ValidateDataTypes(inputInfo, inputOutputSupportedTypes, descriptorName);
2792 ValidateDataTypes(cellStateInInfo, cellStateSupportedTypes, descriptorName);
2793 ValidateDataTypes(outputStateInInfo, inputOutputSupportedTypes, descriptorName);
2795 ValidateDataTypes(cellStateOutInfo, cellStateSupportedTypes, descriptorName);
2796 ValidateDataTypes(outputStateOutInfo, inputOutputSupportedTypes, descriptorName);
2799 ValidateTensorDataTypesMatch(inputInfo, outputStateInInfo, descriptorName,
"input",
"outputStateIn");
2800 ValidateTensorDataTypesMatch(outputStateInInfo, outputStateOutInfo, descriptorName,
2801 "outputStateIn",
"outputStateOut");
2802 ValidateTensorDataTypesMatch(cellStateInInfo, cellStateOutInfo, descriptorName,
"cellStateIn",
"cellStateOut");
2805 ValidateTensorQuantizationSpace(inputInfo, outputStateInInfo, descriptorName,
"input",
"outputStateIn");
2806 ValidateTensorQuantizationSpace(inputInfo, outputStateOutInfo, descriptorName,
"input",
"outputStateOut");
2807 ValidateTensorQuantizationSpace(cellStateInInfo, cellStateOutInfo, descriptorName,
"cellStateIn",
"cellStateOut");
2810 const uint32_t numBatches = inputInfo.GetShape()[0];
2811 const uint32_t inputSize = inputInfo.GetShape()[1];
2812 const uint32_t outputSize = cellStateInInfo.GetShape()[1];
2815 ValidateTensorNumDimNumElem(inputInfo, 2, (numBatches * inputSize), descriptorName +
" input");
2816 ValidateTensorNumDimNumElem(cellStateInInfo, 2, (numBatches * outputSize), descriptorName +
" cellStateIn");
2817 ValidateTensorNumDimNumElem(outputStateInInfo, 2, (numBatches * outputSize), descriptorName +
" outputStateIn");
2818 ValidateTensorNumDimNumElem(cellStateOutInfo, 2, (numBatches * outputSize), descriptorName +
" cellStateOut");
2819 ValidateTensorNumDimNumElem(outputStateOutInfo, 2, (numBatches * outputSize), descriptorName +
" outputStateOut");
2824 ValidateTensorNumDimNumElem(inputToInputWeightsInfo, 2, (outputSize * inputSize),
" InputToInputWeights");
2828 ValidateTensorNumDimNumElem(inputToForgetWeightsInfo, 2, (outputSize * inputSize),
" InputToForgetWeights");
2832 ValidateTensorNumDimNumElem(inputToCellWeightsInfo, 2, (outputSize * inputSize),
" InputToCellWeights");
2836 ValidateTensorNumDimNumElem(inputToOutputWeightsInfo, 2, (outputSize * inputSize),
" InputToOutputWeights");
2840 ValidateTensorNumDimNumElem(recurrentToInputWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToInputWeights");
2844 ValidateTensorNumDimNumElem(recurrentToForgetWeightsInfo, 2, (outputSize * outputSize),
2845 " RecurrentToForgetWeights");
2849 ValidateTensorNumDimNumElem(recurrentToCellWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToCellWeights");
2853 ValidateTensorNumDimNumElem(recurrentToOutputWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToCellWeights");
2856 ValidateDataTypes(inputToInputWeightsInfo, weightsSupportedTypes, descriptorName);
2858 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToForgetWeightsInfo, descriptorName,
2859 "inputToInputWeights",
"inputToForgetWeights");
2860 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToCellWeightsInfo, descriptorName,
2861 "inputToInputWeights",
"inputToCellWeights");
2862 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToOutputWeightsInfo, descriptorName,
2863 "inputToInputWeights",
"inputToOutputWeights");
2865 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToInputWeightsInfo, descriptorName,
2866 "inputToInputWeights",
"recurrentToInputWeights");
2867 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToForgetWeightsInfo, descriptorName,
2868 "inputToInputWeights",
"recurrentToForgeteights");
2869 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToCellWeightsInfo, descriptorName,
2870 "inputToInputWeights",
"recurrentToCellWeights");
2871 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToOutputWeightsInfo, descriptorName,
2872 "inputToInputWeights",
"recurrentToOutputWeights");
2875 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToForgetWeightsInfo,
2876 descriptorName,
"inputToInputWeights",
"inputToForgetWeights");
2877 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToCellWeightsInfo,
2878 descriptorName,
"inputToInputWeights",
"inputToCellWeights");
2879 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToOutputWeightsInfo,
2880 descriptorName,
"inputToInputWeights",
"inputToOutputWeights");
2882 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToInputWeightsInfo,
2883 descriptorName,
"inputToInputWeights",
"recurrentToInputWeights");
2884 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToForgetWeightsInfo,
2885 descriptorName,
"inputToInputWeights",
"recurrentToForgetWeights");
2886 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToCellWeightsInfo,
2887 descriptorName,
"inputToInputWeights",
"recurrentToCellWeights");
2888 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToOutputWeightsInfo,
2889 descriptorName,
"inputToInputWeights",
"recurrentToOutputWeights");
2894 ValidateTensorNumDimNumElem(inputGateBiasInfo, 1, outputSize,
" InputGateBias");
2898 ValidateTensorNumDimNumElem(forgetGateBiasInfo, 1, outputSize,
" ForgetGateBias");
2900 ValidatePointer(
m_CellBias, descriptorName,
"CellBias");
2902 ValidateTensorNumDimNumElem(cellBiasInfo, 1, outputSize,
" CellBias");
2906 ValidateTensorNumDimNumElem(outputGateBiasInfo, 1, outputSize,
" OutputGateBias");
2909 ValidateDataTypes(inputGateBiasInfo, biasSupportedTypes, descriptorName);
2911 ValidateTensorDataTypesMatch(inputGateBiasInfo, forgetGateBiasInfo, descriptorName,
2912 "inputGateBias",
"forgetGateBias");
2913 ValidateTensorDataTypesMatch(inputGateBiasInfo, cellBiasInfo, descriptorName,
2914 "inputGateBias",
"cellBias");
2915 ValidateTensorDataTypesMatch(inputGateBiasInfo, outputGateBiasInfo, descriptorName,
2916 "inputGateBias",
"outputGateBias");
2919 ValidateBiasTensorQuantization(inputGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
2920 ValidateBiasTensorQuantization(forgetGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
2921 ValidateBiasTensorQuantization(cellBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
2922 ValidateBiasTensorQuantization(outputGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
const ConstCpuTensorHandle * m_RecurrentToForgetWeights
const ConstCpuTensorHandle * m_InputGateBias
const ConstCpuTensorHandle * m_InputToCellWeights
std::vector< TensorInfo > m_InputTensorInfos
const ConstCpuTensorHandle * m_ForgetGateBias
const ConstCpuTensorHandle * m_RecurrentToInputWeights
std::vector< TensorInfo > m_OutputTensorInfos
const ConstCpuTensorHandle * m_RecurrentToCellWeights
const ConstCpuTensorHandle * m_RecurrentToOutputWeights
const ConstCpuTensorHandle * m_CellBias
const ConstCpuTensorHandle * m_OutputGateBias
const ConstCpuTensorHandle * m_InputToForgetWeights
const ConstCpuTensorHandle * m_InputToOutputWeights
const ConstCpuTensorHandle * m_InputToInputWeights
const TensorInfo & GetTensorInfo() const