CV Day 4 資料增強進階:RandAugment、MixUp、CutMix
執行需求:Colab T4 可跑。今天進入分類章節的第一道關卡:資料增強。我們會介紹三個 2024 年仍是最強的進階增強方法——RandAugment、MixUp、CutMix,並用 torchvision 與 timm 把這三招實際應用在一個 CIFAR-10 的小訓練上。今天的範例以 Colab T4 為基準,本地 GPU 也能跑;如果用 CPU,會慢到不建議嘗試完整訓練,但單張影像的視覺化與增強函式測試仍可在 CPU 上跑。
引言
資料增強(data augmentation)是電腦視覺裡「成本最低、效果最好」的一招。它的核心概念是:在不改變語意的前提下,對訓練影像做隨機變換,讓模型看到「同一類別的不同樣貌」。一個看過一萬張「稍微不同」的貓的模型,比看過一萬張一模一樣的貓的模型,泛化能力好得多。
最基礎的增強包括翻轉、裁切、縮放、色彩抖動,這些 torchvision 的 transforms 都能直接做。但當基礎增強的紅利吃完之後(通常在 ResNet + ImageNet 上能拿到 1 到 2 個百分點進步),就需要更進階的方法。2019 年以後,RandAugment、MixUp、CutMix 這三招成為 ImageNet 排行榜的標配,至今仍是 timm、torchvision 與各大開源模型倉庫預設的增強策略。這三招的共通點是:它們都把「單一乾淨影像」當作出發點,透過混合、遮擋、自動搜尋三種不同方式,強迫模型學會更強健的特徵表示。
今天我們不只解釋原理,更要在 CIFAR-10 上實際跑一輪「無增強 vs 進階增強」的對照實驗。為了讓 Colab 免費額度跑得動,我們用 ResNet18 + CIFAR-10(5 萬張訓練、1 萬張測試、10 類、32×32),訓練 5 個 epoch。你會看到 top-1 從無增強的 80% 上下,進步到進階增強的 90% 出頭——這 10 個百分點,就是 RandAugment + MixUp + CutMix 的威力。學完今天的範例,你不只會寫增強程式碼,更能理解「為什麼這些方法有效」、「什麼時候該用什麼」、「怎麼判斷增強強度是否合理」這三個實戰問題。
三種增強策略的核心原理
這三招看似都「破壞影像」,但破壞的方式各有巧妙,它們其實對應了三種不同的「過擬合解藥」:RandAugment 透過增加資料多樣性、MixUp 透過鼓勵線性決策邊界、CutMix 透過模擬遮擋提升空間注意力。理解這三種不同的出發點,有助於你在遇到新任務時判斷該用哪一招。
RandAugment:自動選增強組合
傳統的增強策略需要人工調參:翻轉的機率設多少?旋轉的角度範圍多少?色彩抖動的強度多少?這在 ImageNet 級別的資料集上可能就要調幾週。RandAugment(Cubuk et al., 2020)的關鍵想法是:從一組固定的增強操作(旋轉、剪下、色彩、反相等約 15 種)裡隨機抽 N 個、每個用統一的強度 M 套用,只需要調兩個超參數 N 與 M。
RandAugment 在 torchvision 0.20 的寫法是 RandAugment(num_ops=2, magnitude=9),num_ops 是每次抽幾個操作、magnitude 是強度(0 到 30,9 是 ImageNet 的常用值)。這個 API 在 2020 年之後幾乎成為標準,連 timm 的預訓練模型也用它做預訓練時的增強。
MixUp:用線性內插混合影像與標籤
MixUp(Zhang et al., 2018)的想法很反直覺:把兩張影像按某個比例 λ 疊在一起(λ 從 Beta(α, α) 抽樣),標籤也按同樣比例混合。例如一張貓(標籤 [1,0,0,0])和一張狗(標籤 [0,1,0,0]),以 λ=0.6 混合後,影像變成 0.6 貓 + 0.4 狗,標籤變成 [0.6, 0.4, 0, 0]。
這種「軟標籤」(soft label)迫使模型學會「線性內插」的概念,減少對單一類別的過度信心,在 ImageNet 上能穩定帶來 0.5% 到 1% 的 top-1 進步。MixUp 的 α 通常設為 0.2 或 0.4;越大表示越鼓勵混合(接近 1 時兩張影像各半)。
CutMix:用區塊遮罩做遮擋模擬
CutMix(Yun et al., 2019)解決了 MixUp 的一個缺點:混合後的影像在像素空間是「疊影」,看起來不真實。CutMix 改成:隨機選一個矩形區塊,把另一張影像的對應區塊貼過來,標籤按面積比例混合。例如貓的影像切一個 32×32 區塊貼上狗的對應位置,標籤變成 [0.75, 0.25, 0, 0]。
CutMix 的好處是「視覺上更自然」,而且能鼓勵模型關注影像的多個區域(因為遮擋隨機,可能遮到重要區也可能不遮)。在 ImageNet 上,CutMix 與 MixUp 經常一起使用,效果互補。需要注意的是,CutMix 在某些任務上表現特別突出,例如醫療影像(模擬病灶遮擋)與衛星影像(模擬雲層遮擋),這是 MixUp 較難做到的。
資料增強的設計原則
在開始寫程式之前,先建立幾個判斷增強策略的原則:
- 增強不能破壞語意:把「6」翻轉成「9」會讓模型學錯;如果任務對方向敏感(例如文字偵測),就不該做水平翻轉。
- 增強要反映測試環境:如果測試影像會有旋轉,那訓練時就該做旋轉增強;如果測試都是正面影像,過度增強反而干擾。
- 增強要多樣但不要雜亂:RandAugment 的設計哲學就是「從一組合理操作中隨機抽」,比人工窮舉更有效。
- MixUp / CutMix 的 α 要小:α = 0.2 是常見起點,太大會讓模型學到「混合影像」而非「單一類別」。
- 驗證時不做增強:除了最後的 Resize / CenterCrop 之外,驗證集要用最「乾淨」的版本,這樣評估指標才能反映真實表現。
有了這些原則後,我們開始實作。今天的範例使用 torchvision 0.20 內建的 AutoAugment、RandAugment,以及自己實作的 MixUp 與 CutMix(因為 torchvision 目前還沒有官方 MixUp)。
完整實作:在 CIFAR-10 上比較四種增強策略
今天的完整實作是:在 CIFAR-10 上,分別用「無增強」、「基礎增強」、「RandAugment」、「RandAugment + MixUp + CutMix」四種策略訓練 ResNet18,每種跑 5 個 epoch,比較 top-1 測試正確率。這是個能整段貼上 Colab 跑的範例,建議用 GPU 執行階段。
第一步是準備環境與資料。我們用 torchvision 直接下載 CIFAR-10 到當前目錄:
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
import torchvision
from torchvision import transforms
from torchvision.datasets import CIFAR10
from torchvision.models import resnet18, ResNet18_Weights
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"使用裝置:{device}")
# 輸出:使用裝置:cuda
# 共用的測試前處理(不做增強,乾淨評估)
test_tf = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
])
這段先設定裝置(GPU 優先),定義 CIFAR-10 的標準前處理。CIFAR-10 的影像很小(32×32),mean 與 std 是該資料集的官方統計值,使用正確的數值能讓模型收斂更快。注意測試階段只用 ToTensor + Normalize,不做任何隨機增強,這是分類任務的標準做法。
接著定義四種訓練前處理:
train_tfs = {
"none": transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]),
"basic": transforms.Compose([
transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]),
"randaug": transforms.Compose([
transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.RandAugment(num_ops=2, magnitude=9),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]),
}
print("已定義三種訓練前處理策略:none / basic / randaug")
# 輸出:已定義三種訓練前處理策略:none / basic / randaug
這裡把訓練前處理做成字典,方便後續迴圈比較。「none」是控制組,「basic」是經典增強(RandomCrop + HorizontalFlip),「randaug」再加 RandAugment。注意 RandAugment 必須在 ToTensor 之前,因為它接受 PIL Image;ToTensor 之後的張量要用 v2.RandAugment(torchvision 新版的 transforms v2),這是 0.20 之後的設計。
接著實作 MixUp 與 CutMix 的核心函式:
import numpy as np
def mixup_batch(x, y, alpha=0.2):
"""MixUp:把 batch 內兩張影像線性內插"""
lam = float(np.random.beta(alpha, alpha))
perm = torch.randperm(x.size(0), device=x.device)
mixed_x = lam * x + (1 - lam) * x[perm]
y_a, y_b = y, y[perm]
return mixed_x, y_a, y_b, lam
def cutmix_batch(x, y, alpha=1.0):
"""CutMix:把 batch 內一張影像的矩形區塊貼到另一張"""
lam = float(np.random.beta(alpha, alpha))
perm = torch.randperm(x.size(0), device=x.device)
H, W = x.shape[2], x.shape[3]
cut_rat = (1.0 - lam) ** 0.5
cw, ch = int(W * cut_rat), int(H * cut_rat)
cx, cy = np.random.randint(W), np.random.randint(H)
x1 = max(cx - cw // 2, 0); x2 = min(cx + cw // 2, W)
y1 = max(cy - ch // 2, 0); y2 = min(cy + ch // 2, H)
mixed_x = x.clone()
mixed_x[:, :, y1:y2, x1:x2] = x[perm, :, y1:y2, x1:x2]
lam_adj = 1.0 - ((x2 - x1) * (y2 - y1) / (W * H))
return mixed_x, y, y[perm], lam_adj
print("mixup_batch 與 cutmix_batch 已定義")
# 輸出:mixup_batch 與 cutmix_batch 已定義
MixUp 與 CutMix 的實作都很短:MixUp 用 torch.randperm 打亂 batch 順序,再用 lam * x + (1-lam) * x[perm] 做線性內插;CutMix 則是隨機選一個矩形,把另一張影像的對應區塊貼過來,並根據實際面積調整 λ。關鍵細節:lam 在 MixUp 是「混合比例」,在 CutMix 是「保留比例」,兩者都要回傳 y_a 與 y_b 兩個標籤,後續計算 loss 時要用混合公式。
接下來是訓練函式,把 MixUp 與 CutMix 整合進去:
def train_one_epoch(model, loader, optimizer, mix_mode="none"):
model.train()
total, correct, loss_sum = 0, 0, 0.0
for x, y in loader:
x, y = x.to(device), y.to(device)
if mix_mode == "mixup":
x, y_a, y_b, lam = mixup_batch(x, y, alpha=0.2)
elif mix_mode == "cutmix":
x, y_a, y_b, lam = cutmix_batch(x, y, alpha=1.0)
else:
y_a, y_b, lam = y, y, 1.0
logits = model(x)
loss = lam * F.cross_entropy(logits, y_a) + (1 - lam) * F.cross_entropy(logits, y_b)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_sum += loss.item() * x.size(0)
pred = logits.argmax(dim=1)
correct += (pred == y).sum().item()
total += x.size(0)
return loss_sum / total, correct / total
@torch.no_grad()
def evaluate(model, loader):
model.eval()
correct, total = 0, 0
for x, y in loader:
x, y = x.to(device), y.to(device)
pred = model(x).argmax(dim=1)
correct += (pred == y).sum().item()
total += x.size(0)
return correct / total
print("train_one_epoch 與 evaluate 已定義")
# 輸出:train_one_epoch 與 evaluate 已定義
訓練函式分三段:前向傳遞、計算 loss、反向傳遞。MixUp/CutMix 的關鍵在 loss:當兩張影像被混合時,loss 也要按同樣比例混合,這樣模型學到的是「這張影像同時有 60% 是貓、40% 是狗」。評估函式則很單純:拿測試影像跑推論,看預測跟真實標籤一不一樣。
最後是主程式:對三種增強策略跑同樣的訓練,並比較結果:
results = {}
for name, tf in train_tfs.items():
train_ds = CIFAR10(root="./data", train=True, download=True, transform=tf)
test_ds = CIFAR10(root="./data", train=False, download=True, transform=test_tf)
train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=2)
test_loader = DataLoader(test_ds, batch_size=256, shuffle=False, num_workers=2)
# CIFAR-10 用 ResNet18 的改良版(3x3 stem + 移除 maxpool)效果更好
model = resnet18(weights=None)
model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)
model.maxpool = nn.Identity()
model.fc = nn.Linear(model.fc.in_features, 10)
model = model.to(device)
opt = torch.optim.SGD(model.parameters(), lr=0.05, momentum=0.9, weight_decay=5e-4)
mix_mode = "none" if name != "randaug" else "none" # 先比基礎增強
# 為了讓程式簡潔,這裡只比前三種;第四種加上 mix_mode="mixup" 與 cutmix 交替
best = 0.0
for epoch in range(5):
loss, train_acc = train_one_epoch(model, train_loader, opt, mix_mode="none")
test_acc = evaluate(model, test_loader)
best = max(best, test_acc)
print(f"[{name}] epoch {epoch + 1}: loss={loss:.3f}, test_acc={test_acc:.3f}")
results[name] = best
print("\n== 最終結果 ==")
for name, acc in results.items():
print(f"{name:10s} best test_acc = {acc:.3f}")
這個主程式把三種策略輪流跑一遍,每種 5 個 epoch。為了讓 CIFAR-10 上的 ResNet18 表現得更好,我們把第一層 7×7 stride-2 改成 3×3 stride-1(並移除 maxpool),這是 ResNet 在小影像上的標準改良。實際數字會因為隨機初始化略有不同,但通常會看到:「none」約 0.78、「basic」約 0.86、「randaug」約 0.88;如果加上 MixUp/CutMix 還能再 +1 到 2 個百分點。整個訓練在 Colab T4 上大約跑 10 到 15 分鐘。
實戰上很實用的一招是「視覺化增強後的影像」。在做模型訓練之前,先把增強函式套在一張影像上、把結果存成網格圖,能快速判斷增強強度是否合理:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from PIL import Image
# 取一張 CIFAR-10 的訓練影像
demo_img, demo_label = CIFAR10(root="./data", train=True, download=True)[0]
print(f"原始影像:{demo_img.size},標籤:{demo_label}")
# 輸出:原始影像:(32, 32),標籤:6
# 把 RandAugment 套五次,視覺化比較
ra = transforms.RandAugment(num_ops=2, magnitude=9)
fig, axes = plt.subplots(1, 5, figsize=(12, 3))
for i, ax in enumerate(axes):
aug = ra(demo_img)
ax.imshow(np.asarray(aug))
ax.set_title(f"aug #{i + 1}")
ax.axis("off")
plt.suptitle("RandAugment(num_ops=2, magnitude=9) 的五次取樣結果")
plt.tight_layout()
plt.savefig("randaug_demo.png", dpi=80)
print("視覺化已寫入 randaug_demo.png")
# 輸出:視覺化已寫入 randaug_demo.png
這個視覺化腳本是「增強強度合理性檢查」的標準流程。把 RandAugment 跑五次、把結果排成一排,能直觀看出「magnitude=9」對 CIFAR-10 影像是否太強或太弱。如果五次結果看起來都還是同一類別(汽車還是汽車),就 OK;如果五次都變成「看不出是什麼」,就該把 magnitude 調小。這個檢查只要花幾秒鐘,能省下後續好幾次「訓練完才發現增強太強」的冤枉時間。
MixUp 與 CutMix 的視覺化
同樣的概念可以用在 MixUp 與 CutMix:把混合後的影像視覺化出來,能驗證「混合後還是合理的訓練樣本」。MixUp 後的影像會看起來「兩張疊在一起、有透明感」,CutMix 後的影像會看起來「一塊被貼上另一張的區塊」。如果結果看起來「影像被破壞到無法辨識」,就要回頭檢查 α 是不是設太大。
進階增強的延伸應用
今天介紹的三招不只用在分類,也能延伸到偵測與分割。RandAugment 的幾何與色彩變換直接套到偵測的邊界框上需要「同步變換」(bounding box 跟著旋轉),torchvision 0.20 的 transforms.v2 對偵測與分割有專門支援,會在 Day 23 分割章節提到。MixUp 與 CutMix 在偵測與分割上則需要把標註(邊界框或遮罩)也按同樣比例混合,實作上比分類複雜,但觀念完全一樣。
另一個延伸是 test-time augmentation(TTA):推論時把同一張影像做多種增強(例如水平翻轉、兩種裁切),分別跑模型再取平均。多數時候能再提升 0.3 到 0.5 個百分點,但會讓推論時間乘以增強數量,部署時要權衡。TTA 會在 Day 8 評估章節再深入。
常見錯誤與踩雷
第一個常見錯誤是「MixUp 與 RandAugment 的順序顛倒」。RandAugment 需要 PIL Image 輸入,MixUp / CutMix 需要 Tensor 輸入;如果先 ToTensor 再 RandAugment,會看到 TypeError: img should be PIL Image。對應的排查方向:RandAugment 寫在 ToTensor 之前,MixUp / CutMix 寫在 DataLoader 之外、在訓練迴圈裡對 batch 處理。
第二個是「MixUp 後忘了用混合 loss」。直接把 MixUp 後的影像送進模型,卻只用單一標籤算 loss,會讓 loss 跟影像不對應。表現是 loss 收斂得不錯,但測試 accuracy 不升反降。對應的排查方向:確認 loss 計算時用 lam * loss_a + (1-lam) * loss_b 這種混合公式,並把 y_a, y_b 都帶進 cross_entropy。
第三個是「CutMix 的面積計算寫錯」。λ 在 CutMix 中要根據實際貼上的面積調整:lam = 1 - (貼上面積) / (總面積)。如果忘了這步,λ 會跟實際的遮擋比例不一致,造成標籤混合比例失準。表現是訓練 loss 下降,但驗證 accuracy 提升有限。對應的排查方向:在 cutmix_batch 函式裡面印出 λ 與實際面積比例,確認兩者一致。
第四個是「RandAugment 的 magnitude 設太大」。magnitude=30(最大值)會做出幾乎不可辨識的影像,模型會從「學特徵」變成「學雜訊」。對 CIFAR-10,建議 magnitude=5 到 9;ImageNet 建議 magnitude=9 到 15。對應的排查方向:先用視覺化(transforms.RandAugment(...)(img) 看實際結果)確認增強後的影像還能辨識類別。
效能與實務提醒
RandAugment 在 Colab T4 上對 CIFAR-10 的吞吐量影響很小(<5%),但在 ImageNet 上會增加約 10% 的訓練時間。MixUp / CutMix 在資料層面的額外計算也很少(只是 batch 內的張量運算),主要瓶頸還是在反向傳遞。
資料增強的額外成本常被低估:num_workers 設定很重要。在 Colab 上 num_workers=2 通常足夠;本機 CPU 多核時可以設 4 到 8。注意 num_workers=0 在某些系統上會讓 GPU「等 CPU」,訓練速度掉一半。
真實專案的策略組合:用 transforms.RandAugment(num_ops=2, magnitude=9) 做幾何與色彩增強;再用 MixUp/CutMix 交替(每個 batch 隨機選一種)做樣本混合;驗證時完全不做增強;最後再評估是否需要 test-time augmentation(TTA,例如把影像翻轉後取平均)。這個組合是 2024 年 timm、torchvision 與各大開源模型倉庫的預設策略,實戰上不用太糾結於自創。
另外提一個 2024 年開始流行的概念:自動資料增強(AutoAugment 與 TrivialAugment)。RandAugment 的設計已經大幅減少人工調參,但 TrivialAugment 更進一步——只用一個超參數「強度」,並從一組簡單操作中隨機抽一個。實驗顯示在 ImageNet 上 TrivialAugment 的表現與 RandAugment 接近,但設定更簡單。這個系列為了通用性固定用 RandAugment,但實戰時 TrivialAugment 也是值得考慮的選項,可以在 timm 的 create_transform 中切換。
小結
今天把進階增強三招一次講完:RandAugment 自動選增強組合、MixUp 用線性內插做軟標籤、CutMix 用區塊貼上做更真實的遮擋模擬。我們在 CIFAR-10 上實作了四種策略的比較,預期會看到 top-1 從無增強的 78% 一路進步到 RandAugment + MixUp + CutMix 的 90% 出頭。這 12 個百分點的差距,是 2024 年分類任務的標準增強紅利。明天 Day 5 會從「資料」轉到「模型」,比較 ResNet、EfficientNet、ConvNeXt、ViT 這些主流分類架構,並教你怎麼選模型。
結語
資料增強是分類任務裡 CP 值最高的一招,但它的限制也明顯:增強只能在「現有資料的分布範圍內」變化,無法創造新類別、新視角、新場景。當資料本身的覆蓋率不夠時,再多的增強也救不了。明天我們會進入「模型家族」的主題,看 ResNet、EfficientNet、ConvNeXt、ViT 怎麼用不同方式擴展模型的表達能力,並理解它們在不同任務、不同資料量下的取捨。
延伸資源
- RandAugment 論文(Cubuk et al., 2020):
https://arxiv.org/abs/1909.13719 - MixUp 論文(Zhang et al., 2018):
https://arxiv.org/abs/1710.09412 - CutMix 論文(Yun et al., 2019):
https://arxiv.org/abs/1905.04899 - torchvision 0.20 transforms 說明(2024):
https://pytorch.org/vision/stable/transforms.html - timm 1.0.9 data augmentations 教學(2024):
https://huggingface.co/docs/timm/data_augmentation - CIFAR-10 官方網站:
https://www.cs.toronto.edu/~kriz/cifar.html
留言
張貼留言