Convert a TensorFlow tensor, PyTorch tensor, Numpy array or python list to a Numpy array.
(obj)
| 282 | |
| 283 | |
| 284 | def to_numpy(obj): |
| 285 | """ |
| 286 | Convert a TensorFlow tensor, PyTorch tensor, Numpy array or python list to a Numpy array. |
| 287 | """ |
| 288 | |
| 289 | framework_to_numpy = { |
| 290 | "pt": lambda obj: obj.detach().cpu().numpy(), |
| 291 | "tf": lambda obj: obj.numpy(), |
| 292 | "jax": lambda obj: np.asarray(obj), |
| 293 | "np": lambda obj: obj, |
| 294 | } |
| 295 | |
| 296 | if isinstance(obj, (dict, UserDict)): |
| 297 | return {k: to_numpy(v) for k, v in obj.items()} |
| 298 | elif isinstance(obj, (list, tuple)): |
| 299 | return np.array(obj) |
| 300 | |
| 301 | # This gives us a smart order to test the frameworks with the corresponding tests. |
| 302 | framework_to_test_func = _get_frameworks_and_test_func(obj) |
| 303 | for framework, test_func in framework_to_test_func.items(): |
| 304 | if test_func(obj): |
| 305 | return framework_to_numpy[framework](obj) |
| 306 | |
| 307 | return obj |
| 308 | |
| 309 | |
| 310 | class ModelOutput(OrderedDict): |
no test coverage detected