Skip to content

联邦推荐系统

推荐系统是联邦学习最自然的应用场景之一。用户行为数据天然分布在各设备上,且涉及隐私,联邦推荐无需收集用户原始行为即可训练推荐模型。

联邦推荐系统

传统推荐 vs 联邦推荐

维度传统推荐联邦推荐
数据位置中心服务器用户设备
隐私风险高(数据集中)低(数据不出域)
冷启动依赖全局数据更难
实时性批量更新增量更新
合规性难(GDPR)

联邦矩阵分解

python
class FederatedMatrixFactorization:
    """联邦矩阵分解推荐"""

    def __init__(self, num_items, embedding_dim=64):
        self.num_items = num_items
        self.embedding_dim = embedding_dim
        # 全局物品嵌入
        self.item_embeddings = nn.Embedding(num_items, embedding_dim)

    def client_train(self, user_interactions, local_user_embed):
        """客户端本地训练"""
        optimizer = torch.optim.SGD(
            [local_user_embed], lr=0.01
        )

        for item_id, rating in user_interactions:
            item_embed = self.item_embeddings(torch.tensor(item_id))
            pred = torch.dot(local_user_embed, item_embed)
            loss = (pred - rating) ** 2
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()

        return local_user_embed

    def server_update(self, item_grad_updates):
        """服务端更新物品嵌入"""
        for item_id, grad in item_grad_updates.items():
            self.item_embeddings.weight.data[item_id] -= 0.01 * grad

联邦神经网络推荐

python
class FederatedNCF:
    """联邦 Neural Collaborative Filtering"""

    def __init__(self, num_users, num_items, embed_dim=64):
        self.user_embed = nn.Embedding(num_users, embed_dim)
        self.item_embed = nn.Embedding(num_items, embed_dim)
        self.mlp = nn.Sequential(
            nn.Linear(embed_dim * 2, 128),
            nn.ReLU(),
            nn.Linear(128, 64),
            nn.ReLU(),
            nn.Linear(64, 1),
            nn.Sigmoid()
        )

    def local_train(self, user_id, interactions):
        """本地训练用户个性化部分"""
        u_embed = self.user_embed(torch.tensor(user_id))
        total_loss = 0

        for item_id, label in interactions:
            i_embed = self.item_embed(torch.tensor(item_id))
            concat = torch.cat([u_embed, i_embed])
            pred = self.mlp(concat)
            total_loss += F.binary_cross_entropy(pred, torch.tensor([label]))

        return total_loss

通信优化

推荐系统用户量大,通信优化至关重要:

方法压缩比精度损失适用
Top-K10-100x梯度稀疏
量化 (INT8)4x通用
哈希压缩100x+超大规模
蒸馏10x模型压缩

联邦推荐实践

  1. 物品嵌入全局共享,用户嵌入本地保留
  2. 使用负采样减少计算量
  3. 新物品通过内容特征冷启动
  4. DP 保护用户偏好隐私

隐私攻击

研究表明,即使只上传梯度,攻击者仍可通过模型反转推断用户交互过的物品。必须结合 DP 或安全聚合。

相关资源

最近更新