Split Learning
분할 학습이 디바이스 간에 신경망을 분할하여 협업형 AI를 지원하는 방식을 배우고, 개인정보 보호 위험, 학습 워크플로, 애플리케이션, 디자인 선택 사항을 살펴보세요.
스플릿 러닝은 신경망을 두 개 이상의 컴퓨팅 위치로 나누는 분산 머신 러닝 접근 방식입니다. 클라이언트는 모델의 초기 레이어를 통해 개인 입력 데이터를 처리하고, 중간 활성화 값만 서버로 전송하며, 로컬 레이어 학습을 계속하는 데 필요한 그라디언트를 수신합니다. 이를 통해 조직이나 장치는 원본 학습 데이터를 직접 전송하지 않고도 협업할 수 있습니다.
이 접근 방식은 머신 러닝이 프라이버시, 소유권, 대역폭 또는 하드웨어 경계를 넘어 작동해야 할 때 특히 유용합니다. 예를 들어, 병원은 의료 이미지를 로컬에 유지하고 더 강력한 서버는 모델의 연산 집약적인 부분을 실행할 수 있습니다. 그러나 원본 데이터를 로컬에 유지한다고 해서 중간 표현이 여전히 민감한 정보를 드러낼 수 있으므로 데이터 프라이버시가 자동으로 보장되는 것은 아닙니다.
스플릿 러닝의 작동 방식#
신경망은 선택한 컷 레이어에서 분할됩니다. 컷 레이어 이전의 레이어는 클라이언트에서 실행되고, 그 이후의 레이어는 서버에서 실행됩니다. 클라이언트 측 네트워크의 출력은 종종 활성화, 중간 표현 또는 스매시드 데이터라고 합니다.
학습 단계는 다음 순서로 진행됩니다.
- 클라이언트는 원본 입력에서 컷 레이어까지 순방향 패스를 실행합니다.
- 생성된 활성화 값을 서버로 전송합니다.
- 서버는 순방향 패스를 완료하고 손실을 계산합니다.
- 역전파 중에 서버는 컷 레이어 활성화 값에 대한 그라디언트를 계산하여 반환합니다.
- 클라이언트는 해당 그라디언트를 사용하여 로컬 레이어를 업데이트합니다.
이 프로세스는 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 AI 위험 관리 프레임워크는 이러한 위험을 평가하기 위한 더 넓은 프로세스를 제공합니다.
모든 학습 단계에 양방향 통신이 필요할 수 있으므로 대역폭과 지연 시간도 중요합니다. 느리거나 신뢰할 수 없는 클라이언트는 전체 시스템을 지연시킬 수 있으며, 일관되지 않은 데이터 분포는 수렴에 영향을 미칠 수 있습니다.
Ultralytics YOLO는 턴키 스플릿 러닝 오케스트레이션을 제공하지 않습니다. 이를 구현하려면 YOLO 아키텍처를 신중하게 분할하고, 원격 순방향 및 역방향 패스를 조율하며, 문서화된 커스텀 트레이너 워크플로를 확장해야 할 수 있습니다. 데이터가 소유한 하드웨어에만 남아 있으면 되는 프로젝트의 경우 Ultralytics Platform 모델 학습이 스트리밍된 메트릭과 함께 로컬 학습을 지원하지만, 모델 자체가 참가자 간에 분할되지 않으므로 로컬 학습은 스플릿 러닝이 아닙니다.






