| 39 | |
| 40 | |
| 41 | class TestApply(unittest.TestCase): |
| 42 | def _test_apply_impl(self, tensor, pending_transforms, expected_shape): |
| 43 | result = apply_pending(tensor, pending_transforms) |
| 44 | self.assertListEqual(result[1], pending_transforms) |
| 45 | self.assertEqual(result[0].shape, expected_shape) |
| 46 | |
| 47 | def _test_apply_metatensor_impl(self, tensor, pending_transforms, expected_shape, pending_as_parameter): |
| 48 | tensor_ = convert_to_tensor(tensor, track_meta=True) |
| 49 | if pending_as_parameter: |
| 50 | result, transforms = apply_pending(tensor_, pending_transforms) |
| 51 | else: |
| 52 | for p in pending_transforms: |
| 53 | tensor_.push_pending_operation(p) |
| 54 | if not isinstance(p, dict): |
| 55 | return |
| 56 | result, transforms = apply_pending(tensor_) |
| 57 | self.assertEqual(result.shape, expected_shape) |
| 58 | |
| 59 | SINGLE_TRANSFORM_CASES = single_2d_transform_cases() |
| 60 | |
| 61 | def test_apply_single_transform(self): |
| 62 | for case in self.SINGLE_TRANSFORM_CASES: |
| 63 | self._test_apply_impl(*case) |
| 64 | |
| 65 | def test_apply_single_transform_metatensor(self): |
| 66 | for case in self.SINGLE_TRANSFORM_CASES: |
| 67 | self._test_apply_metatensor_impl(*case, False) |
| 68 | |
| 69 | def test_apply_single_transform_metatensor_override(self): |
| 70 | for case in self.SINGLE_TRANSFORM_CASES: |
| 71 | self._test_apply_metatensor_impl(*case, True) |
| 72 | |
| 73 | |
| 74 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected
searching dependent graphs…