Multidimensional iterator in TensorFlow

Viewed 99

I want to sum up the return value of a vectorized function evaluated on a hypergrid:

    l = tf.reshape(tf.linspace(-b, b, n), (n, 1))
    dims = [l] * d
    mesh = tf.meshgrid(*dims)
    y = f(mesh) 
    result = tf.reduce_sum(y)

Unfortunately, mesh becomes so large that it does not fit into the VRAM when calling tf.meshgrid for a high-dimensional input.

Therefore I am looking for a solution similar to np.ndindex that would allow me to generate sub-meshes in TensorFlow. I do not want to work with loops since d varies. Also, I am not sure if it is viable to deal with recursion in Tensorflow 1.15 . Thanks in advance!

1 Answers

Here is a possible solution I have worked out, not sure if it would be actually good for your case or not. The idea is to recursively subdivide the desired mesh according to a given size. Interestingly, this does not work in TensorFlow 2.x, because it tries to make the recursive function into a graph, which would result in an infinite graph - not sure if there might be a work around that. The solution is obviously much slower than doing the computation on the whole grid directly, but the result is technically the same. The problem is the error. If the grid is really big, then it is likely that a reduction over its entirety will have a significant amount of error. That happens in the first place without doing this subdivision, though, and in fact subdividing seems to reduce the error, if anything (at least in some experiments I did).

Anyway, the code ended up a bit long, although conceptually is not too complicated, I hope the comments make it mostly clear.

import tensorflow as tf

def block_reduction_func(block):
    # This function computes the reduction of a block
    return tf.math.reduce_sum(block)

def intermediate_reduction_func(values):
    # This function computes the reduction of an
    # array of intermediate reduction results
    # (in this example is the same)
    return block_reduction_func(values)

def make_block(aa):
    # Makes an actual block from some space slices
    return tf.stack(tf.meshgrid(*aa), axis=-1)

def get_block_slices(aa, i, size):
    # Selects the space slices corresponding to a particular block
    aa2 = []
    for dim, a in enumerate(aa):
        # Number of slices in this level for this dimension
        s = tf.size(a)
        n = s // size
        n += tf.dtypes.cast(s % size > 0, n.dtype)
        # Select dimension slice
        j = i % n
        aa2.append(aa[dim][j * size:(j + 1) * size])
        i //= n
    return aa2

def by_blocks(aa, blocks):
    # Reduces a space by blocks
    if not blocks:
        # When there are no more subdivisions to do
        # just reduce the current block
        res = block_reduction_func(make_block(aa))
        with tf.control_dependencies([]): #([tf.print(res, aa)]):
            return res + 0
    else:
        # Get current division size
        size, *blocks = blocks
        # Get number of blocks in this recursion level
        num_blocks = 1
        for a in aa:
            s = tf.size(a)
            n = s // size
            n += tf.dtypes.cast(s % size > 0, n.dtype)
            num_blocks *= n
        # Array for intermediate results
        ta = tf.TensorArray(aa[0].dtype, num_blocks, element_shape=())
        # Loop through blocks
        _, ta = tf.while_loop(
            lambda i, ta: i < num_blocks,
            lambda i, ta: (i + 1,
                           ta.write(i, by_blocks(get_block_slices(aa, i, size), blocks))),
            [0, ta], parallel_iterations=1)
        # Reduce intermediate results
        values = ta.stack()
        return intermediate_reduction_func(values)

# Test
b = 1.0
n = 100
d = 3
# Recursive divisions of n (can have arbitrary size)
# Divide in blocks of 60, then blocks of 12
blocks = [60, 12]
with tf.Graph().as_default(), tf.Session():
    # Using positive values only in this example
    # so the errors do not overtake the result
    a = tf.linspace(0., b, n)
    aa = [a] * d
    r1 = block_reduction_func(make_block(aa))
    r2 = by_blocks(aa, blocks)
    # Check results (should be 1500000)
    print(r1.eval())
    # 1499943.6
    print(r2.eval())
    # 1499998.8

    # CPU timings
    %timeit r1.eval()
    # 99.3 µs ± 169 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
    %timeit r2.eval()
    # 96.8 ms ± 170 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)

    # GPU timings
    %timeit r1.eval()
    # 195 µs ± 615 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
    %timeit r2.eval()
    # 316 ms ± 1.54 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
Related