| 1171 | } |
| 1172 | |
| 1173 | Status FusedBatchNormExShape(shape_inference::InferenceContext* c) { |
| 1174 | TF_RETURN_IF_ERROR(FusedBatchNormV3Shape(c)); |
| 1175 | |
| 1176 | string data_format_str; |
| 1177 | TF_RETURN_IF_ERROR(c->GetAttr("data_format", &data_format_str)); |
| 1178 | TensorFormat data_format; |
| 1179 | if (!FormatFromString(data_format_str, &data_format)) { |
| 1180 | return errors::InvalidArgument("Invalid data format string: ", |
| 1181 | data_format_str); |
| 1182 | } |
| 1183 | ShapeHandle x; |
| 1184 | TF_RETURN_IF_ERROR(c->WithRank(c->input(0), 4, &x)); |
| 1185 | |
| 1186 | int channel_dim_index = GetTensorFeatureDimIndex(4, data_format); |
| 1187 | DimensionHandle channel_dim = c->Dim(x, channel_dim_index); |
| 1188 | |
| 1189 | // This is a cuDNN implementation constraint. |
| 1190 | if (c->ValueKnown(channel_dim) && c->Value(channel_dim) % 4 != 0) { |
| 1191 | return errors::InvalidArgument( |
| 1192 | "_FusedBatchNormEx channel dimension must be divisible by 4."); |
| 1193 | } |
| 1194 | |
| 1195 | return Status::OK(); |
| 1196 | } |
| 1197 | |
| 1198 | Status FusedBatchNormGradShape(shape_inference::InferenceContext* c) { |
| 1199 | string data_format_str; |
nothing calls this directly
no test coverage detected