Adversarial training для текстовых классификаторов на неразмеченных данных
Adversarial и virtual adversarial training поднимают качество текстовой классификации без ручной разметки. Разбираем, как эти методы устроены и когда за них стоит браться.

Зачем добавлять шум специально
У вас есть 500 размеченных примеров и 50 000 неразмеченных. Классическая история: разметка стоит денег и времени, а сырого текста хоть отбавляй. Adversarial training и его расширение для полуразметки — virtual adversarial training (VAT) — позволяют выжать пользу из неразмеченного массива, не нанимая разметчиков.
Идея контринтуитивная. Вместо того чтобы делать модель устойчивой к случайному шуму, вы генерируете шум специально в том направлении, где модель сильнее всего ошибается, и заставляете её справляться именно с ним. В компьютерном зрении это возмущение пикселей. В тексте — сдвиг эмбеддингов слов, потому что сами токены дискретны и «чуть-чуть изменить букву» смысла не имеет.
Подход описан в работе Миято, Дай и Гудфеллоу Adversarial Training Methods for Semi-Supervised Text Classification (2016–2017). Возмущение накладывается не на входные слова, а на их векторные представления — и уже на этом уровне модель учится быть устойчивой.
Adversarial training: что это на уровне формул
Обычная supervised-модель минимизирует функцию потерь на размеченных данных. Adversarial training добавляет второе слагаемое: потери на намеренно испорченном входе. Возмущение r выбирается так, чтобы максимально увеличить ошибку, но при этом оставаться в пределах небольшой нормы (обычно L2).
Точное направление найти дорого, поэтому его аппроксимируют по градиенту потерь: берут градиент по эмбеддингам, нормируют и масштабируют на маленький коэффициент epsilon. Получается «худшее из близких» возмущение, которое считается за один дополнительный проход градиента.
Virtual adversarial training: тот же приём без меток
Ключевой шаг для полуразметки. AT требует истинную метку, чтобы посчитать ошибку. VAT её не требует. Вместо ошибки относительно правильного ответа он минимизирует расхождение между предсказанием на чистом входе и предсказанием на возмущённом — то есть заставляет модель давать стабильный ответ в окрестности каждой точки.
Мера расхождения — обычно KL-дивергенция1 между двумя распределениями вероятностей классов. Поскольку истинная метка не нужна, VAT работает на любом неразмеченном тексте. Именно это и делает связку supervised loss (на размеченных) плюс VAT loss (на всех) полуконтролируемым обучением.
Модель наказывают не за неправильный ответ, а за то, что её ответ дрожит от микроскопического сдвига входа. Устойчивость к такому дрожанию оказывается сильным сигналом даже без единой метки.
Три близких, но разных метода
| Метод | Нужна метка | Что минимизирует | Где применим |
|---|---|---|---|
| Adversarial training (AT) | Да | Потери на возмущённом входе относительно истинной метки | Только размеченные данные |
| Virtual adversarial (VAT) | Нет | Расхождение предсказаний чистый vs возмущённый вход | Размеченные и неразмеченные |
| Random perturbation (базлайн) | Нет | Расхождение при случайном шуме | Любые, но слабее VAT |
Случайный шум в таблице — важный ориентир. Если VAT не обгоняет обычное зашумление эмбеддингов, значит адверсариальное направление вы посчитали неправильно или epsilon подобран мимо.
Как это собрать на практике
Пайплайн строится вокруг обычной сети над эмбеддингами — в оригинале это LSTM, но с тем же успехом это может быть свёрточный энкодер или трансформер. Возмущение живёт на выходе слоя эмбеддингов.
- Нормализуйте эмбеддинги. Возмущение фиксированной нормы бессмысленно, если у разных слов вектора разного масштаба. В оригинальной работе эмбеддинги нормируют по частотам слов перед добавлением шума.
- Посчитайте чистое предсказание. Прямой проход, сохраняете распределение вероятностей.
- Найдите адверсариальное направление. Один шаг: градиент дивергенции (VAT) или потерь (AT) по эмбеддингам, нормировка, умножение на epsilon.
- Второй проход с возмущением. Считаете loss на испорченном входе.
- Суммируете. Итоговая потеря — supervised loss плюс взвешенный adversarial/VAT loss. Вес и epsilon — главные гиперпараметры.
Пример: шаг VAT на PyTorch
import torch
import torch.nn.functional as F
def vat_loss(model, embeds, xi=1e-6, eps=2.0, n_iter=1):
# embeds: эмбеддинги входа [batch, seq, dim], requires_grad не нужен
with torch.no_grad():
logits_clean = model.from_embeds(embeds)
p_clean = F.softmax(logits_clean, dim=-1)
# случайное стартовое направление
d = torch.randn_like(embeds)
d = F.normalize(d, dim=-1)
for _ in range(n_iter):
d.requires_grad_(True)
logits_hat = model.from_embeds(embeds + xi * d)
p_hat = F.log_softmax(logits_hat, dim=-1)
kl = F.kl_div(p_hat, p_clean, reduction='batchmean')
grad = torch.autograd.grad(kl, d)[0]
d = F.normalize(grad.detach(), dim=-1)
# финальное возмущение
r_adv = eps * d
logits_adv = model.from_embeds(embeds + r_adv)
p_adv = F.log_softmax(logits_adv, dim=-1)
return F.kl_div(p_adv, p_clean, reduction='batchmean')
Здесь from_embeds — часть модели после слоя эмбеддингов. Одной итерации степенного метода (n_iter=1) обычно достаточно: точность аппроксимации направления слабо растёт с числом шагов, а стоимость удваивается.
Что настраивать и на что смотреть
- epsilon — норма возмущения. Слишком маленькое ничего не меняет, слишком большое сбивает обучение. Подбирается по валидации, часто в диапазоне единиц.
- Вес VAT-слагаемого — соотношение supervised и unsupervised частей. При крошечной размеченной выборке его повышают.
- Отношение размеченных к неразмеченным в батче. VAT-loss считается на всех, supervised — только на размеченных. Мешайте их так, чтобы градиенты не заглушали друг друга.
- Стоимость. Каждый шаг — минимум два прямых прохода вместо одного. На больших трансформерах это заметная надбавка к времени обучения.
Практический ориентир из оригинальной работы: VAT давал улучшение на бенчмарках вроде IMDB и других задачах классификации отзывов и новостей относительно чисто supervised-базлайна на том же объёме меток. Точные цифры зависят от датасета и архитектуры — проверяйте на своей задаче, а не переносите чужие проценты вслепую.
1 KL-дивергенция (расхождение Кульбака — Лейблера) — несимметричная мера того, насколько одно распределение вероятностей отличается от другого. Ноль означает полное совпадение; чем больше значение, тем сильнее предсказания расходятся.
Prompt-инженер: Идеальные запросы для Midjourney, ChatGPT и других моделей.
Спросить за 15 ₽Источники: Miyato, Dai, Goodfellow. Adversarial Training Methods for Semi-Supervised Text Classification, Miyato et al. Virtual Adversarial Training: A Regularization Method for Supervised and Semi-Supervised Learning
Частые вопросы
Чем VAT отличается от обычного data augmentation?
Можно ли применять эти методы к трансформерам, а не только к LSTM?
Почему возмущают эмбеддинги, а не сами слова?
Сколько неразмеченных данных нужно, чтобы VAT дал эффект?
VAT сильно замедляет обучение?
Чем проверить, что метод реально работает, а не просто добавляет шум?
Материал носит информационный характер и подготовлен редакцией «Агентуры». Он не является офертой, рекламой или индивидуальной консультацией. Упомянутые продукты, компании и торговые знаки принадлежат их правообладателям. Перед принятием решений, влекущих юридические или финансовые последствия, обратитесь к профильному специалисту.