如何向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
相关产品推荐
相关产品推荐

