Simplifies an object by rounding float numbers, and downcasting tensors/numpy arrays to get simple equality test within tests.
(obj, decimals=3)
| 2084 | |
| 2085 | |
| 2086 | def nested_simplify(obj, decimals=3): |
| 2087 | """ |
| 2088 | Simplifies an object by rounding float numbers, and downcasting tensors/numpy arrays to get simple equality test |
| 2089 | within tests. |
| 2090 | """ |
| 2091 | import numpy as np |
| 2092 | |
| 2093 | if isinstance(obj, list): |
| 2094 | return [nested_simplify(item, decimals) for item in obj] |
| 2095 | if isinstance(obj, tuple): |
| 2096 | return tuple([nested_simplify(item, decimals) for item in obj]) |
| 2097 | elif isinstance(obj, np.ndarray): |
| 2098 | return nested_simplify(obj.tolist()) |
| 2099 | elif isinstance(obj, Mapping): |
| 2100 | return {nested_simplify(k, decimals): nested_simplify(v, decimals) for k, v in obj.items()} |
| 2101 | elif isinstance(obj, (str, int, np.int64)): |
| 2102 | return obj |
| 2103 | elif obj is None: |
| 2104 | return obj |
| 2105 | elif is_torch_available() and isinstance(obj, torch.Tensor): |
| 2106 | return nested_simplify(obj.tolist(), decimals) |
| 2107 | elif is_tf_available() and tf.is_tensor(obj): |
| 2108 | return nested_simplify(obj.numpy().tolist()) |
| 2109 | elif isinstance(obj, float): |
| 2110 | return round(obj, decimals) |
| 2111 | elif isinstance(obj, (np.int32, np.float32)): |
| 2112 | return nested_simplify(obj.item(), decimals) |
| 2113 | else: |
| 2114 | raise Exception(f"Not supported: {type(obj)}") |
| 2115 | |
| 2116 | |
| 2117 | def check_json_file_has_correct_format(file_path): |
no test coverage detected