如何为TensorFlow的map_fn堆叠异尺寸参数,替代multiprocessing的Pool.map
TensorFlow map_fn 多不同形状参数传入方案
tf.map_fn 原生支持多参数并行映射,不需要将不同形状的参数硬堆叠为单个张量,你只需要按参数的位置分别堆叠同位置的所有输入即可。
1. 按参数位置分别构建堆叠张量
你的每一次myfunc调用都接收3个固定结构的参数:第一个是MxN矩阵,第二个是I维向量,第三个是J维向量,同位置的参数形状完全一致,分别堆叠不会有维度冲突:
# 先按参数位置收集所有输入 arg1_list = [arg1] arg2_list = [arg2] arg3_list = [arg3] # 动态添加参数的逻辑不变 if mask[1]: arg1_list.append(arga) arg2_list.append(argb) arg3_list.append(argc) # 分别堆叠同位置的所有参数 stacked_arg1 = tf.stack(arg1_list) # 形状为 [K, M, N],K为调用次数 stacked_arg2 = tf.stack(arg2_list) # 形状为 [K, I] stacked_arg3 = tf.stack(arg3_list) # 形状为 [K, J]
2. 调用tf.map_fn时传入参数元组
将堆叠后的三个张量组成元组传给elems参数,fn会自动按批次接收每组对应位置的参数:
res = tf.map_fn( # 按顺序接收对应位置的单组参数,传入myfunc lambda x: myfunc(x[0], x[1], x[2]), elems=(stacked_arg1, stacked_arg2, stacked_arg3), # 建议显式指定输出的形状和类型,避免自动推导出错 fn_output_signature=tf.TensorSpec(shape=myfunc_output_shape, dtype=myfunc_output_dtype) )
错误方案原因说明
tf.stack要求所有待堆叠的张量形状完全一致,你之前将不同形状的a、b、c放在同一层级堆叠,自然触发维度不匹配报错- RaggedTensor仅适用于维度数相同、仅某一轴长度不同的张量堆叠场景,你的参数维度数都不一致,且
tf.map_fn对嵌套RaggedTensor的支持度极低,自然会抛出类型错误
内容的提问来源于stack exchange,提问作者Patafikss
相关产品推荐
相关产品推荐

