(*args, device=torch.device('cpu'), **kwargs)
| 64 | def decorator(func): |
| 65 | @wraps(func) |
| 66 | def wrapper(*args, device=torch.device('cpu'), **kwargs): |
| 67 | args = list(args) |
| 68 | # get arguments dimensions |
| 69 | args_order, kwargs_order = get_args_order(func, args, kwargs) |
| 70 | args_dim = [dims[i] for i in args_order] |
| 71 | kwargs_dim = {key: dims[i] for key, i in kwargs_order.items()} |
| 72 | # convert to torch tensor |
| 73 | device = get_device(args, kwargs) or device |
| 74 | for i, arg in enumerate(args): |
| 75 | if isinstance(arg, (Number, list, tuple)) and args_dim[i] is not None: |
| 76 | args[i] = torch.tensor(arg, device=device) |
| 77 | for key, arg in kwargs.items(): |
| 78 | if isinstance(arg, (Number, list, tuple)) and kwargs_dim[key] is not None: |
| 79 | kwargs[key] = torch.tensor(arg, device=device) |
| 80 | # broadcast arguments |
| 81 | args, kwargs, spatial = broadcast_args(args, kwargs, args_dim, kwargs_dim) |
| 82 | for i, (arg, arg_dim) in enumerate(zip(args, args_dim)): |
| 83 | if isinstance(arg, torch.Tensor) and arg_dim is not None: |
| 84 | args[i] = arg.reshape([-1, *arg.shape[arg.ndim-arg_dim:]]) |
| 85 | for key, arg in kwargs.items(): |
| 86 | if isinstance(arg, torch.Tensor) and kwargs_dim[key] is not None: |
| 87 | kwargs[key] = arg.reshape([-1, *arg.shape[arg.ndim-kwargs_dim[key]:]]) |
| 88 | # call function |
| 89 | results = func(*args, **kwargs) |
| 90 | type_results = type(results) |
| 91 | results = list(results) if isinstance(results, (tuple, list)) else [results] |
| 92 | # restore spatial dimensions |
| 93 | for i, result in enumerate(results): |
| 94 | results[i] = result.reshape([*spatial, *result.shape[1:]]) |
| 95 | if type_results == tuple: |
| 96 | results = tuple(results) |
| 97 | elif type_results == list: |
| 98 | results = list(results) |
| 99 | else: |
| 100 | results = results[0] |
| 101 | return results |
| 102 | return wrapper |
| 103 | return decorator |
nothing calls this directly
no test coverage detected