3252 const std::string descriptorName{
"QuantizedLstmQueueDescriptor"};
3255 ValidateNumInputs(workloadInfo, descriptorName, 3);
3256 ValidateNumOutputs(workloadInfo, descriptorName, 2);
3266 std::vector<DataType> inputOutputSupportedTypes =
3271 std::vector<DataType> cellStateSupportedTypes =
3276 std::vector<DataType> weightsSupportedTypes =
3281 std::vector<DataType> biasSupportedTypes =
3287 ValidateDataTypes(inputInfo, inputOutputSupportedTypes, descriptorName);
3288 ValidateDataTypes(cellStateInInfo, cellStateSupportedTypes, descriptorName);
3289 ValidateDataTypes(outputStateInInfo, inputOutputSupportedTypes, descriptorName);
3291 ValidateDataTypes(cellStateOutInfo, cellStateSupportedTypes, descriptorName);
3292 ValidateDataTypes(outputStateOutInfo, inputOutputSupportedTypes, descriptorName);
3295 ValidateTensorDataTypesMatch(inputInfo, outputStateInInfo, descriptorName,
"input",
"outputStateIn");
3296 ValidateTensorDataTypesMatch(outputStateInInfo, outputStateOutInfo, descriptorName,
3297 "outputStateIn",
"outputStateOut");
3298 ValidateTensorDataTypesMatch(cellStateInInfo, cellStateOutInfo, descriptorName,
"cellStateIn",
"cellStateOut");
3301 ValidateTensorQuantizationSpace(inputInfo, outputStateInInfo, descriptorName,
"input",
"outputStateIn");
3302 ValidateTensorQuantizationSpace(inputInfo, outputStateOutInfo, descriptorName,
"input",
"outputStateOut");
3303 ValidateTensorQuantizationSpace(cellStateInInfo, cellStateOutInfo, descriptorName,
"cellStateIn",
"cellStateOut");
3306 const uint32_t numBatches = inputInfo.GetShape()[0];
3307 const uint32_t inputSize = inputInfo.GetShape()[1];
3308 const uint32_t outputSize = cellStateInInfo.GetShape()[1];
3311 ValidateTensorNumDimNumElem(inputInfo, 2, (numBatches * inputSize), descriptorName +
" input");
3312 ValidateTensorNumDimNumElem(cellStateInInfo, 2, (numBatches * outputSize), descriptorName +
" cellStateIn");
3313 ValidateTensorNumDimNumElem(outputStateInInfo, 2, (numBatches * outputSize), descriptorName +
" outputStateIn");
3314 ValidateTensorNumDimNumElem(cellStateOutInfo, 2, (numBatches * outputSize), descriptorName +
" cellStateOut");
3315 ValidateTensorNumDimNumElem(outputStateOutInfo, 2, (numBatches * outputSize), descriptorName +
" outputStateOut");
3320 ValidateTensorNumDimNumElem(inputToInputWeightsInfo, 2, (outputSize * inputSize),
" InputToInputWeights");
3324 ValidateTensorNumDimNumElem(inputToForgetWeightsInfo, 2, (outputSize * inputSize),
" InputToForgetWeights");
3328 ValidateTensorNumDimNumElem(inputToCellWeightsInfo, 2, (outputSize * inputSize),
" InputToCellWeights");
3332 ValidateTensorNumDimNumElem(inputToOutputWeightsInfo, 2, (outputSize * inputSize),
" InputToOutputWeights");
3336 ValidateTensorNumDimNumElem(recurrentToInputWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToInputWeights");
3340 ValidateTensorNumDimNumElem(recurrentToForgetWeightsInfo, 2, (outputSize * outputSize),
3341 " RecurrentToForgetWeights");
3345 ValidateTensorNumDimNumElem(recurrentToCellWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToCellWeights");
3349 ValidateTensorNumDimNumElem(recurrentToOutputWeightsInfo, 2, (outputSize * outputSize),
" RecurrentToCellWeights");
3352 ValidateDataTypes(inputToInputWeightsInfo, weightsSupportedTypes, descriptorName);
3354 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToForgetWeightsInfo, descriptorName,
3355 "inputToInputWeights",
"inputToForgetWeights");
3356 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToCellWeightsInfo, descriptorName,
3357 "inputToInputWeights",
"inputToCellWeights");
3358 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, inputToOutputWeightsInfo, descriptorName,
3359 "inputToInputWeights",
"inputToOutputWeights");
3361 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToInputWeightsInfo, descriptorName,
3362 "inputToInputWeights",
"recurrentToInputWeights");
3363 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToForgetWeightsInfo, descriptorName,
3364 "inputToInputWeights",
"recurrentToForgeteights");
3365 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToCellWeightsInfo, descriptorName,
3366 "inputToInputWeights",
"recurrentToCellWeights");
3367 ValidateTensorDataTypesMatch(inputToInputWeightsInfo, recurrentToOutputWeightsInfo, descriptorName,
3368 "inputToInputWeights",
"recurrentToOutputWeights");
3371 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToForgetWeightsInfo,
3372 descriptorName,
"inputToInputWeights",
"inputToForgetWeights");
3373 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToCellWeightsInfo,
3374 descriptorName,
"inputToInputWeights",
"inputToCellWeights");
3375 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, inputToOutputWeightsInfo,
3376 descriptorName,
"inputToInputWeights",
"inputToOutputWeights");
3378 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToInputWeightsInfo,
3379 descriptorName,
"inputToInputWeights",
"recurrentToInputWeights");
3380 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToForgetWeightsInfo,
3381 descriptorName,
"inputToInputWeights",
"recurrentToForgetWeights");
3382 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToCellWeightsInfo,
3383 descriptorName,
"inputToInputWeights",
"recurrentToCellWeights");
3384 ValidateTensorQuantizationSpace(inputToInputWeightsInfo, recurrentToOutputWeightsInfo,
3385 descriptorName,
"inputToInputWeights",
"recurrentToOutputWeights");
3390 ValidateTensorNumDimNumElem(inputGateBiasInfo, 1, outputSize,
" InputGateBias");
3394 ValidateTensorNumDimNumElem(forgetGateBiasInfo, 1, outputSize,
" ForgetGateBias");
3396 ValidatePointer(
m_CellBias, descriptorName,
"CellBias");
3398 ValidateTensorNumDimNumElem(cellBiasInfo, 1, outputSize,
" CellBias");
3402 ValidateTensorNumDimNumElem(outputGateBiasInfo, 1, outputSize,
" OutputGateBias");
3405 ValidateDataTypes(inputGateBiasInfo, biasSupportedTypes, descriptorName);
3407 ValidateTensorDataTypesMatch(inputGateBiasInfo, forgetGateBiasInfo, descriptorName,
3408 "inputGateBias",
"forgetGateBias");
3409 ValidateTensorDataTypesMatch(inputGateBiasInfo, cellBiasInfo, descriptorName,
3410 "inputGateBias",
"cellBias");
3411 ValidateTensorDataTypesMatch(inputGateBiasInfo, outputGateBiasInfo, descriptorName,
3412 "inputGateBias",
"outputGateBias");
3415 ValidateBiasTensorQuantization(inputGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
3416 ValidateBiasTensorQuantization(forgetGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
3417 ValidateBiasTensorQuantization(cellBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
3418 ValidateBiasTensorQuantization(outputGateBiasInfo, inputInfo, inputToInputWeightsInfo, descriptorName);
const ConstTensorHandle * m_InputGateBias
const ConstTensorHandle * m_RecurrentToInputWeights
const TensorInfo & GetTensorInfo() const
std::vector< TensorInfo > m_InputTensorInfos
const ConstTensorHandle * m_InputToForgetWeights
const ConstTensorHandle * m_RecurrentToCellWeights
const ConstTensorHandle * m_ForgetGateBias
std::vector< TensorInfo > m_OutputTensorInfos
const ConstTensorHandle * m_RecurrentToOutputWeights
const ConstTensorHandle * m_OutputGateBias
const ConstTensorHandle * m_RecurrentToForgetWeights
const ConstTensorHandle * m_InputToOutputWeights
const ConstTensorHandle * m_InputToInputWeights
const ConstTensorHandle * m_CellBias
const ConstTensorHandle * m_InputToCellWeights