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

如何向TensorFlow的MMOE模型添加TFRecords中的独热编码特征x

如何在MMOE模型中引入并处理TFRecords中的int型one-hot特征x

1. 修改TFRecords解析逻辑

不管你用tf.data.TFRecordDataset还是TensorFlow Recommenders的数据集类,都要在特征解析描述里加入对x的定义。示例代码如下:

# 定义所有特征的解析规则,新增x的解析
feature_description = {
    # 原有特征保留
    'anc_pcat_map': tf.io.FixedLenFeature([], tf.int32),
    'anc_ptcode_map': tf.io.FixedLenFeature([], tf.int32),
    # 新增int型特征x的解析,shape根据实际情况调整(比如多维度写[2])
    'x': tf.io.FixedLenFeature([], tf.int32)
}

def parse_example(example_proto):
    return tf.io.parse_single_example(example_proto, feature_description)

# 应用解析函数到数据集
dataset = tf.data.TFRecordDataset('your_tfrecords_path.tfrecord').map(parse_example)

如果用tfrs的DatasetBuilder子类,直接在_parse_example方法里添加对应解析逻辑即可。

2. 在模型__init__中添加one-hot编码层

首先确定特征x的总类别数(比如x的取值范围是0~N-1,或者统计过所有可能的取值数量),然后在模型初始化时定义编码层:

class TIGMMOE(tfrs.Model):
    def __init__(self, use_cross_layer, deep_layer_sizes, num_units, num_shared_experts, num_x_classes, projection_dim=None):
        super().__init__()
        
        self.embedding_dimension = 8
        self._embeddings = {}
    
        # 原有embedding层保留
        self._embeddings['category'] = tf.keras.Sequential(
          [tf.keras.layers.Embedding(num_total_pcats + 1, 32)
          ],name='cat_emb')
    
        self._embeddings['ptype'] = tf.keras.Sequential(
          [tf.keras.layers.Embedding(num_total_ptypes + 1, 128)
          ],name='ptype_emb')
        
        # 新增x的one-hot编码层
        # 若x是连续整数取值,直接用CategoryEncoding
        self.x_one_hot = tf.keras.layers.CategoryEncoding(
            num_tokens=num_x_classes + 1,  # +1用于处理未知值(可选)
            output_mode='one_hot',
            name='x_one_hot'
        )
        # 若x是非连续取值,先加Lookup层映射到连续索引:
        # self.x_lookup = tf.keras.layers.IntegerLookup(vocabulary=your_x_vocab_list, mask_token=None)
        # self.x_one_hot = tf.keras.layers.CategoryEncoding(num_tokens=len(your_x_vocab_list), output_mode='one_hot')
        
        # ... 其他原有代码

3. 在call函数中处理特征x并整合到特征集合

在模型的call方法里,读取输入的x特征,完成编码后加入到特征字典中:

def call(self, feat_inputs):
    features = feat_inputs
    anchor_embeddings = {}
    
    # 原有特征处理逻辑保留
    anchor_embeddings['anc_feat_vec'] = features['anc_feat_vec']
    anchor_embeddings['anc_ptcode_map_emb'] = self._embeddings['category'](features['anc_pcat_map'])
    anchor_embeddings['anc_pcat_map_emb'] = self._embeddings['ptype'](features['anc_ptcode_map'])
    
    # 新增特征x的处理
    x_raw = features['x']
    # 若用了Lookup层,先做索引映射:x_raw = self.x_lookup(x_raw)
    x_one_hot_emb = self.x_one_hot(x_raw)
    anchor_embeddings['x_one_hot'] = x_one_hot_emb
    
    # 后续可将所有embedding拼接(比如作为MMOE的输入):
    # concat_embeddings = tf.concat([v for v in anchor_embeddings.values()], axis=-1)
    # ... 其他原有逻辑

注意事项

  • 若不确定x的总类别数,可先遍历TFRecords统计所有取值,或用IntegerLookup.adapt()方法自动学习词汇表。
  • 若x是多维度int特征,需调整解析时的FixedLenFeatureshape,并根据需求处理编码逻辑(如先flatten)。
  • 确保输入pipeline中特征x的名称,与模型call函数中读取的名称完全一致。

内容的提问来源于stack exchange,提问作者Chris_007

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 06:56:03