도메인 적응에서 Adversarial Adaptation과 Self-Training 활용법
라벨이 부족한 타깃 도메인에 모델을 전이하기 위한 Adversarial Domain Adaptation과 Self-Training의 구조, 운영 제어, PyTorch 구현을 정리한다.
2026-08-14 · 최초 발행 2024-04-29
타깃 라벨 없이 전이 성능을 끌어올리는 방법
트랜스퍼 러닝은 소스 도메인에서 습득한 지식을 타깃 도메인으로 옮겨 표본 수와 분포 차이의 부담을 줄이는 접근이다. 여기서는 라벨이 부족한 타깃 도메인으로 일반화 성능을 높이고, 데이터를 비용 효율적으로 활용하면서 모델 재사용을 확대하는 방법을 다룬다.
문제 설정에는 UDA(비지도 도메인 적응), SSDA(준지도 도메인 적응), Closed-set·Partial·Open-set 같은 시나리오가 있다. 공변량 변화(Covariate shift), 라벨 변화(Label shift), 클래스 조건부 변화(Class-conditional shift), 부정 전이(Negative transfer)를 함께 고려해야 한다.
Adversarial Domain Adaptation(ADA)은 특징 추출기가 도메인 불변(feature-invariant) 표현을 학습하도록 만들어 소스와 타깃 분포의 정합(domain alignment)을 유도한다. 특징 추출기 F, 분류기 C, 도메인 판별기 D, GRL(Gradient Reversal Layer)을 사용하며, F는 C의 분류 손실을 낮추는 동시에 D를 속이는 방향으로 갱신된다.
Self-Training(ST)은 타깃의 비라벨 데이터에서 신뢰도 기반 의사 라벨(pseudo label)을 만들고 이를 다시 학습에 쓰는 방식이다. 데이터 자체를 증강하는 효과를 노린다. 교사-학생(Teacher-Student) 구조나 단일 모델의 반복 재라벨링으로 구성할 수 있으며, EMA(Exponential Moving Average) 교사, 온도 스케일링, 임계값 스케줄링을 적용한다.
공통 파이프라인과 안정화 지점
입력은 소스 라벨 데이터 Ds=(xs, ys)와 타깃 비라벨 데이터 Dt=(xt)다. 특징 추출기 F와 분류기 C를 중심으로 학습해 타깃 데이터에 더 잘 맞는 분류기 C를 얻는다. 운영 단계에서는 소스·타깃 배치 비율을 조절하고, 타깃 검증을 대신할 지표가 필요하다.
ADA에서는 F가 백본, C가 헤드, D가 도메인 판별 역할을 맡고 GRL이 연결된다. 손실은 L = Lcls(Ds) + λ·Ladv(Ds, Dt)로 둘 수 있다. 기본 정합은 전역(feature-level) 단위이며, 필요하면 클래스 조건부 정합(conditional alignment)을 더한다.
ST에서는 p(y|x)의 최고 확률이 τ 이상인 샘플만 의사 라벨로 채택한다. 클래스별 임계값과 샘플 수 균형을 조절해야 하며, Teacher EMA, Consistency Regularization, Soft/Hard pseudo label의 혼용이 안정화에 쓰인다.
엔트로피 최소화, 샤프니스 인식 최적화(SAM), 강·약 증강 일관성(FixMatch류), BatchNorm 적응(Tent류), 모멘트 매칭, 비균형 재가중치(CB loss)도 보조 수단이 된다. Confirmation bias와 라벨 쉬프트에 따른 오정합은 역검증(reverse validation), 타깃 엔트로피·불확실성 모니터링으로 다룬다. λ 스케줄링, 도메인 혼합(MixUp), 클래스 조건부 정합, 개방집합 탐지는 부정 전이를 줄이는 데 사용한다.
특징 정합과 의사 라벨이 만나는 흐름
D 손실이 급락하면 D 용량을 줄이거나 정규화를 강화한다. τ가 낮아 타깃 커버리지가 부족할 때는 스케줄링을 상향하고, 불확실성이 상승하면 학습을 중단한 뒤 체크포인트를 롤백한다.
ADA는 판별기를 혼동시키며 특징을 맞춘다
ADA의 입력은 Ds=(xs, ys), Dt=(xt), 초기화한 F·C·D, 그리고 하이퍼파라미터 λ다.
- 소스 배치로 Lcls를 계산하고 F와 C를 갱신한다.
- 소스·타깃 혼합 배치로 D를 학습한다. GRL을 거친 F는 D가 도메인을 구분하기 어렵게 만드는 방향으로 갱신된다.
- λ를 워밍업(예: 0→1)하고, Backbone은 낮게 D는 높게 두 축의 학습률을 관리한다.
이 과정을 거쳐 F는 도메인 불변 표현을, C는 향상된 타깃 일반화 성능을 목표로 한다. 스펙트럴 노름, Gradient penalty, 입력에 y를 포함하는 클래스 조건부 D, MMD·BNM 보조 손실을 함께 적용할 수 있다.
ST는 신뢰도 필터와 교사 갱신으로 적응한다
ST는 사전학습된 F·C, Dt=(xt), 임계값 τ, 교사-학생 업데이트 계수 α를 입력으로 사용한다.
- 교사 T가 xt를 예측하고 신뢰도가 τ 이상인 샘플에서 의사 라벨을 만든다.
- 약·강 증강 간 일관성 손실로 학생 S를 업데이트하며 클래스 균형을 유지한다.
- EMA로
T ← αT + (1−α)S를 갱신하고, τ와 샘플 수 스케줄을 조절한다.
산출물은 타깃 도메인에 특화해 재학습한 모델이다. 온도 스케일링, 불확실성 추정(MC Dropout/DE), 역확률 가중치는 각각 의사 라벨의 안정화와 라벨 쉬프트 보정에 사용된다.
선택 기준은 분포 차이와 의사 라벨 품질이다
| 지표 | Adversarial Domain Adaptation | Self-Training |
|---|---|---|
| 성능 | 분포 차가 크고 라벨 쉬프트가 작을 때 강하다. 클래스 조건부 정합 시 추가 향상을 기대할 수 있다. | 소프트·하드 라벨 품질이 높을 때 강하다. 높은 τ에서는 고정밀·저재현 트레이드오프가 있다. |
| 확장성 | 추가 모듈(D, GRL)로 복잡도가 늘지만 분산 학습 호환은 양호하다. | 라벨링 파이프라인과 필터링이 필요하며, 대규모 타깃 데이터에 선형 확장한다. |
| 일관성 | 전역 정합만 쓰면 클래스 혼선 위험이 있고 조건부 정합으로 개선한다. | 증강 일관성으로 결정경계를 정제하지만 클래스 불균형에서는 편향 위험이 있다. |
| 안정성 | D-F 간 미니맥스가 불안정할 수 있어 λ와 학습률 스케줄이 필요하다. | Confirmation bias가 가능하며 τ·EMA·불확실성 제어로 완화한다. |
| 운영 편의 | 단일 학습 파이프라인에 통합하기 쉽다. | 주기적 재라벨링과 선별 로직을 운영해야 하지만 온라인 적응(TTA)에 용이하다. |
도메인 차이가 드러나는 적용 장면
제조 비전에서는 합성 이미지에서 실사 결함 검출로 옮길 때 ADA로 텍스처와 조명에 불변인 특징을 학습하고, ST로 현장 데이터에 계속 적응할 수 있다.
리테일 이미지 분류에서는 스튜디오 이미지와 유저 업로드 이미지 사이의 차이를 초기 ADA 정합으로 다룬 뒤, 운영 중 ST로 계절과 트렌드 변동을 흡수한다. OCR·문서 분석에서는 도메인별 폰트와 스캔 품질 차이에 BN 적응과 ST를 적용해 저품질 스캔의 커버리지를 넓힌다.
음성 인식에서는 마이크와 환경 변화에 대해 ADA로 잡음 도메인을 정합하고, ST로 환경별 커스텀 파인튜닝을 수행한다. 의료 영상은 기기·기관 간 편차를 대상으로 클래스 조건부 ADA와 불확실성 기반 ST를 조합하며, 규제 환경에서는 로그와 추적성을 확보한다.
배치와 신뢰도 제어를 운영에 포함한다
소스:타깃 배치 비율은 1:1~1:3 범위에서 탐색하고 클래스 균형 리샘플링을 적용한다. 강·약 증강을 섞되 RandAug·CTAugment를 활용하고 Color jitter는 제한한다.
ADA에서는 λ 커브(예: 0→1 Sigmoid), D capacity, 배치 간 고정된 GRL을 관리한다. ST에서는 τ를 0.95→0.8로 스케줄링하고 클래스별 τ를 둘 수 있으며, EMA α는 0.99~0.999 범위에서 사용한다.
타깃이 무라벨인 환경에서는 역검증(reverse validation)과 타깃 엔트로피↓, A-distance↓, BNM↑를 모니터링한다. ECE·ACE를 측정하고 온도 스케일링으로 캘리브레이션을 조정한다. Open-set 탐지(energy score, MSP), 거부 옵션, EM·BBSE 기반 라벨 쉬프트 추정 및 재가중치, 체크포인트 앙상블·EMA·조기 종료도 리스크 완화 수단이다.
PyTorch로 구현하는 ADA와 Self-Training
전제조건: Python 3.10+, PyTorch 2.2+, CUDA 선택. 분류기 예시이며 데이터로더와 백본은 대체 가능하다.
# pip install torch torchvision
import torch
import torch.nn as nn
import torch.nn.functional as F
class GRL(torch.autograd.Function):
@staticmethod
def forward(ctx, x, lambda_):
ctx.lambda_ = lambda_
return x.view_as(x)
@staticmethod
def backward(ctx, grad_output):
return -ctx.lambda_ * grad_output, None
class Feature(nn.Module):
def __init__(self, dim=128):
super().__init__()
self.net = nn.Sequential(nn.Flatten(), nn.Linear(784, dim), nn.ReLU(), nn.Linear(dim, dim))
def forward(self, x): return self.net(x)
class Classifier(nn.Module):
def __init__(self, dim=128, num_classes=10):
super().__init__()
self.fc = nn.Linear(dim, num_classes)
def forward(self, f): return self.fc(f)
class Discriminator(nn.Module):
def __init__(self, dim=128):
super().__init__()
self.net = nn.Sequential(nn.Linear(dim, 128), nn.ReLU(), nn.Linear(128, 1))
def forward(self, f): return self.net(f)
def train_ada_step(x_s, y_s, x_t, Ftr, Cls, Dis, opt_FC, opt_D, lambda_):
Ftr.train(); Cls.train(); Dis.train()
# 1) Supervised on source
f_s = Ftr(x_s); y_hat = Cls(f_s)
loss_cls = F.cross_entropy(y_hat, y_s)
# 2) Adversarial domain loss
f_t = Ftr(x_t).detach() # detach for D update
d_s = Dis(f_s.detach())
d_t = Dis(f_t)
y_dom_s = torch.ones_like(d_s)
y_dom_t = torch.zeros_like(d_t)
loss_d = F.binary_cross_entropy_with_logits(d_s, y_dom_s) + \
F.binary_cross_entropy_with_logits(d_t, y_dom_t)
opt_D.zero_grad(); loss_d.backward(); opt_D.step()
# 3) GRL update for F
f_s = Ftr(x_s); f_t = Ftr(x_t)
f_stacked = torch.cat([f_s, f_t], 0)
logits_dom = Dis(GRL.apply(f_stacked, lambda_))
y_dom = torch.cat([torch.ones_like(logits_dom[:len(f_s)]),
torch.zeros_like(logits_dom[len(f_s):])], 0)
loss_adv = F.binary_cross_entropy_with_logits(logits_dom, y_dom)
loss = loss_cls + lambda_ * (-loss_adv) # maximize adv by minimizing negative
opt_FC.zero_grad(); loss.backward(); opt_FC.step()
return loss_cls.item(), loss_d.item(), loss_adv.item()
@torch.no_grad()
def pseudo_label(x_t, model, tau=0.9, T=1.0):
model.eval()
logits = model['C'](model['F'](x_t))/T
probs = logits.softmax(-1)
conf, y_hat = probs.max(-1)
mask = conf >= tau
return y_hat[mask], mask
def ema_update(teacher, student, alpha=0.99):
for tp, sp in zip(teacher.parameters(), student.parameters()):
tp.data.mul_(alpha).add_(sp.data * (1 - alpha))
# Self-Training step (student S, teacher T)
def train_st_step(x_t_weak, x_t_strong, Tm, Sm, opt_S, tau=0.9):
with torch.no_grad():
logits_t = Tm['C'](Tm['F'](x_t_weak))
probs = logits_t.softmax(-1)
conf, y = probs.max(-1)
mask = conf >= tau
logits_s = Sm['C'](Sm['F'](x_t_strong[mask]))
loss = F.cross_entropy(logits_s, y[mask]) if mask.any() else torch.tensor(0., device=logits_s.device)
opt_S.zero_grad(); loss.backward(); opt_S.step()
return loss.item(), mask.float().mean().item()
적용할 때 λ는 0→1로 워밍업하고 Dis LR는 F/C 대비 25배로 둔다. ST의 τ는 0.95에서 시작해 데이터 품질에 따라 0.850.9로 조절하며, EMA α는 0.99~0.999를 사용한다.
기대할 수 있는 변화와 안전장치
타깃 정확도는 벤치마크 기준으로 데이터 난이도에 따라 520%p 개선될 수 있으며, 캘리브레이션 지표(ECE)는 최대 30% 개선된다. 타깃 무라벨 활용 시나리오에서는 라벨링 비용을 5080% 절감할 수 있다.
신규 도메인 온보딩의 TTM을 줄이고 현장 변화에 대한 회복력을 높이는 효과도 있다. 무라벨·온디바이스 적응은 데이터 거버넌스와 프라이버시 요구를 충족하는 데 활용된다.
초기 이행에서는 ADA로 베이스라인을 안정화하고, 운영 단계에서는 ST로 지속 적응하는 구성이 상호 보완적이다. 이때 불확실성·라벨 쉬프트 모니터링과 안전장치를 함께 둬야 한다.