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

PyTorch中是否可将多个变换函数堆叠为单一函数?

在PyTorch中实现多变换函数的"堆叠"调用

首先明确:torch.stack是专门用于堆叠张量的API,无法直接堆叠函数。不过可以通过自定义逻辑实现你想要的效果——让多个变换函数按对应位置处理输入张量的元素,输出对应结果。

方法一:自定义可调用类

这个类会存储你要堆叠的所有变换函数,调用时自动将输入张量的每个元素匹配到对应函数处理:

import torch

class FuncStack:
    def __init__(self, funcs):
        self.funcs = funcs
        
    def __call__(self, x):
        # 检查输入长度和函数数量是否匹配
        assert len(x) == len(self.funcs), "输入张量的长度必须和函数数量一致"
        # 逐个元素应用对应函数
        results = [func(elem) for elem, func in zip(x, self.funcs)]
        return torch.tensor(results)

# 测试示例
f_stack = FuncStack([lambda x: x+1, lambda x: x*2, lambda x: x*x])

print(f_stack(torch.tensor([0, 0, 0])))
# >>> tensor([1, 0, 0])

print(f_stack(torch.tensor([3, 3, 3])))
# >>> tensor([4, 6, 9])

print(f_stack(torch.tensor([3, 0, 0])))
# >>> tensor([4, 0, 0])

方法二:高效向量化实现

如果你的变换函数都是可向量化的操作,这种方式避免了显式循环,处理大张量时效率更高:

import torch

def create_func_stack(funcs):
    def apply_funcs(x):
        assert len(x) == len(funcs), "输入张量的长度必须和函数数量一致"
        res = x.clone()
        for idx, func in enumerate(funcs):
            res[idx] = func(res[idx])
        return res
    return apply_funcs

# 测试示例
f_stack = create_func_stack([lambda x: x+1, lambda x: x*2, lambda x: x*x])

print(f_stack(torch.tensor([0, 0, 0])))
# >>> tensor([1, 0, 0])

print(f_stack(torch.tensor([3, 3, 3])))
# >>> tensor([4, 6, 9])

print(f_stack(torch.tensor([3, 0, 0])))
# >>> tensor([4, 0, 0])

额外场景说明

如果你的需求是让每个函数作用于整个输入张量,再把结果堆叠成一个多维张量(比如输入[3,3,3],输出[[4,4,4],[6,6,6],[9,9,9]]),可以用更简单的实现:

def stack_funcs(funcs):
    def apply(x):
        return torch.stack([func(x) for func in funcs])
    return apply

# 测试
f_stack = stack_funcs([lambda x: x+1, lambda x: x*2, lambda x: x*x])
print(f_stack(torch.tensor([3,3,3])))
# >>> tensor([[4, 4, 4],
#             [6, 6, 6],
#             [9, 9, 9]])

内容的提问来源于stack exchange,提问作者Jake Levi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 18:40:30