(self, image, fused, freeze_mode)
| 42 | class BNTest(test.TestCase): |
| 43 | |
| 44 | def _simple_model(self, image, fused, freeze_mode): |
| 45 | output_channels, kernel_size = 2, 3 |
| 46 | conv = conv_layers.conv2d( |
| 47 | image, |
| 48 | output_channels, |
| 49 | kernel_size, |
| 50 | use_bias=False, |
| 51 | kernel_initializer=init_ops.ones_initializer()) |
| 52 | bn_layer = normalization_layers.BatchNormalization(fused=fused) |
| 53 | bn_layer._bessels_correction_test_only = False |
| 54 | training = not freeze_mode |
| 55 | bn = bn_layer.apply(conv, training=training) |
| 56 | loss = math_ops.reduce_sum(math_ops.abs(bn)) |
| 57 | optimizer = gradient_descent.GradientDescentOptimizer(0.01) |
| 58 | if not freeze_mode: |
| 59 | update_ops = ops.get_collection(ops.GraphKeys.UPDATE_OPS) |
| 60 | with ops.control_dependencies(update_ops): |
| 61 | train_op = optimizer.minimize(loss) |
| 62 | else: |
| 63 | train_op = optimizer.minimize(loss) |
| 64 | saver = saver_lib.Saver(write_version=saver_pb2.SaverDef.V2) |
| 65 | return loss, train_op, saver |
| 66 | |
| 67 | def _train(self, |
| 68 | checkpoint_path, |
no test coverage detected