(tensors_list)
| 450 | |
| 451 | |
| 452 | def _as_tensor_list_list(tensors_list): |
| 453 | if not tensors_list: |
| 454 | raise ValueError("Expected at least one set of tensors") |
| 455 | if isinstance(tensors_list[0], dict): |
| 456 | expected_keys = set(tensors_list[0].keys()) |
| 457 | for tensors in tensors_list[1:]: |
| 458 | if set(tensors.keys()) != expected_keys: |
| 459 | raise ValueError("All dictionaries in tensors_list must have " |
| 460 | "the same keys") |
| 461 | return [_as_tensor_list(tensors) for tensors in tensors_list] |
| 462 | else: |
| 463 | return tensors_list |
| 464 | |
| 465 | |
| 466 | def _as_original_type(original_tensors, tensor_list): |
no test coverage detected