TensorFlow 2.4中如何构建<feature_name:feature_weight>映射字典
TensorFlow feature_column 特征名与模型权重映射方案
前置修正
你示例中获取权重的代码存在错误:DenseFeatures层仅负责特征预处理转换,没有可训练参数,LR的权重实际存储在第二层的Dense层中,正确获取权重的方式如下:
# 从Dense层获取权重,shape为(总特征维度, 1),拉平为1维数组 weights = my_model.layers[1].get_weights()[0].flatten()
DenseFeatures层的输出特征顺序和你传入的my_features列表顺序严格一致,只需要按顺序遍历每个特征列生成对应维度的特征名,再和权重数组按位置匹配即可实现映射。
完整实现代码
import tensorflow as tf from itertools import product def build_feature_weight_map(feature_columns, weights): feature_names = [] for col in feature_columns: # 处理被IndicatorColumn包装的词汇表分类列 if isinstance(col, tf.feature_column.IndicatorColumn): cat_col = col.categorical_column if isinstance(cat_col, tf.feature_column.VocabularyListCategoricalColumn): col_key = cat_col.key vocab_size = len(cat_col.vocabulary_list) # 按词汇表顺序生成特征名 for i in range(vocab_size): feature_names.append(f"{col_key}_{i}") # 处理被IndicatorColumn包装的交叉特征列 elif isinstance(cat_col, tf.feature_column.CrossedColumn): cross_keys = "_X_".join([k.key for k in cat_col.keys]) # 你的场景中交叉组合共4*4=16种,小于hash_bucket_size=50无哈希冲突 weight_vocab = [2,3,0,1] vol_vocab = [3,4,1,2] all_pairs = list(product(weight_vocab, vol_vocab)) # 按交叉列的哈希规则对组合排序,匹配DenseFeatures输出顺序 def get_pair_hash(pair): return cat_col._hash_bucket_function(str(pair).encode(), cat_col.hash_bucket_size) sorted_pairs = sorted(all_pairs, key=get_pair_hash) for w_val, v_val in sorted_pairs: feature_names.append(f"{cross_keys}_{w_val}_{v_val}") # 生成特征名-权重映射字典 return dict(zip(feature_names, weights)) # 调用方法获取映射字典 feature_weight_map = build_feature_weight_map(my_features, weights)
注意事项
- 如果交叉列的
hash_bucket_size小于实际特征组合数,存在哈希冲突时,可直接按桶编号命名交叉特征为{交叉列名}_{桶id}即可 - 新增其他类型特征列(如数值列、分桶列)时,只需在函数中新增对应类型的判断逻辑,按列的维度生成对应特征名即可
内容的提问来源于stack exchange,提问作者Bread
相关产品推荐
相关产品推荐

