如何使用tf.math.segment_sum创建Keras层并解决秩不匹配报错
问题原因
- 维度不符合OP要求:
tf.keras.Input(shape=(1,))生成的是二维张量,形状为(batch_size, 1),但tf.math.segment_sum的入参规则为:- 待分段求和的数据张量第一维为样本序列长度
- 分段ID参数
segment_ids必须是一维张量,长度和待求和数据的第一维长度完全相等
你传入的待求和值、分段ID都是形状[?,1]的二维张量,直接触发秩不匹配报错。
- 分段ID不符合OP规则:
tf.math.segment_sum要求分段ID必须从0开始、连续递增,你示例中的ID从1开始,就算维度匹配也会生成多余的全0空分段。 - 额外注意:原生
tf.math.segment_sum不识别batch维度,会把整个输入batch的所有样本当成同一段序列做分段计算,多batch训练时会出现不同样本数据混淆的问题。
修正方案
先通过tf.squeeze去掉两个输入张量长度为1的末尾维度,把二维张量转为符合要求的一维张量,同时把ID偏移1位转为从0开始编号。修正后可直接跑通你示例场景的代码如下:
import tensorflow as tf import pandas as pd df = pd.DataFrame({'id': [1, 1, 2, 2, 3, 3, 4, 4], 'x_1': [1, 0, 0, 0, 0, 1, 1, 1], 'target': [1, 1, 0, 0, 1, 1, 2, 2]}) input_ = tf.keras.Input((1,), name='X', dtype=tf.float32) cid = tf.keras.Input(shape=(1,), dtype='int64', name='id') summed = tf.keras.layers.Lambda( lambda x: tf.math.segment_sum( tf.squeeze(x[0], axis=-1), tf.squeeze(x[1], axis=-1) - 1 ), name='segment_sum' )([input_, cid]) model = tf.keras.Model(inputs=[input_, cid], outputs=[summed]) # 验证结果 print(model.predict([ df['x_1'].values.reshape(-1, 1), df['id'].values.reshape(-1, 1) ])) # 输出 [1. 0. 1. 2.],和按id分组求和的预期结果完全一致
多batch训练适配说明
如果需要在训练时按batch独立计算(不同batch的分段逻辑互不干扰),不要直接使用上述基础实现,需要先给每个batch内的分段ID加上对应batch的偏移量,避免跨batch的分段混淆;如果你的分段ID不是连续从0开始的,替换为tf.math.unsorted_segment_sum即可适配非连续ID场景。
内容的提问来源于stack exchange,提问作者Frits
相关产品推荐
相关产品推荐

