Tensor Fusion과 Cross-Modal Retrieval로 설계하는 멀티모달 검색
Tensor Fusion과 공유 임베딩 기반 Cross-Modal Retrieval의 학습, 검색 인덱스, 결측·서빙 운영 전략을 정리한다.
2026-08-14 · 최초 발행 2024-04-29
융합 모델과 교차모달 검색이 다루는 문제
텍스트, 이미지, 오디오처럼 표현 방식이 다른 데이터는 단순히 이어 붙이는 것만으로 충분하지 않다. 예측 과제에서는 모달리티 사이의 상호작용을 포착해야 하고, 검색 과제에서는 쿼리와 대상이 달라도 비교 가능한 공통 공간이 필요하다.
멀티모달 융합은 여러 모달리티의 표현을 결합해 다운스트림 과제의 시너지를 얻는 기법을 가리킨다. Early·Intermediate·Late Fusion부터 Bilinear/Tensor Fusion, Attention/Gating, MoE 기반 방식까지 선택지가 넓다. 이 가운데 Tensor Fusion은 모달리티 간 외적(outer product)으로 고차 텐서를 구성해 더 풍부한 상호작용을 학습한다.
Cross-Modal Retrieval은 서로 다른 모달리티의 쿼리와 타깃을 하나의 임베딩 공간에 사상한 뒤, 근사 최근접 탐색(ANN)으로 Top-K 결과를 찾는 구조다. InfoNCE·CLIP-style 대조학습, 트리플렛·랭킹 손실, 하드 네거티브 마이닝은 양의 쌍의 정렬(Alignment)과 임베딩 분포의 균일성(Uniformity)을 함께 다루는 데 사용된다.
상호작용 표현의 비용을 제어하는 방법
Tensor Fusion은 2개 이상 모달리티의 조합 특징을 고차 텐서로 표현한다. 외적 기반 구조 덕분에 비선형 상호작용을 학습할 수 있지만, 계산량과 파라미터가 빠르게 커지는 문제가 뒤따른다.
저랭크 분해(Low-Rank Factorization), Bilinear Pooling(MFB/MFH), Factorized Tensor(Decomposed CP/Tucker)는 이 비용을 낮추기 위한 방법이다. 융합 성능을 확보하면서도 모델 크기와 연산량을 통제해야 할 때 저랭크 Tensor Fusion이 선택지가 된다.
검색 모델에서는 모달리티별 인코더와 투영 헤드가 공유 임베딩 공간을 만든다. L2 정규화와 템퍼러처 스케일링을 적용하고, 배치 구성과 하드 네거티브 비율을 조정해 정렬과 균일성의 균형을 맞춘다.
지도 목표인 분류·회귀와 대조학습을 같이 최적화할 수도 있다. 이때 융합 헤드와 검색 헤드를 병렬로 두고, 모달리티 드롭아웃, MixGen/EDA, Batch-hard Mining으로 견고성을 보완한다. FP16+Grad-Scaling과 Gradient Checkpointing은 메모리 압축에 사용된다.
데이터부터 인덱스 교체까지의 운영 경로
텍스트는 토크나이징하고, 이미지는 CNN/ViT 백본으로 처리하며, 오디오는 Mel-Spectrogram으로 정규화한다. 병렬 데이터의 타임스탬프 정렬도 이 단계에 포함된다.
결측 모달리티는 Zero-imputation + Gating, Late Fusion 폴백, Teacher-free Distillation으로 대응할 수 있다. 단일 모달 상태에서도 견딜 수 있도록 설계하지 않으면 실제 입력 조건에서 검색과 예측 품질이 흔들린다.
검색 서빙은 HNSW, IVF-PQ, ScaNN/FAISS 같은 벡터 인덱스를 이용해 대용량 검색 지연을 낮춘다. 캐시와 재랭킹용 크로스 인코더는 정밀도 보완에 쓰인다. 인덱스를 갱신할 때는 스냅샷을 만들고 교체하는 배포 방식이 필요하며, 드리프트 모니터링과 주기적 리빌드로 품질을 유지한다.
융합 전략을 고를 때의 트레이드오프
| 융합 전략 | 성능(정확도/F1) | 확장성(메모리/속도) | 일관성(결측/잡음) | 안정성(NaN/폭주) | 운영 편의(배포/튜닝) |
|---|---|---|---|---|---|
| Early Concat | 중 | 상 | 중 | 상 | 상 |
| Bilinear Pooling(MFB/MFH) | 상 | 중 | 중 | 중 | 중 |
| Tensor Fusion(저랭크) | 상 | 중 | 상 | 중 | 중 |
| Attention/Gating | 상 | 중 | 상 | 중 | 중 |
| Late Fusion(앙상블) | 중 | 상 | 상 | 상 | 상 |
데이터와 백본 수준에 따라 상대 평가는 달라질 수 있다. 최신 정보 확인이 필요하다.
검색과 분류가 함께 쓰이는 장면
이커머스에서는 상품 이미지와 텍스트 설명을 결합해 CTR·전환을 개선하고, 텍스트→이미지와 이미지→텍스트 검색을 제공할 수 있다. 온라인 ANN 인덱스와 배치 리빌드를 함께 운용하며, 재랭킹에는 가격·재고·노출 제약을 반영한다.
미디어와 UGC 관리에서는 영상 썸네일, 자막, 오디오를 공동 임베딩으로 만들 수 있다. 불법·중복 콘텐츠 탐지와 저작권 관리 자동화가 대상이며, 하드 네거티브 마이닝은 유사한 노이즈를 분리하는 데 쓰인다.
의료 영상과 리포트를 매칭하는 경우에는 보고서로 영상을 검색해 유사 사례를 탐색하고, 융합 분류로 병변 검출 민감도를 높일 수 있다. 개인정보 익명화, 감사 로그, 모델 해석성 관리는 이 구조에서 함께 다뤄야 한다.
고객 지원 환경에서는 스크린샷, 에러 로그, 사용자 텍스트를 함께 분석해 해결 문서를 교차 검색할 수 있다. 장애 패턴을 임베딩하면 근본 원인 분석도 가속할 수 있다.
학습과 서빙에서 확인할 운영 조건
학습은 정제된 텍스트·이미지·오디오 병렬 데이터를 타임스탬프에 맞춰 준비하는 데서 시작한다. 모달 인코더를 전이학습하고 Tensor Fusion 또는 저랭크 융합 헤드와 대조학습 헤드를 병렬로 구성한다. 다운스트림 로스(예: Cross-Entropy)와 대조 로스(InfoNCE)의 가중합을 최적화한다.
결측 모달리티에는 Gating 폴백을 적용하고, NaN이 발생하면 배치를 건너뛰며 Grad Clip을 적용한다. OOM 상황에서는 AMP와 Checkpointing을 사용한다.
서빙에서는 신규 콘텐츠를 배치 또는 스트리밍으로 받아 임베딩을 만들고 ANN 인덱스에 증분 추가한다. 주간 스냅샷 재빌드 뒤에는 Top-K 후보를 재랭킹하고 정책 필터링을 거친다. 스냅샷 파일은 원자적으로 교체하고, 롤백 포인트를 남긴 상태에서 카나리 검증 후 트래픽을 전환한다.
PyTorch로 구현하는 저랭크 Tensor Fusion
전제조건은 Python 3.10, PyTorch 2.3+, torchvision 0.18+, faiss 1.7.4이며 GPU는 선택 사항이다. 예시는 CPU에서도 실행할 수 있다.
# pip install torch torchvision
import torch
import torch.nn as nn
import torch.nn.functional as F
class LowRankTensorFusion(nn.Module):
def __init__(self, d_t, d_i, d_a, h=128, rank=8):
super().__init__()
# 각 모달 입력에 bias 1을 concat 예정 -> (d+1)
self.Z_t = nn.Parameter(torch.randn(rank, h, d_t + 1) * 0.02)
self.Z_i = nn.Parameter(torch.randn(rank, h, d_i + 1) * 0.02)
self.Z_a = nn.Parameter(torch.randn(rank, h, d_a + 1) * 0.02)
self.out = nn.Linear(h, h)
def forward(self, t, i, a):
# t: [B, d_t], i: [B, d_i], a: [B, d_a]
B = t.size(0)
t = torch.cat([t, torch.ones(B, 1, device=t.device)], dim=1)
i = torch.cat([i, torch.ones(B, 1, device=i.device)], dim=1)
a = torch.cat([a, torch.ones(B, 1, device=a.device)], dim=1)
# einsum: rank r, hidden h
ht = torch.einsum('rhd,bd->brh', self.Z_t, t) # [B, R, H]
hi = torch.einsum('rhd,bd->brh', self.Z_i, i) # [B, R, H]
ha = torch.einsum('rhd,bd->brh', self.Z_a, a) # [B, R, H]
h = (ht * hi * ha).sum(dim=1) # [B, H]
h = F.relu(self.out(h))
return h
# 예시 입력
B = 4
t = torch.randn(B, 256) # text
i = torch.randn(B, 512) # image
a = torch.randn(B, 128) # audio
fusion = LowRankTensorFusion(256, 512, 128, h=128, rank=8)
y = fusion(t, i, a) # [B, 128]
print(y.shape)
대조학습과 FAISS 인덱스를 연결하는 예시
# pip install torch torchvision faiss-cpu
import torch
import torch.nn as nn
import faiss
import numpy as np
class ProjectionHead(nn.Module):
def __init__(self, in_dim, out_dim=256):
super().__init__()
self.net = nn.Sequential(
nn.Linear(in_dim, out_dim),
nn.ReLU(),
nn.Linear(out_dim, out_dim)
)
self.scale = nn.Parameter(torch.tensor(10.0)) # temperature의 역수
def forward(self, x):
x = self.net(x)
x = nn.functional.normalize(x, dim=-1)
return x, self.scale
def contrastive_loss(z_txt, z_img, scale):
sim = scale * (z_txt @ z_img.T) # [B, B]
labels = torch.arange(sim.size(0), device=sim.device)
loss_t = nn.functional.cross_entropy(sim, labels)
loss_i = nn.functional.cross_entropy(sim.T, labels)
return (loss_t + loss_i) / 2
# dummy encoders
txt_enc = nn.Linear(768, 512)
img_enc = nn.Linear(1024, 512)
txt_proj = ProjectionHead(512, 256)
img_proj = ProjectionHead(512, 256)
B = 64
txt = torch.randn(B, 768)
img = torch.randn(B, 1024)
opt = torch.optim.AdamW(list(txt_enc.parameters()) + list(img_enc.parameters()) +
list(txt_proj.parameters()) + list(img_proj.parameters()), lr=2e-4)
txt_h = txt_enc(txt)
img_h = img_enc(img)
z_t, s_t = txt_proj(txt_h)
z_i, s_i = img_proj(img_h)
loss = contrastive_loss(z_t, z_i, (s_t + s_i)/2)
loss.backward()
nn.utils.clip_grad_norm_(list(txt_enc.parameters()) + list(img_enc.parameters()), max_norm=1.0)
opt.step()
# 인덱스 빌드 (HNSW 예시)
emb_matrix = z_i.detach().cpu().numpy().astype('float32') # gallery (image)
d = emb_matrix.shape[1]
index = faiss.IndexHNSWFlat(d, 32) # M=32
index.hnsw.efConstruction = 200
index.add(emb_matrix)
# 검색 (텍스트 쿼리)
q = z_t.detach().cpu().numpy().astype('float32')
index.hnsw.efSearch = 128
D, I = index.search(q, k=5) # Top-5 인덱스
print(I[0], D[0])
배치 크기, 음수 비율, 온도(=1/scale)를 조정하고, 하드 네거티브는 동일 카테고리 안의 비정합 샘플에서 채집한다. 인덱스는 HNSW(고정밀·온라인 증분), IVF-PQ(대용량·저메모리), 하이브리드 구조(HNSW on PQ)를 검토할 수 있다.
새 인덱스 스냅샷은 빌드한 뒤 원자적으로 교체하고, 헬스체크와 트래픽 전환 사이에 롤백용 스냅샷을 유지한다.
기대할 수 있는 변화와 확장 방향
융합 분류·랭킹은 단일 모달 대비 Accuracy/F1 +38%p, AUC +25%p를 기대할 수 있다. 검색에서는 Recall@10 +1025%p, zero-shot Top-1 +37%p 개선이 가능하다. HNSW 튜닝에서 efSearch 64→128은 mRecall 상향에 쓰이며, IVF-PQ는 메모리를 ~70% 절감할 수 있다.
노이즈와 결측에 대한 예측·검색 일관성을 확보하고, 임베딩과 인덱스를 분리한 아키텍처로 데이터 스케일 증가에 따라 선형에 가까운 확장성을 노릴 수 있다. 멀티태스킹은 라벨 비용 절감과 제로샷 전이성 강화에도 연결된다.
초기에는 Bilinear/Attention 융합과 CLIP-style 대조학습으로 시작한 뒤, 데이터와 트래픽이 늘어나면 Tensor Fusion 및 HNSW/IVF-PQ 하이브리드 구조로 확장할 수 있다.