Truncate `tensors` at `limit` (even if it's a nested list/tuple/dict of tensors).
(tensors, limit)
| 96 | |
| 97 | |
| 98 | def nested_truncate(tensors, limit): |
| 99 | "Truncate `tensors` at `limit` (even if it's a nested list/tuple/dict of tensors)." |
| 100 | if isinstance(tensors, (list, tuple)): |
| 101 | return type(tensors)(nested_truncate(t, limit) for t in tensors) |
| 102 | elif isinstance(tensors, dict): |
| 103 | return type(tensors)({k: nested_truncate(v, limit) for k, v in tensors.items()}) |
| 104 | elif tensors is None: |
| 105 | return None |
| 106 | return tensors[:limit] |
| 107 | |
| 108 | |
| 109 | def _secs2timedelta(secs): |