(self)
| 338 | class KerasAccuracyTest(test.TestCase): |
| 339 | |
| 340 | def test_accuracy(self): |
| 341 | acc_obj = metrics.Accuracy(name='my_acc') |
| 342 | |
| 343 | # check config |
| 344 | self.assertEqual(acc_obj.name, 'my_acc') |
| 345 | self.assertTrue(acc_obj.stateful) |
| 346 | self.assertEqual(len(acc_obj.variables), 2) |
| 347 | self.assertEqual(acc_obj.dtype, dtypes.float32) |
| 348 | self.evaluate(variables.variables_initializer(acc_obj.variables)) |
| 349 | |
| 350 | # verify that correct value is returned |
| 351 | update_op = acc_obj.update_state([[1], [2], [3], [4]], [[1], [2], [3], [4]]) |
| 352 | self.evaluate(update_op) |
| 353 | result = self.evaluate(acc_obj.result()) |
| 354 | self.assertEqual(result, 1) # 2/2 |
| 355 | |
| 356 | # Check save and restore config |
| 357 | a2 = metrics.Accuracy.from_config(acc_obj.get_config()) |
| 358 | self.assertEqual(a2.name, 'my_acc') |
| 359 | self.assertTrue(a2.stateful) |
| 360 | self.assertEqual(len(a2.variables), 2) |
| 361 | self.assertEqual(a2.dtype, dtypes.float32) |
| 362 | |
| 363 | # check with sample_weight |
| 364 | result_t = acc_obj([[2], [1]], [[2], [0]], sample_weight=[[0.5], [0.2]]) |
| 365 | result = self.evaluate(result_t) |
| 366 | self.assertAlmostEqual(result, 0.96, 2) # 4.5/4.7 |
| 367 | |
| 368 | def test_accuracy_ragged(self): |
| 369 | acc_obj = metrics.Accuracy(name='my_acc') |
nothing calls this directly
no test coverage detected