(self)
| 166 | # TODO(b/120949004): Re-enable garbage collection check |
| 167 | # @test_util.run_in_graph_and_eager_modes(assert_no_eager_garbage=True) |
| 168 | def test_mean(self): |
| 169 | m = metrics.Mean(name='my_mean') |
| 170 | |
| 171 | # check config |
| 172 | self.assertEqual(m.name, 'my_mean') |
| 173 | self.assertTrue(m.stateful) |
| 174 | self.assertEqual(m.dtype, dtypes.float32) |
| 175 | self.assertEqual(len(m.variables), 2) |
| 176 | self.evaluate(variables.variables_initializer(m.variables)) |
| 177 | |
| 178 | # check initial state |
| 179 | self.assertEqual(self.evaluate(m.total), 0) |
| 180 | self.assertEqual(self.evaluate(m.count), 0) |
| 181 | |
| 182 | # check __call__() |
| 183 | self.assertEqual(self.evaluate(m(100)), 100) |
| 184 | self.assertEqual(self.evaluate(m.total), 100) |
| 185 | self.assertEqual(self.evaluate(m.count), 1) |
| 186 | |
| 187 | # check update_state() and result() + state accumulation + tensor input |
| 188 | update_op = m.update_state(ops.convert_n_to_tensor([1, 5])) |
| 189 | self.evaluate(update_op) |
| 190 | self.assertAlmostEqual(self.evaluate(m.result()), 106 / 3, 2) |
| 191 | self.assertEqual(self.evaluate(m.total), 106) # 100 + 1 + 5 |
| 192 | self.assertEqual(self.evaluate(m.count), 3) |
| 193 | |
| 194 | # check reset_states() |
| 195 | m.reset_states() |
| 196 | self.assertEqual(self.evaluate(m.total), 0) |
| 197 | self.assertEqual(self.evaluate(m.count), 0) |
| 198 | |
| 199 | # Check save and restore config |
| 200 | m2 = metrics.Mean.from_config(m.get_config()) |
| 201 | self.assertEqual(m2.name, 'my_mean') |
| 202 | self.assertTrue(m2.stateful) |
| 203 | self.assertEqual(m2.dtype, dtypes.float32) |
| 204 | self.assertEqual(len(m2.variables), 2) |
| 205 | |
| 206 | def test_mean_with_sample_weight(self): |
| 207 | m = metrics.Mean(dtype=dtypes.float64) |
nothing calls this directly
no test coverage detected