| 335 | #undef REGISTER_TEST |
| 336 | |
| 337 | TEST_F(MklRemapperTest, FuseBatchNormWithRelu) { |
| 338 | using ::tensorflow::ops::Placeholder; |
| 339 | |
| 340 | for (bool is_training : {true, false}) { |
| 341 | for (bool has_side_input : {true, false}) { |
| 342 | tensorflow::Scope s = tensorflow::Scope::NewRootScope(); |
| 343 | |
| 344 | const int num_channels = 24; |
| 345 | |
| 346 | TensorShape channel_shape({num_channels}); |
| 347 | TensorShape empty_shape({0}); |
| 348 | |
| 349 | auto input = |
| 350 | Placeholder(s.WithOpName("input"), DT_FLOAT, |
| 351 | ops::Placeholder::Shape({2, 8, 8, num_channels})); |
| 352 | auto input_cast = ops::Cast(s.WithOpName("input_cast"), input, DT_FLOAT); |
| 353 | auto scale = Placeholder(s.WithOpName("scale"), DT_FLOAT); |
| 354 | auto offset = Placeholder(s.WithOpName("offset"), DT_FLOAT); |
| 355 | auto mean = Placeholder(s.WithOpName("mean"), DT_FLOAT); |
| 356 | auto var = Placeholder(s.WithOpName("var"), DT_FLOAT); |
| 357 | |
| 358 | float epsilon = 0.1f; |
| 359 | auto fbn = |
| 360 | ops::FusedBatchNormV3(s.WithOpName("fused_batch_norm"), input_cast, |
| 361 | scale, offset, mean, var, |
| 362 | ops::FusedBatchNormV3::IsTraining(is_training) |
| 363 | .Epsilon(epsilon) |
| 364 | .DataFormat("NHWC")); |
| 365 | |
| 366 | if (has_side_input) { |
| 367 | auto side_input = |
| 368 | Placeholder(s.WithOpName("side_input"), DT_FLOAT, |
| 369 | ops::Placeholder::Shape({2, 8, 8, num_channels})); |
| 370 | auto side_input_cast = |
| 371 | ops::Cast(s.WithOpName("side_input_cast"), side_input, DT_FLOAT); |
| 372 | auto add = ops::Add(s.WithOpName("add"), fbn.y, side_input_cast); |
| 373 | auto relu = ops::Relu(s.WithOpName("relu"), add); |
| 374 | } else { |
| 375 | auto relu = ops::Relu(s.WithOpName("relu"), fbn.y); |
| 376 | } |
| 377 | |
| 378 | auto input_t = GenerateRandomTensor<DT_FLOAT>({2, 8, 8, num_channels}); |
| 379 | auto scale_t = GenerateRandomTensor<DT_FLOAT>(channel_shape); |
| 380 | auto offset_t = GenerateRandomTensor<DT_FLOAT>(channel_shape); |
| 381 | auto mean_t = GenerateRandomTensor<DT_FLOAT>(is_training ? empty_shape |
| 382 | : channel_shape); |
| 383 | auto var_t = GenerateRandomTensor<DT_FLOAT>(is_training ? empty_shape |
| 384 | : channel_shape); |
| 385 | auto side_input_t = |
| 386 | GenerateRandomTensor<DT_FLOAT>({2, 8, 8, num_channels}); |
| 387 | |
| 388 | GrapplerItem item; |
| 389 | item.fetch = {"relu"}; |
| 390 | if (has_side_input) |
| 391 | item.feed = {{"input", input_t}, {"scale", scale_t}, |
| 392 | {"offset", offset_t}, {"mean", mean_t}, |
| 393 | {"var", var_t}, {"side_input", side_input_t}}; |
| 394 | else |
nothing calls this directly
no test coverage detected