(self)
| 598 | |
| 599 | @test_util.run_deprecated_v1 |
| 600 | def testGetLossesFor(self): |
| 601 | |
| 602 | class MyLayer(base_layers.Layer): |
| 603 | |
| 604 | def build(self, input_shape): |
| 605 | self.a = self.add_variable('a', |
| 606 | (), |
| 607 | dtypes.float32, |
| 608 | trainable=False) |
| 609 | self.b = self.add_variable('b', |
| 610 | (), |
| 611 | dtypes.float32, |
| 612 | trainable=False) |
| 613 | self.add_loss(self.a) |
| 614 | self.built = True |
| 615 | |
| 616 | def call(self, inputs): |
| 617 | self.add_loss(inputs, inputs=True) |
| 618 | return inputs + 1 |
| 619 | |
| 620 | layer = MyLayer() |
| 621 | inputs = array_ops.placeholder(dtypes.float32, (), 'inputs') |
| 622 | intermediate_inputs = inputs + 1 |
| 623 | outputs = layer.apply(intermediate_inputs) |
| 624 | |
| 625 | self.assertEqual(len(layer.losses), 2) |
| 626 | self.assertEqual(len(layer.get_losses_for(None)), 1) |
| 627 | self.assertEqual(len(layer.get_losses_for([inputs])), 1) |
| 628 | self.assertEqual(len(layer.get_losses_for([intermediate_inputs])), 1) |
| 629 | self.assertEqual(len(layer.get_losses_for([outputs])), 0) |
| 630 | |
| 631 | # Call same layer on new input, creating one more conditional loss |
| 632 | inputs = array_ops.placeholder(dtypes.float32, (), 'inputs') |
| 633 | intermediate_inputs = inputs + 1 |
| 634 | outputs = layer.apply(intermediate_inputs) |
| 635 | |
| 636 | self.assertEqual(len(layer.losses), 3) |
| 637 | self.assertEqual(len(layer.get_losses_for(None)), 1) |
| 638 | # Check that we are successfully filtering out irrelevant losses |
| 639 | self.assertEqual(len(layer.get_losses_for([inputs])), 1) |
| 640 | self.assertEqual(len(layer.get_losses_for([intermediate_inputs])), 1) |
| 641 | self.assertEqual(len(layer.get_losses_for([outputs])), 0) |
| 642 | |
| 643 | |
| 644 | class IdentityLayer(base_layers.Layer): |
nothing calls this directly
no test coverage detected