| 92 | set_determinism(None) |
| 93 | |
| 94 | def check_match(self, in1, in2): |
| 95 | if isinstance(in1, dict): |
| 96 | self.assertTrue(isinstance(in2, dict)) |
| 97 | for (k1, v1), (k2, v2) in zip(in1.items(), in2.items()): |
| 98 | if isinstance(k1, Enum) and isinstance(k2, Enum): |
| 99 | k1, k2 = k1.value, k2.value |
| 100 | self.check_match(k1, k2) |
| 101 | # Transform ids won't match for windows with multiprocessing, so don't check values |
| 102 | if k1 == TraceKeys.ID and sys.platform in ["darwin", "win32"]: |
| 103 | continue |
| 104 | if not (isinstance(k1, str) and k1.endswith("_transforms")): |
| 105 | self.check_match(v1, v2) # transform stack not necessarily match |
| 106 | elif isinstance(in1, (list, tuple)): |
| 107 | for l1, l2 in zip(in1, in2): |
| 108 | self.check_match(l1, l2) |
| 109 | elif isinstance(in1, (str, int)): |
| 110 | self.assertEqual(in1, in2) |
| 111 | elif isinstance(in1, (torch.Tensor, np.ndarray)): |
| 112 | np.testing.assert_array_equal(in1, in2) |
| 113 | else: |
| 114 | raise RuntimeError(f"Not sure how to compare types. type(in1): {type(in1)}, type(in2): {type(in2)}") |
| 115 | |
| 116 | def check_decollate(self, dataset): |
| 117 | batch_size = 2 |