1. 项目概述:基于LSTM的广告点击率预测系统
在互联网广告领域,点击率(CTR)预测一直是个核心问题。我最近完成了一个基于Flask框架和LSTM深度学习的广告点击率预测系统,这个毕设项目不仅实现了从数据收集到模型部署的全流程,还创新性地将用户行为序列建模应用于CTR预测场景。相比传统逻辑回归等静态模型,LSTM能够捕捉用户点击行为的时间依赖性,在实际测试中AUC指标提升了约15%。
这个系统的独特之处在于它采用了双通道输入架构:一方面处理用户静态特征(如 demographics),另一方面通过LSTM处理用户历史行为序列。在电商广告场景的测试中,这种架构对"看了又看"这类连续行为模式的预测准确率尤为突出。下面我将详细拆解这个项目的技术实现和踩坑经验。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 系统架构设计
2.1 整体技术栈选型
系统采用Python 3.8作为开发语言,主要基于以下技术栈:
- Web框架:Flask(轻量级,适合快速原型开发)
- 深度学习框架:TensorFlow 2.4 + Keras
- 前端:Bootstrap 5 + ECharts
- 数据库:MySQL 8.0(结构化数据)+ Redis(实时特征缓存)
- 部署:Docker + Nginx
选择Flask而非Django的考虑在于:
- 广告预测系统不需要Django自带的全功能Admin等组件
- Flask的轻量级特性更适合快速迭代模型API
- 可以更灵活地集成TensorFlow Serving
2.2 数据流设计
系统数据处理流程分为离线训练和在线预测两条管道:
离线训练流程:
python复制def offline_training_pipeline():
# 1. 数据采集
raw_data = collect_data_from_hive()
# 2. 特征工程
feature_engineer = FeatureEngineer()
processed_data = feature_engineer.transform(raw_data)
# 3. 模型训练
model = LSTMModel()
model.train(processed_data)
# 4. 模型导出
model.export_to_saved_model()
在线预测流程:
python复制@app.route('/predict', methods=['POST'])
def predict():
# 1. 获取实时请求数据
request_data = request.get_json()
# 2. 实时特征抽取
features = realtime_feature_extract(request_data)
# 3. 加载模型预测
predictor = load_model()
result = predictor.predict(features)
# 4. 返回结果
return jsonify({'ctr': result[0]})
2.3 核心模块划分
系统主要包含以下模块:
- 数据采集模块:负责从各业务系统收集用户行为日志
- 特征工程模块:处理特征编码、归一化、序列填充等
- 模型训练模块:LSTM网络构建与训练
- 预测服务模块:提供RESTful API接口
- 监控看板模块:展示模型性能指标
3. 关键算法实现
3.1 LSTM模型架构设计
模型采用Embedding+BiLSTM的混合结构,核心代码如下:
python复制def build_lstm_model():
# 输入层
user_input = Input(shape=(None,), name='user_seq')
# Embedding层
embedding = Embedding(input_dim=10000,
output_dim=64,
mask_zero=True)(user_input)
# BiLSTM层
lstm_out = Bidirectional(LSTM(units=128,
return_sequences=False))(embedding)
# 全连接层
dense = Dense(64, activation='relu')(lstm_out)
# 输出层
output = Dense(1, activation='sigmoid')(dense)
return Model(inputs=user_input, outputs=output)
模型训练的关键参数配置:
- 优化器:Adam(lr=0.001)
- 损失函数:BinaryCrossentropy
- 评估指标:AUC
- Batch size:256
- Epochs:20(早停策略)
3.2 特征工程实践
3.2.1 用户行为序列构建
用户行为序列的处理是项目的关键难点。我们采用滑动窗口方法构建序列:
python复制def build_behavior_sequence(raw_logs, window_size=10):
# 按时间排序
sorted_logs = sorted(raw_logs, key=lambda x: x['timestamp'])
sequences = []
for i in range(len(sorted_logs) - window_size):
# 提取窗口内行为
window = sorted_logs[i:i+window_size]
seq = [log['item_id'] for log in window]
label = sorted_logs[i+window_size]['click']
sequences.append((seq, label))
return sequences
3.2.2 特征编码策略
针对不同类型特征采用不同编码方式:
- 用户ID:Embedding
- 类别特征:One-Hot
- 数值特征:分桶+Embedding
- 时间特征:周期编码
注意:对于稀疏特征(如item_id),务必使用动态权重裁剪(AdaptiveEmbedding)来防止内存爆炸
3.3 模型优化技巧
3.3.1 样本权重调整
针对正负样本不均衡问题(CTR通常<5%),采用样本加权:
python复制def get_sample_weight(y_true):
pos_weight = len(y_true) / sum(y_true) - 1
return np.where(y_true == 1, pos_weight, 1)
3.3.2 自定义评估指标
实现AUC和RIG(Relative Information Gain)指标:
python复制class CustomMetrics(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
y_pred = self.model.predict(x_val)
auc = roc_auc_score(y_val, y_pred)
rig = 1 - log_loss(y_val, y_pred) / log_loss(y_val, [y_val.mean()]*len(y_val))
print(f"\nVal AUC: {auc:.4f}, RIG: {rig:.4f}")
4. 系统实现细节
4.1 Flask API设计
预测API的关键实现要点:
python复制@app.route('/api/v1/predict', methods=['POST'])
def predict():
try:
# 参数校验
data = request.get_json()
validate_request(data)
# 特征抽取
features = feature_pipeline.transform(data)
# 模型预测
with graph.as_default():
set_session(sess)
proba = model.predict([features])[0][0]
# 返回结果
return jsonify({
'status': 'success',
'ctr': float(proba),
'model_version': current_version
})
except Exception as e:
return jsonify({'status': 'error', 'message': str(e)}), 400
4.2 性能优化实践
4.2.1 预测加速技巧
- 模型预热:服务启动时预先加载模型
python复制# 服务启动时
global model, graph, sess
model = load_model()
graph = tf.get_default_graph()
sess = tf.keras.backend.get_session()
- 批量预测:支持批量请求处理
python复制def batch_predict(features_list):
# 将多个请求特征堆叠为batch
batch = np.stack(features_list)
with graph.as_default():
set_session(sess)
return model.predict(batch)
4.2.2 缓存策略
- 特征缓存:使用Redis缓存热门商品特征
python复制def get_item_features(item_id):
redis_key = f"item:{item_id}"
features = redis_client.get(redis_key)
if not features:
features = db.query_item_features(item_id)
redis_client.setex(redis_key, 3600, pickle.dumps(features))
return pickle.loads(features)
4.3 前端交互实现
使用ECharts实现点击率预测可视化:
javascript复制function renderCTRChart(predictions) {
const chart = echarts.init(document.getElementById('ctr-chart'));
const option = {
tooltip: {
trigger: 'axis',
formatter: params => {
return `广告ID: ${params[0].data[0]}<br/>
预测CTR: ${(params[0].data[1]*100).toFixed(2)}%`;
}
},
xAxis: {type: 'value'},
yAxis: {type: 'category', data: predictions.map(p => p.ad_id)},
series: [{
type: 'bar',
data: predictions.map(p => [p.ad_id, p.ctr])
}]
};
chart.setOption(option);
}
5. 部署与监控
5.1 Docker化部署
Dockerfile关键配置:
dockerfile复制FROM tensorflow/tensorflow:2.4.1-gpu
# 安装依赖
RUN pip install flask gunicorn redis mysqlclient
# 复制代码
COPY . /app
WORKDIR /app
# 暴露端口
EXPOSE 8000
# 启动命令
CMD ["gunicorn", "-b :8000", "--workers=4", "--threads=2", "app:app"]
使用docker-compose编排服务:
yaml复制version: '3'
services:
web:
build: .
ports:
- "8000:8000"
depends_on:
- redis
- mysql
redis:
image: redis:6
ports:
- "6379:6379"
mysql:
image: mysql:8.0
environment:
MYSQL_ROOT_PASSWORD: ${DB_PASSWORD}
ports:
- "3306:3306"
5.2 监控系统搭建
使用Prometheus+Grafana监控:
- 添加Flask监控中间件
python复制from prometheus_flask_exporter import PrometheusMetrics
metrics = PrometheusMetrics(app)
metrics.info('app_info', 'CTR Prediction Service', version='1.0')
- 配置Grafana看板监控:
- QPS
- 预测延迟
- 模型AUC变化
- 特征缓存命中率
6. 常见问题与解决方案
6.1 冷启动问题
问题表现:新用户/新广告缺乏历史数据,预测不准
解决方案:
- 基于内容相似度的兜底策略
python复制def cold_start_predict(item):
# 获取内容特征
content_feat = get_content_features(item)
# 找到最相似的热门商品
similar_items = find_similar_items(content_feat)
# 返回相似商品的平均CTR
return np.mean([i.avg_ctr for i in similar_items])
- 渐进式特征填充机制
6.2 数据漂移问题
问题表现:用户行为模式变化导致模型效果下降
解决方案:
- 建立模型性能监控报警
- 实现自动化retraining pipeline
python复制def retrain_pipeline():
while True:
# 每天检查模型性能
time.sleep(86400)
# 如果AUC下降超过阈值
if check_performance_drop():
# 触发重新训练
train_new_model()
# AB测试新模型
if validate_new_model():
# 上线新模型
deploy_new_model()
6.3 线上服务稳定性
问题表现:预测延迟波动,服务超时
优化措施:
- 实现请求队列和限流
python复制from flask_limiter import Limiter
limiter = Limiter(app, key_func=get_remote_address)
@app.route('/predict')
@limiter.limit("100/minute")
def predict():
...
- GPU资源动态分配策略
7. 项目总结与改进方向
这个基于LSTM的广告点击率预测系统在测试环境中取得了不错的效果,相比基线逻辑回归模型,AUC提升了0.15,线上AB测试显示点击率提升了8%。但在实际落地过程中也暴露出一些问题:
- 实时性不足:当前系统对用户最新行为的响应有约5分钟延迟
- 多模态特征利用不足:未充分挖掘广告图片/文本内容信息
- 可解释性差:业务方难以理解LSTM的黑盒决策
后续改进方向:
- 引入流式计算框架(如Flink)提升实时性
- 增加视觉/文本特征提取模块
- 开发模型解释工具(如SHAP值分析)
这个项目让我深刻体会到,工业级的推荐系统不仅需要好的算法,更需要考虑工程实现、系统稳定性和业务需求的全方位平衡。特别是在处理用户行为序列时,如何平衡序列长度和计算效率是个需要持续优化的课题。
