Split Learning
了解拆分学习如何将神经网络分布在多个设备上以支持协同人工智能,同时探讨隐私风险、训练工作流、应用和设计选择。
拆分学习是一种分布式机器学习方法,将神经网络分割到两个或多个计算位置之间。客户端通过模型的早期层处理私有输入数据,仅将中间激活发送到服务器,并接收继续训练其本地层所需的梯度。这允许组织或设备在不直接传输原始训练数据的情况下进行协作。
当机器学习必须跨隐私、所有权、带宽或硬件边界运行时,这种方法尤其相关。例如,医院可以在本地保留医学影像,而更强大的服务器则运行模型的计算密集型部分。然而,将原始数据保留在本地并不能自动保证数据隐私,因为中间表征可能仍然会泄露敏感信息。
拆分学习的工作原理#
神经网络在选定的切割层处被分割。切割层之前的层在客户端上执行,而切割层之后的层在服务器上执行。客户端网络输出通常被称为激活、中间表征或粉碎数据。
训练步骤遵循以下顺序:
- 客户端运行从原始输入到切割层的前向传播。
- 客户端将生成的激活发送到服务器。
- 服务器完成前向传播并计算损失。
- 在反向传播期间,服务器计算切割层激活的梯度并将其返回。
- 客户端使用该梯度更新其本地层。
此过程依赖于诸如PyTorch自动微分等系统实现的相同链式法则。不同之处在于,激活和梯度在训练期间会跨越网络边界。
以下单进程示例模拟了该边界:
import torch
from torch import nn
client_model = nn.Sequential(nn.Linear(8, 16), nn.ReLU())
server_model = nn.Sequential(nn.Linear(16, 2))
client_optimizer = torch.optim.SGD(client_model.parameters(), lr=0.01)
server_optimizer = torch.optim.SGD(server_model.parameters(), lr=0.01)
inputs = torch.randn(4, 8)
targets = torch.tensor([0, 1, 0, 1])
client_optimizer.zero_grad()
server_optimizer.zero_grad()
client_activations = client_model(inputs)
sent_activations = client_activations.detach().requires_grad_()
predictions = server_model(sent_activations)
loss = nn.CrossEntropyLoss()(predictions, targets)
loss.backward()
client_activations.backward(sent_activations.grad)
server_optimizer.step()
client_optimizer.step()
print(loss.item())分离 client_activations 表示将它们发送到另一个系统。返回的激活梯度重新连接两个半部分以进行优化。生产实现必须增加网络连接、身份验证、加密、故障处理和隐私控制。
拆分学习与相关方法#
拆分学习属于更广泛的分布式训练领域,但它以不同的方式划分计算。
- 联邦学习: 每个参与者通常训练一个完整的本地模型并发送模型更新以进行聚合。拆分学习仅赋予每个参与者模型的一部分,并交换中间激活和梯度。
- 流水线并行: 两种方法都在不同设备上放置不同的层。流水线并行主要在受信任的环境中提高规模或硬件利用率,而拆分学习通常将数据所有者与计算提供者分离。
- 数据并行训练: 诸如PyTorch DistributedDataParallel和TensorFlow分布式训练之类的框架会复制模型并同步更新。它们通常不会将早期层排他性地保留在原始数据旁边。
拆分学习还可以支持垂直分区的数据,其中组织为匹配的记录保留不同的特征。SecretFlow拆分学习工作流说明了这种安排。
实际应用#
-
协同医学影像: 医院可以训练共享的计算机视觉系统,同时将X射线或扫描件保留在其自己的基础架构内。每个医院在本地运行第一层,中央服务器从中间特征完成训练。MIT拆分学习概述使用放射科中心来解释这种架构。
-
资源受限的工业摄像机: 工厂摄像机或网关可以在本地运行紧凑的特征提取器,而服务器则为目标检测训练其余层。这可以减少原始视频传输和客户端计算,使该方法适用于跨多个设施运行的边缘AI系统。
益处、风险和设计选择#
切割层决定了客户端工作负载、服务器工作负载、通信量和信息暴露之间的平衡。较早的切割减少了客户端计算,但可能会产生较大的、类似输入的激活。较晚的切割可以创建更抽象的特征,但需要更强的客户端硬件。
中间激活和梯度可能仍然容易受到重建、推断或操纵的影响。因此,团队应评估访问控制、加密传输、激活保护、审计日志记录和参与者信任,而不是将拆分学习视为完整的隐私解决方案。NIST隐私框架和NIST人工智能风险管理框架为评估这些风险提供了更广泛的流程。
带宽和延迟也很重要,因为每个训练步骤都可能需要双向通信。缓慢或不可靠的客户端可能会延迟整个系统,而不一致的数据分布可能会影响收敛。
Ultralytics YOLO不提供现成的拆分学习编排。实现它需要仔细划分YOLO架构、协调远程前向和反向传播,并可能扩展已记录的自定义训练器工作流。对于仅需要数据保留在自有硬件上的项目,Ultralytics平台模型训练支持带有流式指标的本地训练,但本地训练并非拆分学习,因为模型本身并未在参与者之间进行分割。






