You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.28 09:54:01