tf.keras.layers.CategoryEncoding的multi_hot输出模式行为与定义问询
问题
请帮忙理解tf.keras.layers.CategoryEncoding中multi hot encoding的定义,以及output_mode='multi_hot'的运行行为。
背景
参考通用multi-hot编码的定义:
如果使用multi-hot编码,首先会对类别做标签编码,用单个数字代表对应类别(例如1代表'dog'),随后将数值标签转换为长度为log2(5)=3的二进制向量。
示例如下:'cat' = [0,0,0] 'dog' = [0,0,1] 'fish' = [0,1,0] 'bird' = [0,1,1] 'ant' = [1,0,0]
tf.keras.layers.CategoryEncoding的行为说明
官方文档说明num_tokens是该层支持的token总数量。
参数说明
num_tokens
该层支持的token总数量,输入的所有整数必须满足0 <= value < num_tokens,否则会抛出错误。
output_mode
- "one_hot":将输入的每个独立元素编码为长度等于num_tokens的数组,元素对应索引位置为1。如果最后一维大小为1,则在该维度上编码;如果最后一维大小不为1,会新增一个维度存放编码输出。
- "multi_hot":将输入的每个样本编码为单个长度等于num_tokens的数组,样本中出现的每个词汇项对应位置为1。将最后一维视为样本维度,若输入形状为(..., sample_length),则输出形状为(..., num_tokens)。
根据上述multi hot编码的通用定义,我原本预期tf.keras.layers.CategoryEncoding(num_tokens=5, output_mode="multi_hot")会将5个token编码为长度为3的数组。
但官方文档说明"multi_hot"将每个样本编码为单个长度等于num_tokens的数组,样本中出现的每个词汇项对应位置为1,实际运行效果也符合该描述:
dataset = tf.data.Dataset.from_tensor_slices(tf.constant(['cat', 'dog', 'fish', 'bird'])) lookup = tf.keras.layers.StringLookup(max_tokens=5, oov_token='[UNK]') lookup.adapt(dataset) lookup.get_vocabulary() --- ['[UNK]', 'fish', 'dog', 'cat', 'bird'] mhe = tf.keras.layers.CategoryEncoding(num_tokens=lookup.vocabulary_size(), output_mode="multi_hot") print(f"cat: {mhe(lookup(tf.constant('cat'))).numpy()}") print(f"dog: {mhe(lookup(tf.constant('dog'))).numpy()}") --- cat: [0. 0. 0. 1. 0.] dog: [0. 0. 1. 0. 0.]
单类别输入下该结果和One Hot Encoding没有差异:
ohe = tf.keras.layers.CategoryEncoding(num_tokens=lookup.vocabulary_size(), output_mode="one_hot") print(f"cat: {ohe(lookup(tf.constant('cat'))).numpy()}") print(f"dog: {ohe(lookup(tf.constant('dog'))).numpy()}") --- cat: [0. 0. 0. 1. 0.] dog: [0. 0. 1. 0. 0.]
多值输入场景下,multi_hot会将所有输入值合并编码:
print(ohe(lookup(tf.constant(['cat', 'dog']))).numpy()) --- [[0. 0. 0. 1. 0.] [0. 0. 1. 0. 0.]] print(mhe(lookup(tf.constant(['cat', 'dog']))).numpy()) --- [0. 0. 1. 1. 0.]
如果要对多个样本分别做multi_hot编码,需要输入二维数组:
print(mhe(lookup(tf.constant([['cat'], ['dog']]))).numpy()) --- [[0. 0. 0. 1. 0.] [0. 0. 1. 0. 0.]]
显然tf.keras.layers.CategoryEncoding定义的multi hot encoding和通用的multi-hot编码定义并不一致。
相关参考
- TensorFlow官方issue编号:#52892
内容的提问来源于stack exchange,提问作者mon
相关产品推荐
相关产品推荐

