| 55 | m.apply_to_modules(assert_training) |
| 56 | |
| 57 | def test_module_attributes(self): |
| 58 | class Model(nn.Module): |
| 59 | def __init__(self): |
| 60 | super().__init__() |
| 61 | self.val = None |
| 62 | self.initialize() |
| 63 | |
| 64 | def initialize(self): |
| 65 | self.val = mx.array(1.0) |
| 66 | |
| 67 | model = Model() |
| 68 | self.assertTrue(mx.array_equal(model.val, mx.array(1.0))) |
| 69 | |
| 70 | model.val = None |
| 71 | self.assertEqual(model.val, None) |
| 72 | |
| 73 | model.val = mx.array([3]) |
| 74 | self.assertEqual(model.val.item(), 3) |
| 75 | |
| 76 | def test_model_with_dict(self): |
| 77 | class DictModule(nn.Module): |