Recursively moves all tensors in obj to the specified device. Args: obj: Object to move to device - can be a tensor, list, tuple, dict or any nested combination device: Target device (e.g. 'cuda', 'cpu', torch.device('cuda:0') etc.) Returns: Same object structure wi
(obj, device)
| 66 | |
| 67 | |
| 68 | def to_device(obj, device): |
| 69 | """Recursively moves all tensors in obj to the specified device. |
| 70 | |
| 71 | Args: |
| 72 | obj: Object to move to device - can be a tensor, list, tuple, dict or any nested combination |
| 73 | device: Target device (e.g. 'cuda', 'cpu', torch.device('cuda:0') etc.) |
| 74 | |
| 75 | Returns: |
| 76 | Same object structure with all contained tensors moved to specified device |
| 77 | """ |
| 78 | to_fn = lambda x: x.to(device) |
| 79 | return optree.tree_map(to_fn, obj, is_leaf=torch.is_tensor, none_is_leaf=False) |
| 80 | |
| 81 | |
| 82 | def expand_right(tensor, target_shape): |