You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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__拦截所有属性访问:
    1. 如果访问的是已存在的函数(如run_cuda、gen_noise),则返回一个包装函数,自动注入device参数,并合并链式设置的参数和调用时的参数。
    2. 如果访问的是use_开头的方法,则生成参数设置器,将参数存入内部字典后返回自身,实现链式调用。
  • 无需手动封装每个目标函数,只要函数存在于全局作用域,就能通过代理对象调用,完全实现动态扩展。

内容的提问来源于stack exchange,提问作者CC-white

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 05:35:22