【问题标题】:Efficient tensor contraction in pythonpython中的高效张量收缩
【发布时间】:2017-02-03 23:11:34
【问题描述】:

我有一个张量列表Lndarray 对象),每个都有几个索引。我需要根据连接图收缩这些指数。

连接被编码在((m,i),(n,j))形式的元组列表中,表示“将张量L[m]i-th索引与j-张量的第 th 索引 L[n].

如何处理非平凡的连接图?第一个问题是,一旦我收缩了一对索引,结果就是一个不属于L 列表的新张量。但即使我解决了这个问题(例如,通过为所有张量的所有索引提供唯一标识符),仍然存在一个问题,即可以选择任何顺序来执行收缩,并且某些选择会在计算中产生不必要的巨大野兽(即使最终结果很小)。有什么建议吗?

【问题讨论】:

    标签: python numpy vectorization numpy-einsum


    【解决方案1】:

    除了内存方面的考虑之外,我相信您可以通过一次调用 einsum 来完成收缩,尽管您需要进行一些预处理。我不完全确定您所说的“当我收缩一对索引时,结果是一个不属于列表 L 的新张量”是什么意思,但我认为进行收缩一步就可以解决这个问题。

    我建议使用 einsum 的替代数字索引语法:

    einsum(op0, sublist0, op1, sublist1, ..., [sublistout])
    

    所以你需要做的是将索引编码为整数序列。首先,您首先需要设置一系列唯一索引,并保留另一个副本用作sublistout。然后,遍历您的连接图,您需要在必要时将收缩索引设置为相同的索引,同时从sublistout 中删除收缩索引。

    import numpy as np
    
    def contract_all(tensors,conns):
        '''
        Contract the tensors inside the list tensors
        according to the connectivities in conns
    
        Example input:
        tensors = [np.random.rand(2,3),np.random.rand(3,4,5),np.random.rand(3,4)]
        conns = [((0,1),(2,0)), ((1,1),(2,1))]
        returned shape in this case is (2,3,5)
        '''
    
        ndims = [t.ndim for t in tensors]
        totdims = sum(ndims)
        dims0 = np.arange(totdims)
        # keep track of sublistout throughout
        sublistout = set(dims0.tolist())
        # cut up the index array according to tensors
        # (throw away empty list at the end)
        inds = np.split(dims0,np.cumsum(ndims))[:-1]
        # we also need to convert to a list, otherwise einsum chokes
        inds = [ind.tolist() for ind in inds]
    
        # if there were no contractions, we'd call
        # np.einsum(*zip(tensors,inds),sublistout)
    
        # instead we need to loop over the connectivity graph
        # and manipulate the indices
        for (m,i),(n,j) in conns:
            # tensors[m][i] contracted with tensors[n][j]
    
            # remove the old indices from sublistout which is a set
            sublistout -= {inds[m][i],inds[n][j]}
    
            # contract the indices
            inds[n][j] = inds[m][i]
    
        # zip and flatten the tensors and indices
        args = [subarg for arg in zip(tensors,inds) for subarg in arg]
    
        # assuming there are no multiple contractions, we're done here
        return np.einsum(*args,sublistout)
    

    一个简单的例子:

    >>> tensors = [np.random.rand(2,3), np.random.rand(4,3)]
    >>> conns = [((0,1),(1,1))]
    >>> contract_all(tensors,conns)
    array([[ 1.51970003,  1.06482209,  1.61478989,  1.86329518],
           [ 1.16334367,  0.60125945,  1.00275992,  1.43578448]])
    >>> np.einsum('ij,kj',tensors[0],tensors[1])
    array([[ 1.51970003,  1.06482209,  1.61478989,  1.86329518],
           [ 1.16334367,  0.60125945,  1.00275992,  1.43578448]])
    

    如果有多个收缩,循环中的物流会变得有点复杂,因为我们需要处理所有重复。然而,逻辑是相同的。此外,上述显然缺少检查以确保相应的索引可以收缩。

    事后我意识到不必指定默认的sublistouteinsum 无论如何都会使用该顺序。我决定将这个变量留在代码中,因为如果我们想要一个重要的输出索引顺序,我们必须适当地处理这个变量,它可能会派上用场。


    至于收缩顺序的优化,您可以在 1.12 版的 np.einsum 中进行内部优化(正如 @hpaulj 在现已删除的评论中所指出的那样)。此版本向np.einsum 引入了optimize 可选关键字参数,允许选择以内存为代价减少计算时间的收缩顺序。传递 'greedy''optimal' 作为 optimize 关键字将使 numpy 选择按尺寸大小大致递减顺序的收缩顺序。

    optimize 关键字可用的选项来自显然未记录的(就在线文档而言;help() 幸运的是)函数np.einsum_path

    einsum_path(subscripts, *operands, optimize='greedy')
    
    Evaluates the lowest cost contraction order for an einsum expression by
    considering the creation of intermediate arrays.
    

    来自np.einsum_path 的输出收缩路径也可以用作np.einsumoptimize 参数的输入。在您的问题中,您担心使用了过多的内存,因此我怀疑默认情况下没有优化(可能会更长的运行时间和更小的内存占用)。

    【讨论】:

    • 在最近的一个 SO 问题中,我发现 optimize='optimal' 让我计算一个更大的 einsum 数组,stackoverflow.com/questions/41942115/…。到目前为止,这是我对新功能的唯一体验。
    • 太好了,我周末试试这个!顺便说一句,今天我写了一个计算最佳(内存方面)收缩顺序的函数,所以已经涵盖了。
    • 因此,如果我理解正确,我可以告诉 einsum 执行宫缩的顺序。为此,我可以使用 np.einsum_path 的输出,或者我可以让我自己编写的优化函数的输出具有相同的格式。我会比较一下。
    • 哇!使用optimization=True,我在第一个实际示例中得到了大约 1000 倍的改进!
    • 是的,我的意思是optimize=True,我不能再编辑评论了。顺便说一句,非常感谢!
    【解决方案2】:

    也许有帮助:看看https://arxiv.org/abs/1402.0939,它描述了一个有效框架,用于解决在单个函数ncon(...) 中收缩所谓的张量网络的问题。据我所知,它的实现可直接用于 Matlab(可在链接中找到)和 Python3 (https://github.com/mhauru/ncon)。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2016-10-29
      • 1970-01-01
      • 1970-01-01
      • 2016-06-20
      • 2011-11-08
      • 2018-10-15
      • 1970-01-01
      相关资源
      最近更新 更多