数据增强:简单变换不够用了之后
数据增强:简单变换不够用了之后一旦进项目,好看的架构图就没那么管用了。
但这桥没那么好走,我踩过不少坑,也观察过别人怎么踩坑,下面把过程整理一下。
从图像增强开始踩坑
简单的几何变换
最早接触图像增强时,我以为就是左翻转右翻转、上下颠倒、随机旋转。代码写起来确实简单:
import torchvision.transforms as transforms
from PIL import Image
transform = transforms.Compose([
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomVerticalFlip(p=0.3),
transforms.RandomRotation(degrees=15),
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
])
image = Image.open('sample.jpg')
augmented_images = [transform(image) for _ in range(5)]
跑起来是没问题的,但很快发现几个坑。
第一个坑出现在旋转上。某次做医学影像分类时,旋转了 45 度,结果模型学习时总在一个奇怪的角落转圈。后来查了半天才发现,旋转后图像边缘会留黑边,而这些黑边被模型当成了特征。解决方案很简单:
transform = transforms.Compose([
transforms.RandomRotation(degrees=15, fill=(0, 0, 0)), # 明确指定填充
transforms.Resize((224, 224)), # 统一尺寸,避免黑边比例变化
])
第二个坑是翻转的概率设置。最初我把所有翻转概率都设为 0.5,觉得反正随机。但有些场景下,翻转是有意义的——比如手势识别,左右翻转后手势含义完全变了。最后改成:
# 手势识别任务中,只对无关紧要的样本翻转
def smart_flip(image, label):
if label in ['neutral', 'ok', 'thumbs_up']:
if random.random() < 0.5:
image = image.transpose(Image.FLIP_LEFT_RIGHT)
return image, label
颜色增强的坑
颜色增强看着像是"调个参数就行",但实际用起来需要调得比较克制。我最开始把 ColorJitter 的参数设得挺大,以为这样能让模型更鲁棒,结果训练出来模型在测试集上性能反而下降了。
原因很简单:增强过度,正样本被改得面目全非,模型学不到真正有用的颜色特征。后来我看了一些论文和实践经验,发现颜色增强的关键是"轻度、多种":
color_transform = transforms.Compose([
transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.05),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
参数小一点,但组合起来效果更好。而且一定要记住,这些增强都应该在归一化之前做,不然会破坏分布。
混合与切分
Mixup 和 CutMix 是后来才接触到的增强方法。核心思路是把两张图片的信息混合起来,让模型学习到"样本之间的过渡"。
Mixup 的实现:
def mixup_data(x, y, alpha=1.0):
if alpha > 0:
lam = np.random.beta(alpha, alpha)
else:
lam = 1
batch_size = x.size()[0]
index = torch.randperm(batch_size)
mixed_x = lam * x + (1 - lam) * x[index]
y_a, y_b = y, y[index]
return mixed_x, y_a, y_b, lam
CutMix 则是图像区域级别的混合:
def cutmix_data(x, y, beta=1.0):
lam = np.random.beta(beta, beta)
batch_size = x.size()[0]
index = torch.randperm(batch_size)
bbx1, bby1, bbx2, bby2 = rand_bbox(x.size(), lam)
x[:, :, bbx1:bbx2, bby1:bby2] = x[index, :, bbx1:bbx2, bby1:bby2]
lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (x.size()[-1] * x.size()[-2]))
y_a, y_b = y, y[index]
return x, y_a, y_b, lam
这两个方法效果好,但训练时要调整 loss 函数:
def mixup_criterion(criterion, pred, y_a, y_b, lam):
return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)
我的实践建议是:先不用这些方法,基础模型稳定了再尝试。Mixup 会让训练曲线看着不太正常,收敛变慢,不要慌。
文本增强的坑
同义词替换
文本增强比图像增强要麻烦得多,因为文字不像图像,稍微改一下意思可能就全变了。
最早我用的是简单的同义词替换:
from nltk.corpus import wordnet
import random
def synonym_replacement(text, n=2):
words = text.split()
new_words = words.copy()
random_word_list = list(set([word for word in words if wordnet.synsets(word)]))
random.shuffle(random_word_list)
num_replaced = 0
for random_word in random_word_list:
synonyms = wordnet.synsets(random_word)
if len(synonyms) >= 1:
synonym = synonyms[0].lemmas()[0].name()
new_words = [synonym if word == random_word else word for word in new_words]
num_replaced += 1
if num_replaced >= n:
break
return ' '.join(new_words)
问题很快来了:替换后的句子有时很怪。比如"我喜欢喝咖啡"可能被改成"我爱好喝咖啡",意思还能接受,但"这个项目很重要"改成"这个项目很紧要"就有点生硬。
后来我发现问题在于同义词词典太粗糙,而且没有上下文判断。现在更推荐使用基于上下文的替换:
from transformers import pipeline
generator = pipeline("text2text-generation", model="google/flan-t5-small")
def contextual_replacement(text):
prompt = f"Rewrite this sentence with different words but keep the meaning: {text}"
result = generator(prompt, max_length=len(text.split()) + 10, num_return_sequences=3)
return [item['generated_text'] for item in result]
这样生成的文本会更自然,但代价是速度变慢、需要外部模型。算力紧张时,还是要回到简单的同义词替换,只是要更保守。
回译
回译是另一个流行的文本增强方法。思路是把中文翻译成英文,再翻译回中文,这样能得到一个意思相近但表达不同的句子:
from googletrans import Translator
translator = Translator()
def back_translate(text, src='zh-cn', temp_lang='en'):
# 先翻译到临时语言
temp_translation = translator.translate(text, src=src, dest=temp_lang).text
# 再翻译回来
back_translation = translator.translate(temp_translation, src=temp_lang, dest=src).text
return back_translation
这个方法效果好,但也有坑。最明显的是免费 API 有调用限制,而且翻译质量不稳定。我发现最好在模型训练前把所有文本都回译一遍存下来,而不是实时调用。
另一个问题是对短文本效果不佳。“好的"这个词回译后可能变成"行"或者"可以”,但"好的好"这种语境信息就丢了。
随机删除和交换
随机删除和随机交换是更简单粗暴的增强方法:
def random_deletion(text, p=0.1):
words = text.split()
if len(words) == 1:
return text
new_words = [word for word in words if random.random() > p]
if len(new_words) == 0:
return random.choice(words)
return ' '.join(new_words)
def random_swap(text, n=1):
words = text.split()
for _ in range(n):
if len(words) < 2:
break
idx1, idx2 = random.sample(range(len(words)), 2)
words[idx1], words[idx2] = words[idx2], words[idx1]
return ' '.join(words)
这两个方法看似简单,但需要小心。删除概率太大,句子可能失去关键信息;交换次数太多,句子可能变得不知所云。我的建议是:删除概率不要超过 0.1,交换次数不要超过 1 次。
生成式增强
GAN 生成样本
当传统增强方法都不够用时,生成式方法上场了。GAN 是最早用的,用起来要麻烦一点:
import torch
import torch.nn as nn
class SimpleGenerator(nn.Module):
def __init__(self, latent_dim, img_shape):
super().__init__()
self.img_shape = img_shape
def block(in_feat, out_feat, normalize=True):
layers = [nn.Linear(in_feat, out_feat)]
if normalize:
layers.append(nn.BatchNorm1d(out_feat, 0.8))
layers.append(nn.LeakyReLU(0.2, inplace=True))
return layers
self.model = nn.Sequential(
*block(latent_dim, 128, normalize=False),
*block(128, 256),
*block(256, 512),
*block(512, 1024),
nn.Linear(1024, int(np.prod(img_shape))),
nn.Tanh()
)
def forward(self, z):
img = self.model(z)
img = img.view(img.size(0), *self.img_shape)
return img
训练 GAN 的坑比训练正常分类模型还要多。最常见的是模式崩溃,生成器一直生成同一种图片。这个问题的解决方案比较多,但都不完美:
# 添加噪声的判别器
class NoisyDiscriminator(nn.Module):
def __init__(self, img_shape):
super().__init__()
self.model = nn.Sequential(
nn.Linear(int(np.prod(img_shape)), 512),
nn.LeakyReLU(0.2, inplace=True),
nn.Dropout(0.3), # 添加 dropout
nn.Linear(512, 256),
nn.LeakyReLU(0.2, inplace=True),
nn.Dropout(0.3),
nn.Linear(256, 1),
nn.Sigmoid()
)
def forward(self, img):
img_flat = img.view(img.size(0), -1)
validity = self.model(img_flat)
return validity
还有训练不稳定的问题,判别器太强或太弱都不行。我的经验是:不要指望 GAN 能完美工作,它更像是一个"能生成一些可用的样本"的工具,而不是精确控制的样本生成器。
Diffusion 模型
最近开始尝试用 Diffusion 模型做数据增强,这个比 GAN 更稳定一些:
from diffusers import StableDiffusionPipeline
import torch
pipe = StableDiffusionPipeline.from_pretrained(
"runwayml/stable-diffusion-v1-5",
torch_dtype=torch.float16,
).to("cuda")
def generate_samples(prompt, num_samples=4):
images = []
for _ in range(num_samples):
image = pipe(prompt, num_inference_steps=20, guidance_scale=7.5).images[0]
images.append(image)
return images
用于数据增强时,关键是设计好 prompt。比如要做手写数字增强,prompt 应该尽量具体:
# 手写数字增强
digit_prompt = "handwritten digit {digit}, clear handwriting, black ink on white background, isolated, centered"
samples = generate_samples(digit_prompt.format(digit=3), num_samples=10)
Diffusion 模型的问题是算力需求大,而且生成的样本质量受 prompt 影响很大。我通常会在训练前批量生成好一批样本,而不是实时生成。
条件生成
更精细的增强是使用条件生成,即根据原始样本生成增强版本:
from transformers import BlipProcessor, BlipForConditionalGeneration
from PIL import Image
processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-base")
model = BlipForConditionalGeneration.from_pretrained("Salesforce/blip-image-captioning-base")
def caption_to_variation(image_path, variation_count=3):
image = Image.open(image_path).convert('RGB')
inputs = processor(image, return_tensors="pt")
# 生成描述
out = model.generate(**inputs)
caption = processor.decode(out[0], skip_special_tokens=True)
# 基于描述生成变体
variations = []
for i in range(variation_count):
variation_prompt = f"{caption}, variation {i+1}, different lighting and angle"
variation = pipe(variation_prompt, num_inference_steps=15).images[0]
variations.append(variation)
return variations
这种方法效果比较好,但计算开销也更大。适合数据集很小但算力充足的场景。
实践中的判断
什么时候该增强
不是所有场景都需要数据增强。我的经验是:
- 样本数少于 1000 时,先考虑收集更多真实数据
- 样本数在 1000-10000 时,开始考虑增强
- 样本数超过 10000 时,增强更多是为了鲁棒性而非数量
还有一个判断维度是数据多样性。如果数据已经很多样化(不同场景、不同角度、不同光照),增强的收益会降低。
增强程度的把握
增强过度的危害有时比不增强还大。我见过有人把旋转角度设成 90 度,结果模型学到了"旋转 90 度后还能分类"的无用能力。
一般原则是:增强后的样本应该仍然能被人类识别。如果人类都看不出来是什么,模型也学不到有用的东西。
# 安全的图像增强参数
safe_augmentation = transforms.Compose([
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomRotation(degrees=10), # 角度不要太大
transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.15), # 变化幅度温和
transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)), # 平移比例控制在 10% 以内
])
监控增强效果
增强不是做完就完了,要监控效果。我的做法是:
- 在验证集上对比有无增强的性能
- 可视化增强后的样本,确保合理性
- 检查模型是否学到了增强带来的偏差(比如总认为左边翻转的样本更常见)
# 检查增强分布
def check_augmentation_bias(dataset, transform, num_samples=1000):
transformed_labels = []
for i in range(min(num_samples, len(dataset))):
image, label = dataset[i]
transformed = transform(image)
transformed_labels.append(label)
label_counts = Counter(transformed_labels)
original_counts = Counter([dataset[i][1] for i in range(min(num_samples, len(dataset)))])
# 比较分布是否一致
for label in set(list(label_counts.keys()) + list(original_counts.keys())):
print(f"Label {label}: original={original_counts[label]}, augmented={label_counts[label]}")
一些教训
不要过度依赖增强
数据增强能缓解数据不足的问题,但不能解决根本。真正解决数据不足的方式还是收集更多真实数据。增强更像是"让现有数据发挥更大价值"的技巧,而不是替代品。
我见过有人只用了几百张样本,加上大量的增强,就指望模型能工作。结果模型在真实场景下表现很差,因为增强出来的分布和真实场景还是有差距。
增强方法要适配任务
不是所有增强方法都适合所有任务。比如医学影像诊断中,随意旋转图像可能改变病理特征;而自然场景分类中,旋转则更安全。
文本也是如此。情感分析任务中,简单的同义词替换可能改变情感色彩;而文本分类任务中,同义词替换则更安全。
要有基线对比
每次引入新的增强方法,都要有基线对比。否则你不知道性能提升是因为增强有效,还是因为其他原因(比如运气好、训练时间长)。
我的做法是:先训练一个只有基础数据增强的模型,得到基线性能;然后再尝试高级增强方法,对比性能变化。
收尾
数据增强不是什么魔法技巧,它更像是在有限资源下的妥协。简单的方法往往最可靠,复杂的方法需要更多调参和验证。
真正重要的是理解你的数据和任务,选择合适的增强策略,而不是盲目堆砌方法。有时候,最好的增强就是不做增强,而是好好整理和分析现有数据。
最后说一句:数据增强能解决一些问题,但解决不了所有问题。模型学不好,可能不是数据不够,而是模型架构不对、特征工程不够、或者训练策略有问题。这些问题用增强是治不好的。
版权声明: 本文首发于 指尖魔法屋-数据增强:简单变换不够用了之后(https://blog.thinkmoon.cn/post/175-data-augmentation-from-simple-to-generative/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。