Python:如何实现绑定device参数的动态扩展调用方式?
动态封装PyTorch函数实现免重复传参及链式调用
需求背景
现有两个PyTorch相关函数:
from torch import device def run_cuda(device: device, count: int): ... def gen_noise(device: device, width: int, height: int): ...
当前调用时必须每次手动传入device参数:
device = DEVICE run_cuda(device, count=8) gen_noise(device, width=128, height=128)
希望实现use_device(device)函数,返回的对象可以直接调用上述函数且无需重复传入device,同时支持链式设置其他参数,示例调用方式如下:
device = DEVICE # 多行调用 device_user = use_device(device) device_user.run_cuda(count=8) device_user.gen_noise(width=128, height=128) # 链式调用 use_device(device).use_dimension(512,512).use_iteration(8).gen_noise()
不想手动封装device类,询问Python是否支持这种动态扩展方法。
实现方案
可以通过Python的__getattr__魔法方法实现动态代理,无需手动封装所有函数,同时支持链式参数设置。具体代码如下:
完整实现代码
from torch import device # 原函数保持不变 def run_cuda(device: device, count: int): print(f"执行run_cuda: 设备={device}, 次数={count}") def gen_noise(device: device, width: int, height: int): print(f"生成噪声: 设备={device}, 尺寸={width}x{height}") class DeviceProxy: def __init__(self, target_device): self._device = target_device self._stored_params = {} # 存储链式设置的参数 def __getattr__(self, attr_name): # 处理目标函数调用 if attr_name in globals() and callable(globals()[attr_name]): target_func = globals()[attr_name] def wrapped_func(**kwargs): # 合并预存参数与调用时传入的参数 combined_params = {**self._stored_params, **kwargs} # 自动注入device参数 combined_params['device'] = self._device return target_func(**combined_params) return wrapped_func # 处理链式参数设置方法(以use_开头) elif attr_name.startswith('use_'): param_key = attr_name[4:] def param_setter(*args): # 针对dimension特殊处理,映射为width和height if param_key == 'dimension' and len(args) == 2: self._stored_params['width'] = args[0] self._stored_params['height'] = args[1] elif len(args) == 1: self._stored_params[param_key] = args[0] return self # 返回自身支持链式调用 return param_setter # 不存在的属性抛出异常 raise AttributeError(f"'DeviceProxy' 对象没有属性 '{attr_name}'") def use_device(target_device): return DeviceProxy(target_device)
测试调用示例
import torch DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 多行调用测试 device_proxy = use_device(DEVICE) device_proxy.run_cuda(count=8) device_proxy.gen_noise(width=128, height=128) # 链式调用测试 use_device(DEVICE).use_dimension(512,512).use_iteration(8).gen_noise()
原理说明
DeviceProxy类通过__getattr__拦截所有属性访问:- 如果访问的是已存在的函数(如
run_cuda、gen_noise),则返回一个包装函数,自动注入device参数,并合并链式设置的参数和调用时的参数。 - 如果访问的是
use_开头的方法,则生成参数设置器,将参数存入内部字典后返回自身,实现链式调用。
- 如果访问的是已存在的函数(如
- 无需手动封装每个目标函数,只要函数存在于全局作用域,就能通过代理对象调用,完全实现动态扩展。
内容的提问来源于stack exchange,提问作者CC-white
相关产品推荐
相关产品推荐

