多模態學習(Multimodal Learning)是一種整合來自多種不同類型資料的學習方式,其核心概念在於讓模型同時接收並融合來自不同模態(modality)的資訊,例如圖像與文字、聲音與影像、感測器數值與語音等,從而做出比僅使用單一模態更完整的判斷。在現實世界中,人類對環境的感知本來就是多模態的,而多模態學習正是試圖讓神經網路以類似的方式運作。

組合多種模態的方式之一,是將它們都用來分類。下列是一個使用 MNIST 資料集的範例,其中影像的部分是遮去部分內容後的資料集影像,文字部分則為了範例簡單,是用 one-hot 向量加上噪音干擾;在範例中你將看到,單獨使用一種模態的時候,因為資訊被遮住或干擾,因此得到的準確度可能較低,但是使用兩種模態時,因為資訊有可能互補,故而準確度有機會提升:

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

transform = transforms.ToTensor()
train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
test_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=512, shuffle=True)
test_loader = DataLoader(test_set, batch_size=512, shuffle=False)

IMAGE_KEEP_ROWS = 10
TEXT_NOISE_STD = 0.5
EMBED_DIM = 64
EPOCHS = 5


def mask_image(x):
	x = x.clone()
	x[:, :, IMAGE_KEEP_ROWS:, :] = 0.0
	return x


def noisy_text(labels, std=TEXT_NOISE_STD):
	one_hot = torch.zeros(labels.size(0), 10, device=labels.device)
	one_hot.scatter_(1, labels.unsqueeze(1), 1.0)
	return one_hot + torch.randn_like(one_hot) * std


class ImageEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Flatten(),
			nn.Linear(28 * 28, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
			nn.ReLU(),
		)

	def forward(self, x):
		return self.net(x)


class TextEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Linear(10, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
			nn.ReLU(),
		)

	def forward(self, x):
		return self.net(x)


class ImageOnly(nn.Module):
	def __init__(self):
		super().__init__()
		self.encoder = ImageEncoder()
		self.classifier = nn.Linear(EMBED_DIM, 10)

	def forward(self, x):
		return self.classifier(self.encoder(x))


class TextOnly(nn.Module):
	def __init__(self):
		super().__init__()
		self.encoder = TextEncoder()
		self.classifier = nn.Linear(EMBED_DIM, 10)

	def forward(self, x):
		return self.classifier(self.encoder(x))


class Fusion(nn.Module):
	def __init__(self):
		super().__init__()
		self.img_encoder = ImageEncoder()
		self.txt_encoder = TextEncoder()
		self.classifier = nn.Linear(EMBED_DIM * 2, 10)

	def forward(self, img, txt):
		return self.classifier(torch.cat([self.img_encoder(img), self.txt_encoder(txt)], dim=1))


def count_params(model):
	return sum(p.numel() for p in model.parameters())


def evaluate(model, mode):
	model.eval()
	correct = total = 0
	with torch.no_grad():
		for x, y in test_loader:
			x, y = x.to(device), y.to(device)
			if mode == "img":
				pred = model(mask_image(x)).argmax(dim=1)
			elif mode == "txt":
				pred = model(noisy_text(y)).argmax(dim=1)
			else:
				pred = model(mask_image(x), noisy_text(y)).argmax(dim=1)
			correct += (pred == y).sum().item()
			total += y.size(0)
	return 100 * correct / total


criterion = nn.CrossEntropyLoss()

img_model = ImageOnly().to(device)
print(f"ImageOnly #param: {count_params(img_model)}")
optimizer = optim.Adam(img_model.parameters(), lr=1e-3)
img_model.train()
for epoch in range(EPOCHS):
	print(f'\tEpoch {epoch+1}/{EPOCHS}')
	for x, y in train_loader:
		x, y = x.to(device), y.to(device)
		optimizer.zero_grad()
		criterion(img_model(mask_image(x)), y).backward()
		optimizer.step()
img_model.eval()
print(f"ImageOnly Test Accuracy: {evaluate(img_model, 'img'):.2f}%")

txt_model = TextOnly().to(device)
print(f"TextOnly #param: {count_params(txt_model)}")
optimizer = optim.Adam(txt_model.parameters(), lr=1e-3)
txt_model.train()
for epoch in range(EPOCHS):
	print(f'\tEpoch {epoch+1}/{EPOCHS}')
	for x, y in train_loader:
		x, y = x.to(device), y.to(device)
		optimizer.zero_grad()
		criterion(txt_model(noisy_text(y)), y).backward()
		optimizer.step()
txt_model.eval()
print(f"TextOnly Test Accuracy: {evaluate(txt_model, 'txt'):.2f}%")

fus_model = Fusion().to(device)
print(f"Fusion #param: {count_params(fus_model)}")
optimizer = optim.Adam(fus_model.parameters(), lr=1e-3)
fus_model.train()
for epoch in range(EPOCHS):
	print(f'\tEpoch {epoch+1}/{EPOCHS}')
	for x, y in train_loader:
		x, y = x.to(device), y.to(device)
		optimizer.zero_grad()
		criterion(fus_model(mask_image(x), noisy_text(y)), y).backward()
		optimizer.step()
fus_model.eval()
print(f"Fusion Test Accuracy: {evaluate(fus_model, 'img+txt'):.2f}%")

在上述範例中,我們是將影像和文字資料分別編碼後,再一起餵給 classification head,然後才計算 loss,這種方式通常稱為 early fusion;而根據準確度及執行效率等各種需求,也有人會將多種模態各別做分類並計算 loss 以後,再用平均等方式組合,這種方式通常稱為 late fusion;你也可以根據自己的好奇心或需求,嘗試不同的融合方式。

多模態學習的另一種經典場景是對比式學習(contrastive learning)。這種做法的目的是讓同類樣本在空間中的距離互相靠近,而不同類的樣本則互相拉遠。我們一樣使用 MNIST 示範如下:

import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

transform = transforms.ToTensor()
train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
test_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=512, shuffle=True)
test_loader = DataLoader(test_set, batch_size=512, shuffle=False)

EMBED_DIM = 2
EPOCHS = 10
TEMPERATURE = 0.1


def make_text(labels):
	one_hot = torch.zeros(labels.size(0), 10, device=labels.device)
	one_hot.scatter_(1, labels.unsqueeze(1), 1.0)
	return one_hot


class ImageEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Flatten(),
			nn.Linear(28 * 28, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
		)

	def forward(self, x):
		return self.net(x)


class TextEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Linear(10, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
		)

	def forward(self, x):
		return self.net(x)


def count_params(model):
	return sum(p.numel() for p in model.parameters())


def info_nce_loss(img_emb, txt_emb, temperature=TEMPERATURE):
	img_emb = F.normalize(img_emb, dim=1)
	txt_emb = F.normalize(txt_emb, dim=1)
	logits = img_emb @ txt_emb.T / temperature
	labels = torch.arange(logits.size(0), device=logits.device)
	loss_i2t = F.cross_entropy(logits, labels)
	loss_t2i = F.cross_entropy(logits.T, labels)
	return (loss_i2t + loss_t2i) / 2


def get_embeddings(img_enc, txt_enc):
	img_enc.eval()
	txt_enc.eval()
	all_img, all_txt, all_labels = [], [], []
	with torch.no_grad():
		for x, y in test_loader:
			x, y = x.to(device), y.to(device)
			all_img.append(F.normalize(img_enc(x), dim=1))
			all_txt.append(F.normalize(txt_enc(make_text(y)), dim=1))
			all_labels.append(y)
	return (
		torch.cat(all_img).cpu().numpy(),
		torch.cat(all_txt).cpu().numpy(),
		torch.cat(all_labels).cpu().numpy(),
	)


def recall_at_1(img_emb, txt_emb, labels):
	sim = torch.tensor(img_emb) @ torch.tensor(txt_emb).T
	top1 = sim.argmax(dim=1).numpy()
	return (labels[top1] == labels).mean()


img_enc = ImageEncoder().to(device)
txt_enc = TextEncoder().to(device)
print(f"ImageEncoder #param: {count_params(img_enc)}")
print(f"TextEncoder #param: {count_params(txt_enc)}")

optimizer = optim.Adam(list(img_enc.parameters()) + list(txt_enc.parameters()), lr=1e-3)

for epoch in range(EPOCHS):
	img_enc.train()
	txt_enc.train()
	total_loss = 0
	for x, y in train_loader:
		x, y = x.to(device), y.to(device)
		optimizer.zero_grad()
		loss = info_nce_loss(img_enc(x), txt_enc(make_text(y)))
		loss.backward()
		optimizer.step()
		total_loss += loss.item()
	print(f"Epoch {epoch+1:2d} | Loss: {total_loss / len(train_loader):.4f}")

all_img, all_txt, all_labels = get_embeddings(img_enc, txt_enc)
print(f"Recall@1: {recall_at_1(all_img, all_txt, all_labels):.4f}")

samples_per_class = 100
idx = np.concatenate([
	np.where(all_labels == c)[0][:samples_per_class] for c in range(10)
])
img_sub = all_img[idx]
txt_sub = all_txt[idx]
labels_sub = all_labels[idx]

colors = plt.cm.tab10(np.linspace(0, 1, 10))
plt.figure()
for c in range(10):
	mask = labels_sub == c
	plt.scatter(
		img_sub[mask, 0], img_sub[mask, 1],
		marker="o", facecolors="none", edgecolors=colors[c], label=str(c), linewidths=1, s=15, alpha=0.7)
	plt.scatter(
		txt_sub[mask, 0], txt_sub[mask, 1],
		marker="x", color=colors[c], linewidths=1, s=15
	)
plt.legend()
plt.show()

在上述範例中:

如果在你的應用場景中,多種模態並不總是齊全,會一下缺東西下缺西的話,則經典的處理方式之一,是在訓練時用隨機遮蔽的方式模擬模態缺失;而缺的那一部分,則通常就全部填 0 再餵入模型。一個簡單的範例如下:

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

transform = transforms.ToTensor()
train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
test_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=512, shuffle=True)
test_loader = DataLoader(test_set, batch_size=512, shuffle=False)

IMAGE_KEEP_ROWS = 10
TEXT_NOISE_STD = 0.5
EMBED_DIM = 64
EPOCHS = 5
DROP_PROB = 0.5


def mask_image(x):
	x = x.clone()
	x[:, :, IMAGE_KEEP_ROWS:, :] = 0.0
	return x


def noisy_text(labels, std=TEXT_NOISE_STD):
	one_hot = torch.zeros(labels.size(0), 10, device=labels.device)
	one_hot.scatter_(1, labels.unsqueeze(1), 1.0)
	return one_hot + torch.randn_like(one_hot) * std


def drop_modality(img, txt, drop_prob=DROP_PROB):
	img = img.clone()
	txt = txt.clone()
	b = img.size(0)
	drop_img = torch.rand(b, device=img.device) < drop_prob
	drop_txt = torch.rand(b, device=txt.device) < drop_prob
	both_dropped = drop_img & drop_txt
	keep_img = torch.rand(b, device=img.device) < 0.5
	drop_img[both_dropped] = ~keep_img[both_dropped]
	drop_txt[both_dropped] = keep_img[both_dropped]
	img[drop_img] = 0.0
	txt[drop_txt] = 0.0
	return img, txt


class ImageEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Flatten(),
			nn.Linear(28 * 28, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
			nn.ReLU(),
		)

	def forward(self, x):
		return self.net(x)


class TextEncoder(nn.Module):
	def __init__(self):
		super().__init__()
		self.net = nn.Sequential(
			nn.Linear(10, 256),
			nn.ReLU(),
			nn.Linear(256, 128),
			nn.ReLU(),
			nn.Linear(128, EMBED_DIM),
			nn.ReLU(),
		)

	def forward(self, x):
		return self.net(x)


class Fusion(nn.Module):
	def __init__(self):
		super().__init__()
		self.img_encoder = ImageEncoder()
		self.txt_encoder = TextEncoder()
		self.classifier = nn.Linear(EMBED_DIM * 2, 10)

	def forward(self, img, txt):
		return self.classifier(torch.cat([self.img_encoder(img), self.txt_encoder(txt)], dim=1))


def count_params(model):
	return sum(p.numel() for p in model.parameters())


def evaluate(model, mode):
	model.eval()
	correct = total = 0
	with torch.no_grad():
		for x, y in test_loader:
			x, y = x.to(device), y.to(device)
			xi = mask_image(x)
			t = noisy_text(y)
			if mode == "img_only":
				t = torch.zeros_like(t)
			elif mode == "txt_only":
				xi = torch.zeros_like(xi)
			pred = model(xi, t).argmax(dim=1)
			correct += (pred == y).sum().item()
			total += y.size(0)
	return 100 * correct / total


model = Fusion().to(device)
print(f"Fusion #param: {count_params(model)}")
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

model.train()
for epoch in range(EPOCHS):
	print(f'\tEpoch {epoch+1}/{EPOCHS}')
	for x, y in train_loader:
		x, y = x.to(device), y.to(device)
		xi, t = drop_modality(mask_image(x), noisy_text(y))
		optimizer.zero_grad()
		criterion(model(xi, t), y).backward()
		optimizer.step()

model.eval()
print(f"Both modalities: {evaluate(model, 'both'):.2f}%")
print(f"Image only     : {evaluate(model, 'img_only'):.2f}%")
print(f"Text only      : {evaluate(model, 'txt_only'):.2f}%")

在上述範例中,為了與第一個範例做比較,所以隨機遮蔽資料又被加回來了。你可以看到雙模態資料皆齊全時的效果比較好,而某一模態缺失時,則效果倒退回相當於只有使用單模態做辨識的結果。雖然就效果來說,此範例跟第一個範例差不多,但使用雙模態的好處在於,我們不需要訓練三顆不同的模型。

多模態學習還有其他許多經典方式與應用,此處簡單說明如下: