[MLIR] Add minor fix for filter dimensions. This commit addresses a bug where incorrect filter dimensions were being read and used. Corrected the bug and the relevant test cases were also checked. Signed-off-by: Prateek Gupta <prateek@polymagelabs.com>
diff --git a/tensorflow/compiler/mlir/tensorflow/ir/tf_ops_a_m.cc b/tensorflow/compiler/mlir/tensorflow/ir/tf_ops_a_m.cc index edf90f1..41f628b 100644 --- a/tensorflow/compiler/mlir/tensorflow/ir/tf_ops_a_m.cc +++ b/tensorflow/compiler/mlir/tensorflow/ir/tf_ops_a_m.cc
@@ -1549,9 +1549,9 @@ return_shape[GetTensorBatchDimIndex(num_dims, format)] = input_ty.getShape()[GetTensorBatchDimIndex(num_dims, format)]; return_shape[GetTensorFeatureDimIndex(num_dims, format)] = - filter_ty.getShape()[GetFilterTensorInnerInputChannelsDimIndex( - num_dims, tensorflow::FilterTensorFormat::FORMAT_HWIO)]; - + filter_ty.getShape()[GetFilterTensorOutputChannelsDimIndex( + num_dims, tensorflow::FORMAT_HWIO)]; + inferredReturnTypes.assign( {RankedTensorType::get(return_shape, input_ty.getElementType())}); return success();