多模態學習(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()在上述範例中:
- 因為此處示範的不是「資訊互補」,並且為了視覺化所設定的二維空間,本身對資料的分辨能力可能不足,所以此範例沒有把資料做遮罩或干擾。
- 此處使用的 loss function 是 NT-Xent (Normalized Temperature-scaled Cross Entropy) loss,由 Chen et al. 於 2020 年在 SimCLR 論文中提出,是對比學習領域最廣泛使用的 loss function 之一。其核心概念是:在一個 batch 內,每筆資料的圖像與對應文字構成正樣本對,同一 batch 內其餘所有配對則視為負樣本;loss 的目標是讓正樣本對的相似度高於所有負樣本對。溫度參數則代表模型對差異的敏感度,溫度愈小則模型愈敏感,訓練訊號也愈強,但愈容易不穩定。
- 輸出的 embedding 會落在單位圓上,是因為我們在繪製之前,先做了正規化。此外,為了視覺效果等考量,每個類別只取 100 個投影點來繪製。
- 對比式學習本身跟多模態沒有必然關聯,也可以用在單一模態上,例如人臉辨識中,同一人的不同照片要當作正樣本,而不同人的照片要當做負樣本。
如果在你的應用場景中,多種模態並不總是齊全,會一下缺東西下缺西的話,則經典的處理方式之一,是在訓練時用隨機遮蔽的方式模擬模態缺失;而缺的那一部分,則通常就全部填 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}%")在上述範例中,為了與第一個範例做比較,所以隨機遮蔽資料又被加回來了。你可以看到雙模態資料皆齊全時的效果比較好,而某一模態缺失時,則效果倒退回相當於只有使用單模態做辨識的結果。雖然就效果來說,此範例跟第一個範例差不多,但使用雙模態的好處在於,我們不需要訓練三顆不同的模型。
多模態學習還有其他許多經典方式與應用,此處簡單說明如下:
- Cross-modal Attention:本篇示範的融合方式,都是讓兩個模態各自獨立編碼後才合併;而 cross-modal attention 則是讓一個模態的 embedding 作為 query,去查詢另一個模態的 key 與 value,使兩個模態在 encoder 中段就能互相交換資訊,是目前視覺語言模型最常採用的融合機制。
- CLIP(Contrastive Language-Image Pretraining):由 OpenAI 於 2021 年提出,訓練概念與本篇第二個範例相同,是以對比學習讓影像與文字 embedding 對齊,但規模擴展至四億個網路圖文配對;且大規模訓練帶來天然的零樣本泛化能力,你不太需要針對特定任務微調,也可以得到不錯的分類效果。
- 多模態生成:Stable Diffusion 等模型,會以 Diffusion Model 作為影像生成骨幹,並以 CLIP 的文字 encoder 提供條件訊號,再透過 cross-attention 將文字描述注入 U-Net 的各解碼層,來達成依據文字生成影像的效果。其中,條件注入的方式與 CVAE 的類別條件概念相通,只是條件從離散標籤擴展為連續的文字 embedding。