如何在Keras中扩展Ragged Tensor的维度?
解决方案
要给Ragged Tensor添加非不规则维度,无需使用tf.map_fn,直接调用tf.expand_dims指定正确轴即可,Ragged Tensor原生支持在非不规则维度上扩展:
import tensorflow as tf num_features = 4 x = tf.ragged.constant([ [[0,0,0,0],[0,0,0,0],[1,0,0,0],[0,0,0,0]], [[0,0,0,0],[0,0,0,0],[0,0,0,0],[1,0,0,0],[0,0,0,0]]], dtype = tf.float32, inner_shape=(num_features,)) # 在轴1位置扩展维度,得到目标形状[批量大小, 1, 时间步数, 特征数] y = tf.expand_dims(x, axis=1) print(y.shape) # 输出: (2, 1, None, 4)
如果要在Keras层中集成,直接用Lambda层即可:
expand_layer = tf.keras.layers.Lambda(lambda x: tf.expand_dims(x, axis=1)) y = expand_layer(x)
错误原因
你之前用tf.map_fn的问题在于:map_fn会逐个处理Ragged Tensor的样本(每个样本是形状为[时间步数, 特征数]的普通Tensor),扩展后得到[1, 时间步数, 特征数]的Tensor,但map_fn默认尝试将这些结果堆叠为常规Tensor,而不同样本的时间步数不一致,导致无法堆叠成规则张量,也未正确识别输出为Ragged Tensor。直接对整个Ragged Tensor调用tf.expand_dims会保留其不规则结构,因为扩展的是批量维度后的非不规则轴,不会破坏Ragged Tensor的内部结构。
原问题详情
需求目标
使用Keras为Ragged Tensor添加一个非不规则维度:
- 初始Ragged Tensor形状为
[批量大小, 时间步数, 特征数] - 期望最终Ragged Tensor形状为
[批量大小, 1, 时间步数, 特征数]
(此举旨在对每个样本执行时间卷积,若有相关实现方案欢迎分享)
尝试方案
尝试结合tf.map_fn与调用tf.expand_dims的Lambda layer,但出现不规则维度大小不兼容的错误。添加tf.TensorSpec作为fn_output_signature未解决问题,添加tf.RaggedTensorSpec也无效。
复现代码(基于TensorFlow 2.15.0、Python 3.11)
import tensorflow as tf num_features = 4 x = tf.ragged.constant([ [[0,0,0,0],[0,0,0,0],[1,0,0,0],[0,0,0,0]], [[0,0,0,0],[0,0,0,0],[0,0,0,0],[1,0,0,0],[0,0,0,0]]], dtype = tf.float32, inner_shape=(num_features,)) expandDims = tf.keras.layers.Lambda( lambda x: tf.expand_dims(x,axis=0)) # 期望y为形状(2,1,None,4)的Ragged Tensor,执行报错 y = tf.map_fn(expandDims,x) # 同样报错 #y = tf.map_fn(expandDims,x,fn_output_signature = tf.TensorSpec(shape = (1,None,num_features)))
报错信息
2024-01-07 09:52:03.342983: W tensorflow/core/framework/op_kernel.cc:1839] OP_REQUIRES failed at ragged_tensor_from_variant_op.cc:333 : INVALID_ARGUMENT: All flat_values must have compatible shapes. Shape at index 0: [4,4]. Shape at index 1: [5,4]. If you are using tf.map_fn, then you may need to specify an explicit fn_output_signature with appropriate ragged_rank, and/or convert output tensors to RaggedTensors. Traceback (most recent call last): File "c:\Users\apples\Documents\tensforflow probability course\test.py", line 18, in <module> y = tf.map_fn(expandDims,x) ^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\util\deprecation.py", line 660, in new_func return func(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\util\deprecation.py", line 588, in new_func return func(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\map_fn.py", line 637, in map_fn_v2 return map_fn( ^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\util\deprecation.py", line 588, in new_func return func(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\map_fn.py", line 516, in map_fn result_flat = _result_batchable_to_flat(result_batchable, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\map_fn.py", line 607, in _result_batchable_to_flat spec._batch(batch_size)._from_compatible_tensor_list( File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\ragged\ragged_tensor.py", line 2601, in _from_compatible_tensor_list result = RaggedTensor._from_variant( # pylint: disable=protected-access ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\ragged\ragged_tensor.py", line 2028, in _from_variant result = gen_ragged_conversion_ops.ragged_tensor_from_variant( ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\ops\gen_ragged_conversion_ops.py", line 77, in ragged_tensor_from_variant _ops.raise_from_not_ok_status(e, name) File "C:\Program Files\Python\Python311\Lib\site-packages\tensorflow\python\framework\ops.py", line 5883, in raise_from_not_ok_status raise core._status_to_exception(e) from None # pylint: disable=protected-access ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ tensorflow.python.framework.errors_impl.InvalidArgumentError: {{function_node __wrapped__RaggedTensorFromVariant_output_ragged_rank_1_device_/job:localhost/replica:0/task:0/device:CPU:0}} All flat_values must have compatible shapes. Shape at index 0: [4,4]. Shape at index 1: [5,4]. If you are using tf.map_fn, then you may need to specify an explicit fn_output_signature with appropriate ragged_rank, and/or convert output tensors to RaggedTensors. [Op:RaggedTensorFromVariant] name:
补充说明
已放弃将Ragged Tensor作为tf.keras.layers.Conv1D输入的尝试,改为填充Ragged Tensor使其成为常规张量。
内容的提问来源于stack exchange,提问作者Eric Zarahn
相关产品推荐
相关产品推荐

