TensorFlow Recommenders模型添加经纬度特征的实现方法
TFRS架构下接入经纬度特征的通用实现思路
- 第一步:经纬度特征预处理
原始经纬度是范围固定的连续值(经度范围[-180, 180]、纬度范围[-90,90]),直接输入会和文本分支输出的embedding数值尺度差异过大,导致训练不稳定,需要先做特征变换:- 基础方案:用
tf.keras.layers.Normalization对经纬度做标准化,提前用训练集的经纬度数据调用adapt()拟合均值方差即可,用法和当前代码中使用的TextVectorization层完全一致。 - 优化方案:做周期性编码,将经纬度转换为
sin(经度*π/180)、cos(经度*π/180)、sin(纬度*π/180)、cos(纬度*π/180)四个特征,解决经纬度的环形边界问题(比如经度179°和-179°实际地理距离极近,但原始数值差极大)。
- 基础方案:用
- 第二步:扩展QueryModel结构,新增坐标特征编码分支
在现有文本特征处理分支的基础上,新增独立的坐标特征处理分支:将预处理后的坐标特征输入2-3层全连接层做非线性投影,输出维度可以自行调整(通常设为16/32维,和文本embedding维度适配即可)。最后在call()方法中,将文本embedding和坐标特征的投影向量拼接,作为Query模型的最终输出。 - 第三步:修改上层RatingsModel的特征传参
现有RatingsModel调用Query模型时仅传入了query_features字段,需要额外把输入特征字典里的经度、纬度字段一并传入Query模型,保证坐标分支能拿到对应数据。
修改后的参考代码
修改后的QueryModel:
class QueryModel(tf.keras.Model): def __init__(self, train_coords): super().__init__() max_tokens = 10_000 # 原有文本特征处理分支 self.query_features_vectorizer = tf.keras.layers.TextVectorization( max_tokens=max_tokens) self.query_features_embedding = tf.keras.Sequential([ self.query_features_vectorizer, tf.keras.layers.Embedding(max_tokens, 64, mask_zero=True), tf.keras.layers.GlobalAveragePooling1D(), ]) self.query_features_vectorizer.adapt(query_features) # 新增经纬度处理分支 self.coord_norm = tf.keras.layers.Normalization(axis=-1) # 用训练集经纬度拟合归一化参数,train_coords形状为(num_samples, 2) self.coord_norm.adapt(train_coords) # 坐标特征投影层,输出维度可自行调整 self.coord_encoder = tf.keras.Sequential([ tf.keras.layers.Dense(32, activation="relu"), tf.keras.layers.Dense(16) ]) def call(self, inputs): # 拼接文本embedding和坐标特征embedding text_emb = self.query_features_embedding(inputs["query_features"]) # 从输入中取出经纬度,拼接为[batch_size, 2]的张量 raw_coords = tf.stack([inputs["lon"], inputs["lat"]], axis=1) coord_emb = self.coord_encoder(self.coord_norm(raw_coords)) return tf.concat([text_emb, coord_emb], axis=1)
RatingsModel中需要修改的传参部分:
def call(self, features): # 新增传入经纬度字段到Query模型 query_embeddings = self.query_model({ "query_features": features["query_features"], "lon": features["lon"], "lat": features["lat"], }) warehouse_embeddings = self.candidate_model({ "warehouse_id": features["warehouse_id"], }) return ( self.rating_model( tf.concat([query_embeddings, warehouse_embeddings], axis=1) ), )
可选优化点
- 如果候选仓库本身也维护了经纬度属性,除了在Query侧加用户选择的坐标特征外,还可以在候选模型侧加入仓库坐标的编码分支,或者在拼接Query、候选embedding输入评分模型前,额外拼接查询点和仓库点的球面距离(Haversine距离)作为交叉特征,能大幅提升地理位置相关的推荐效果。
- 坐标投影层的维度、激活函数可以根据数据集规模灵活调整,不需要固定参数。
- 训练和推理阶段要保证经纬度字段是数值类型,不要传入字符串格式的坐标值。
内容的提问来源于stack exchange,提问作者Даниэль Сеидов
相关产品推荐
相关产品推荐

