mirror of
https://github.com/facebookresearch/ReAgent.git
synced 2026-06-16 12:44:41 +00:00
Summary: Need more tests before landing the refactor diffs: D22702504 (https://github.com/facebookresearch/ReAgent/commit/1b470c489d19c33beab88b8ea2e79843d4d31f28), D23123762 (https://github.com/facebookresearch/ReAgent/commit/76829287265bc39f879f3bc1d946a1374c5e1141), D23124179 (https://github.com/facebookresearch/ReAgent/commit/b28f84aa013be00194508f52498160592cb37e9d), D23219012 (https://github.com/facebookresearch/ReAgent/commit/e404c5772ea4118105c2eb136ca96ad5ca8e01db) Back out to a version based on D23155753. Check our team diff history: https://fburl.com/diffs/ppsgazgj Reviewed By: kittipatv Differential Revision: D23270626 fbshipit-source-id: 14653066bb3924a987a54650a51241895b321c8e
50 lines
1.4 KiB
Python
50 lines
1.4 KiB
Python
#!/usr/bin/env python3
|
|
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
|
|
|
|
|
import unittest
|
|
|
|
from reagent.core.observers import ValueListObserver
|
|
from reagent.core.tracker import observable
|
|
|
|
|
|
class TestObservable(unittest.TestCase):
|
|
def test_observable(self):
|
|
@observable(td_loss=float, str_val=str)
|
|
class DummyClass:
|
|
def __init__(self, a, b, c=10):
|
|
super().__init__()
|
|
self.a = a
|
|
self.b = b
|
|
self.c = c
|
|
|
|
def do_something(self, i):
|
|
self.notify_observers(td_loss=i, str_val="not_used")
|
|
|
|
instance = DummyClass(1, 2)
|
|
self.assertIsInstance(instance, DummyClass)
|
|
self.assertEqual(instance.a, 1)
|
|
self.assertEqual(instance.b, 2)
|
|
self.assertEqual(instance.c, 10)
|
|
|
|
observers = [ValueListObserver("td_loss") for _i in range(3)]
|
|
instance.add_observers(observers)
|
|
# Adding twice should not result in double update
|
|
instance.add_observer(observers[0])
|
|
|
|
for i in range(10):
|
|
instance.do_something(float(i))
|
|
|
|
for observer in observers:
|
|
self.assertEqual(observer.values, [float(i) for i in range(10)])
|
|
|
|
def test_no_observable_values(self):
|
|
try:
|
|
|
|
@observable()
|
|
class NoObservableValues:
|
|
pass
|
|
|
|
except AssertionError:
|
|
pass
|