展示多塔神经网络在银行推荐场景的端到端架构设计与PyTorch实现,适合ML工程师参考。
构建一个基于深度学习、具有可解释性的下一最佳产品推荐系统,有助于银行机构预测客户接下来需要哪种产品。银行拥有海量客户数据,包括交易历史、产品持有记录、人口统计特征和行为模式。然而,如何将这些数据转化为可执行的个性化产品推荐,仍然是一项重大挑战。传统的基于规则的系统和协同过滤方法,往往无法捕捉客户产品采用历程中复杂的时序模式。
本文介绍一个使用 Amazon SageMaker AI 和 PyTorch 构建的下一最佳产品(Next-Best-Product,NBP)推荐系统背后的架构与设计决策。我们将解释采用多塔神经网络架构的原因、学习型注意力机制如何提供客户级可解释性,以及 AWS 服务如何协同工作,将这一解决方案从研究阶段推进到生产环境。本文是一篇架构概览,而不是分步部署指南。无论你是在为金融服务还是其他拥有异构客户数据的领域构建推荐系统,本文介绍的架构模式都可以帮助你设计出更准确、更易于解释的模型。
要跟随本文中的架构模式和代码示例进行实践,你需要:
拥有一个 AWS 账户,并具备使用 SageMaker AI、Amazon Simple Storage Service(Amazon S3)、AWS Glue 和 Amazon CloudWatch 的权限
该解决方案需要一个 AWS Identity and Access Management(IAM)执行角色,并授予其访问以下 AWS 服务的权限。请按照最小权限原则创建策略,并将权限范围限定为该解决方案所需的资源。
SageMaker AI – 创建、描述、启动、停止和删除训练任务、处理任务、批量转换任务、模型、端点、端点配置、流水线、实验和监控计划。使用 InvokeEndpoint 进行实时推理。
Amazon S3 – 对数据存储桶具有读写权限。创建和删除存储桶。列出、上传、下载和删除对象。
AWS Glue – 创建、运行和删除 ETL 任务。创建和删除爬网程序。创建和删除 Data Catalog 数据库及表。
CloudWatch – 对日志组、日志流和指标具有读写权限。在清理过程中删除日志组。
IAM – 创建和删除角色。附加和分离策略。PassRole(仅限指定的执行角色 ARN,并将范围限定为 sagemaker.amazonaws.com 和 glue.amazonaws.com)。
有关为 SageMaker AI 编写最小权限 IAM 策略的指导,请参阅 SageMaker AI 的基于身份的策略示例。
熟悉 Python 3.11+ 和 PyTorch
创建该推荐系统所需的软件包:Python 3.11+、PyTorch 2.9+、Pandas 2.3+、NumPy 2.3+、scikit-learn 1.7+、Dask 2025.11+
我们建议使用虚拟环境,并在部署前使用 pip-audit 等工具扫描依赖项中的已知漏洞。
基本了解深度学习概念(嵌入、循环网络、注意力机制)
注意:部署该解决方案会创建需要付费的 AWS 资源,包括 SageMaker AI 训练任务(ml.g5.12xlarge GPU 实例)、SageMaker AI 端点、Amazon S3 存储和 AWS Glue 任务。请按照本文末尾的清理说明操作,以避免持续产生费用。
该解决方案采用多塔深度学习架构,其中包含四个专用神经网络塔,每个塔分别处理客户数据的一个不同方面。各个塔通过学习型注意力机制进行融合,从而同时实现较高的准确率和客户级可解释性。
下图展示了该解决方案的高层架构。
该架构旨在解决银行业的一项常见挑战:从多个产品类别(例如信用卡、存款、保险、贷款和抵押贷款)中预测客户接下来最有可能购买哪种产品,同时提供满足监管要求的可解释结果。
下表总结了该解决方案中的技术选型及其作用。
该解决方案使用 PyTorch,以支持动态计算图(处理可变长度序列时需要 pack_padded_sequence)、在多个架构阶段进行快速迭代,以及与 SageMaker AI 训练任务和推理容器进行原生集成。
为什么在 Amazon S3 上使用 Parquet?
该解决方案以 Snappy 压缩的 Parquet 格式将数据存储在 Amazon Simple Storage Service(Amazon S3)上。Parquet 的列式格式支持列裁剪(仅读取宽表文件中的部分列)、谓词下推(跳过无关的行组)、相比 CSV 实现 3–5 倍的压缩率,并且能够保留数据类型,无须在每次读取时重新解析。
为什么使用 AWS Glue 执行 ETL?
该项目使用运行在 PySpark 上的 AWS Glue 任务,实现无服务器、自动扩缩的数据处理。AWS Glue 提供原生 Spark 集成、用于灵活模式的 DynamicFrame API、Data Catalog 自动注册、用于增量处理的任务书签,以及按 DPU 计费带来的成本效益。
数据流水线架构
数据流水线包含两个阶段:首先使用 AWS Glue ETL 任务统一数据,然后使用 Amazon SageMaker Processing 任务执行机器学习专用的特征工程。
使用 AWS Glue 统一数据
银行数据通常来自多个源系统,并且模式并不一致。AWS Glue ETL 任务会规范化模式,将原始交易类型映射为统一的服务类别,把所有数据合并为每位客户的一条按时间排序的记录,并构建时序特征。处理后的输出以 Parquet 格式写入 Amazon S3,并注册到 AWS Glue Data Catalog。
使用 Amazon SageMaker Processing 执行机器学习专用的特征工程
AWS Glue 任务生成统一历史记录后,Amazon SageMaker Processing 任务会为每位客户创建产品采用序列,使用 Dask 进行并行处理,计算不同时间窗口(7 天、30 天、60 天、180 天和 365 天)的交易聚合,并将序列填充到固定长度,作为模型输入。
处理大规模数据
对于超出可用内存容量的大型数据集,该解决方案采用并行分块处理策略:使用 PyArrow 检查元数据,使用 ProcessPoolExecutor 并行处理数据块,在各批次之间显式执行垃圾回收,并采用增量合并来避免内存使用量骤增。
import gc
from concurrent.futures import ProcessPoolExecutor
chunksize = 5_000_000
n_workers = 4
for batch_start in range(0, total_chunks, n_workers):
with ProcessPoolExecutor(max_workers=n_workers) as executor:
futures = [
executor.submit(process_chunk_range, input_path, output_path, i, start_row, end_row)
for i in range(batch_start, min(batch_start + n_workers, total_chunks))
]
for future in futures:
future.result()
gc.collect() # Force garbage collection between batches
注意:处理真实客户数据时,请参阅“安全注意事项”部分,了解有关个人身份信息(PII)处理、法规合规和数据治理的指导。
该模型采用多塔方法,每个塔专门处理一种客户数据,随后通过基于注意力的融合机制将其整合。
为什么选择多塔架构而不是单一网络?
不同类型的客户数据具有完全不同的结构。序列是由离散 ID 组成的有序列表。交易数据是数值聚合结果。人口统计数据是分类特征与数值特征的混合。行为分群则是分类编码。
强制让这些数据经过相同的网络层会浪费模型容量。因此,该架构使用四个专用塔,每个塔都针对其数据类型进行了专门设计:
序列塔:捕捉时序模式
序列塔使用一个两层门控循环单元(Gated Recurrent Unit,GRU)处理客户的产品采用历史。它是架构的核心组件,因为它不仅能捕捉客户持有哪些产品,还能捕捉客户采用这些产品的先后顺序。
class SequenceTower(nn.Module):
def __init__(self, num_products, embedding_dim=32, hidden_dim=64, dropout=0.2):
super().__init__()
self.embedding = nn.Embedding(num_products + 1, embedding_dim, padding_idx=0)
self.gru = nn.GRU(
input_size=embedding_dim, hidden_size=hidden_dim,
num_layers=2, batch_first=True, dropout=dropout
)
self.active_count_layer = nn.Sequential(
nn.Linear(1, hidden_dim // 2), nn.ReLU(), nn.Dropout(dropout)
)
self.fusion = nn.Sequential(
nn.Linear(hidden_dim + hidden_dim // 2, hidden_dim),
nn.ReLU(), nn.Dropout(dropout)
)
def forward(self, sequence, seq_length, active_count): embedded = self.embedding(sequence) packed = nn.utils.rnn.pack_padded_sequence( embedded, seq_length.cpu().clamp(min=1), batch_first=True, enforce_sorted=False ) _, hidden = self.gru(packed) seq_features = hidden[-1] active_features = self.active_count_layer(active_count) return self.fusion(torch.cat([seq_features, active_features], dim=1))
GRU 有两个门(重置门、更新门),而 LSTM 有三个门(输入门、遗忘门、输出门),因此参数量大约减少了 33%。对于较短的序列(不超过 20 个元素),GRU 的性能与 LSTM 相当,但训练速度更快。更新门的插值机制还形成了一条类似残差连接的天然梯度路径。
为什么使用 pack_padded_sequence?
客户序列的长度各不相同。打包操作会让 GRU 忽略填充 token,防止模型从补零位置学习到噪声。
塔注意力机制:兼具可解释性的学习式融合
该架构并未采用简单拼接,而是使用学习式注意力机制来融合各塔的输出。正是这一设计提供了针对每位客户的可解释性,无须使用 SHAP 或 LIME 等事后解释方法。
```python
class TowerAttentionMechanism(nn.Module):
def __init__(self, hidden_dim=64, num_heads=4, dropout=0.1):
super().__init__()
self.tower_attention = nn.MultiheadAttention(
embed_dim=hidden_dim, num_heads=num_heads,
dropout=dropout, batch_first=True
)
self.context_weighting = nn.Sequential(
nn.Linear(hidden_dim * 4, 4), nn.Softmax(dim=1)
)
def forward(self, tower_outputs):
stacked = torch.stack(tower_outputs, dim=1) # [batch, 4, 64]
attended, _ = self.tower_attention(stacked, stacked, stacked)
stacked = stacked + attended # Residual connection
concat = torch.cat(tower_outputs, dim=1) # [batch, 256]
tower_weights = self.context_weighting(concat) # [batch, 4]
weighted_outputs = [
tower_outputs[i] * tower_weights[:, i:i+1]
for i in range(4)
]
return weighted_outputs, tower_weights
塔权重因客户而异。交易历史丰富的客户会获得较高的交易塔权重,而交易较少但人口统计特征明确的新客户则会获得较高的客户塔权重。这种自适应能力既提高了准确率,也为客户经理和监管机构提供了自然的可解释性。
上下文感知融合:使用残差块实现稳定训练
加权后的塔输出会经过一个带有残差连接的融合网络。残差连接有助于训练过程中的梯度流动,并允许网络在不需要额外深度时学习恒等映射。
class ContextAwareFusion(nn.Module):
def __init__(self, hidden_dim=64, dropout=0.2):
super().__init__()
self.initial_projection = nn.Linear(hidden_dim * 4, hidden_dim)
self.fusion1 = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim * 2), nn.LayerNorm(hidden_dim * 2),
nn.ReLU(), nn.Dropout(dropout), nn.Linear(hidden_dim * 2, hidden_dim)
)
self.layer_norm1 = nn.LayerNorm(hidden_dim)
self.fusion2 = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Dropout(dropout)
)
self.layer_norm2 = nn.LayerNorm(hidden_dim)
def forward(self, weighted_outputs):
concat = torch.cat(weighted_outputs, dim=1)
projected = self.initial_projection(concat)
out1 = self.layer_norm1(projected + self.fusion1(projected)) # Residual
out2 = self.layer_norm2(out1 + self.fusion2(out1)) # Residual
return out2
特征重要性模块:内置可解释性
银行业监管机构要求模型具有可解释性。该架构没有依赖事后解释方法,而是包含了一个特征重要性模块,在前向传播过程中生成针对每位客户的重要性分数,且这些分数之和为 1.0。
class FeatureImportanceModule(nn.Module):
def __init__(self, hidden_dim=64):
super().__init__()
self.feature_contribution = nn.Sequential(
nn.Linear(hidden_dim, 4), nn.Softmax(dim=1)
)
def forward(self, fused_features, tower_weights):
feature_importance = self.feature_contribution(fused_features)
return feature_importance * tower_weights
这会生成如下输出:“对于这位客户,40% 的推荐依据来自其产品序列,30% 来自交易模式,20% 来自人口统计特征,10% 来自行为细分。”客户经理可以利用这些信息,为与每位客户的沟通量身定制话术。
下表汇总了训练配置以及每项选择背后的理由。
所有随机种子(PyTorch、NumPy、CUDA)均已固定,以确保不同训练运行之间完全可复现。
训练过程使用 SageMaker AI 和 ml.g5.12xlarge 实例。SageMaker AI Python SDK 提供了一个 PyTorch Estimator,它会打包训练代码、预置 GPU 实例、执行训练,并自动将模型工件存储到 Amazon S3。
from sagemaker.pytorch import PyTorch
estimator = PyTorch(
entry_point='train.py',
source_dir='src/',
role=role,
instance_count=1,
instance_type='ml.g5.12xlarge',
framework_version='2.5.0',
py_version='py311',
hyperparameters={
'epochs': 50,
'batch_size': 32,
'learning_rate': 0.001,
},
)
模型使用与业务价值直接对应的指标进行评估:
生产模型在所有指标上都取得了出色表现,正确产品始终出现在排名前三的推荐结果中。特征重要性模块证实,序列塔(产品采用历史)贡献的信号最多,其次依次为交易模式、客户人口统计特征和行为细分。
推理与部署
推理流水线使用 SageMaker AI,同时支持批量评分和实时预测。
对于批量评分,SageMaker AI Batch Transform 每晚处理整个客户群,为每位客户生成带有可解释性分数的 top-k 推荐。结果以 JSON 格式存储在 Amazon S3 中,供 CRM 系统和客户经理仪表板使用。
对于实时预测,SageMaker AI 实时端点会在客户登录手机银行应用,或客户经理打开客户资料时,按需提供推荐。
每条推荐包括:
产品 ID 和概率分数。
特征重要性明细:每个塔的贡献百分比。
置信度指标:根据概率分布的熵计算。
def generate_batch_recommendations(model, dataloader, top_k=5):
with torch.no_grad():
for batch in dataloader:
outputs, feature_importance = model(
batch['sequence'], batch['seq_length'],
batch['active_count'], batch['transaction'],
batch['customer'], batch['behavioral']
)
probabilities = F.softmax(outputs, dim=1).cpu().numpy()
# Generate top-k recommendations with explainability scores
有关使用 IAM 身份验证、速率限制和虚拟私有云(VPC)隔离来保护端点的指导,请参阅“安全注意事项”部分。
注意:本文中的代码片段用于说明架构模式,并未达到生产就绪状态。在生产环境中,请添加输入验证(张量形状、NaN 检查、序列长度边界等)、错误处理和推理日志记录。请配置 Amazon SageMaker Model Monitor,以便在输入分布发生漂移时发出警报。
为什么选择多塔架构,而不是单一网络?单一网络需要同时学习如何处理序列、聚合交易、编码人口统计特征,以及解释行为细分。独立的塔可以各自专注于其对应的数据类型,然后由基于注意力的融合机制学习如何针对每位客户以最佳方式组合它们。
为什么在生产环境中选择 GRU,而不是 Transformer?Transformer 擅长处理长序列(100 个以上的元素)。对于不超过 20 个元素的序列,GRU 已经足够;与 Transformer 的注意力图相比,GRU 具有更清晰的可解释性,同时避免了注意力机制的二次方计算开销,并且生成的模型更小(约 5 MB,而 Transformer 约为 15 MB)。
为什么使用学习得到的塔权重,而不是拼接?采用拼接时,模型会对每位客户的所有塔一视同仁。采用学习得到的注意力权重后,模型可以针对每位客户进行自适应调整:交易历史丰富的客户会获得较高的交易塔权重,而新客户则会获得较高的人口统计特征塔权重。
为什么要使用时间窗口交易特征?不同的时间窗口捕捉不同的信号:7 天窗口捕捉即时意图,30 天窗口捕捉月度模式,180 天窗口捕捉季节性模式,365 天窗口捕捉年度模式。在过去 7 天交易频率突然增加的客户,其意图与 365 天内活动稳定的客户不同。
所有随机种子在 PyTorch、NumPy 和 CUDA 中固定,以确保训练的确定性。模型工件、超参数和数据版本通过 Amazon SageMaker Experiments 追踪。
Amazon SageMaker Model Monitor 检测数据漂移(输入特征分布的变化)、模型漂移(预测质量的降低)和潜在偏差(对人口统计特征的过度依赖)。
Amazon SageMaker Pipelines 工作流每月使用最新客户数据重新训练模型。管道自动化完整工作流:数据处理、训练、评估和条件部署(仅在指标优于当前生产模型时部署)。
使用真实银行数据部署此解决方案时,实施最小权限 IAM 角色,限制仅对所需的 SageMaker AI、Amazon S3、AWS Glue 和 CloudWatch 操作的访问。
使用 AWS Key Management Service(AWS KMS)客户托管密钥对 S3 存储桶和 SageMaker AI 训练卷进行静态数据加密,并通过 S3 存储桶策略对所有传输中数据强制执行 TLS。
在没有互联网网关的私有 VPC 子网中部署训练作业和端点,对 AWS 服务通信使用 VPC 端点(AWS PrivateLink),并在训练作业上设置 enable_network_isolation=True。
使用 AWS Signature Version 4 签名保护推理端点,考虑使用 Amazon Cognito 和 Amazon API Gateway 进行消费者身份验证和速率限制。
对于数据治理,评估监管义务(PCI-DSS、GDPR、CCPA),实施数据最小化,并定义保留策略。
启用 AWS CloudTrail 进行 API 审计日志记录,启用 S3 版本控制以确保模型工件完整性。
有关完整实施指导,请参阅 SageMaker AI 安全文档。
为避免持续的 AWS 费用,在完成此解决方案的评估后删除以下资源。
您可以通过 AWS 管理控制台或使用 AWS 命令行界面(AWS CLI)删除这些资源。
警告:删除这些资源是不可逆转的。在继续之前:
本文展示了如何使用 PyTorch 和 SageMaker AI 为银行构建下一个最佳产品推荐系统。具有学习型注意力融合的多塔架构在提供银行监管机构所需的可解释性的同时,实现了高预测准确度。
关键要点是:
您可以将此多塔架构适配到您自己的产品目录和客户数据。要开始,请探索 SageMaker AI 文档。有关 AWS 上 PyTorch 的更多示例,请参阅 PyTorch on AWS。如果您需要帮助为您的组织构建推荐系统,请联系 AWS 代表。