MCPcopy Create free account
hub / github.com/InternRobotics/G2VLM / wrapper

Function wrapper

eval_code/recons/models/moge/utils3d/torch/_helpers.py:66–101  ·  view source on GitHub ↗
(*args, device=torch.device('cpu'), **kwargs)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 4

get_args_orderFunction · 0.70
get_deviceFunction · 0.70
broadcast_argsFunction · 0.70
deviceMethod · 0.45

Tested by

no test coverage detected