Введение и обзор
Эта книга предназначена для того, чтобы помочь инженерам машинного обучения (ML) и специалистам по данным успешно пройти собеседования по проектированию ML-систем с акцентом на генеративный ИИ (GenAI). Она дополняет предыдущую книгу «ML System Design Interview» [1], охватывающую фундаментальные темы: системы поиска и рекомендаций. Данная книга исследует приложения GenAI и уникальные задачи проектирования подобных систем. Она также служит руководством для тех, кто хочет понять, как GenAI применяется на практике.
В этой главе рассматриваются две ключевые темы. Во-первых, даётся обзор GenAI — его фундаментальных концепций и приложений. Затем вводится комплексный фреймворк для построения ML-систем, необходимый как для реальных приложений, так и для подготовки к собеседованиям. Этот фреймворк послужит основой для разработки популярных GenAI-систем в последующих главах.
Приступим.
Обзор GenAI
Что такое ИИ и ML?
ИИ — это раздел информатики, направленный на создание систем, способных выполнять задачи, обычно требующие человеческого интеллекта: рассуждение, планирование, решение проблем. ML — это подраздел ИИ, использующий алгоритмы для обучения на данных, а не на основе предопределённых правил. Эти алгоритмы анализируют данные, выявляют паттерны и делают предсказания или генерируют новый контент на основе усвоенных закономерностей. Такие приложения, как рекомендательные системы, обнаружение мошенничества, автономные транспортные средства и чат-боты, как правило, работают на ML-моделях.
ML-модели в целом делятся на две категории:
- Дискриминативные
- Генеративные
Дискриминативные модели
Дискриминативные модели классифицируют данные, обучаясь различиям между классами на основе входных признаков. Формально они изучают условные вероятности P(Y|X), где Y — целевая переменная, а X — входные признаки.
Дискриминативные модели применяются как для классификации — определения класса входных данных, так и для регрессии — предсказания непрерывного значения. Например, в системе обнаружения мошенничества дискриминативная модель классифицирует транзакции как легитимные или мошеннические, анализируя такие признаки, как сумма транзакции и история покупок. Аналогично в рекомендациях фильмов модель предсказывает оценку пользователя на основе его исторических взаимодействий.
Распространённые алгоритмы дискриминативных моделей включают:
- Логистическая регрессия: линейная модель, предсказывающая вероятность бинарного исхода на основе входных признаков.
- Метод опорных векторов (SVM): находит гиперплоскости, наилучшим образом разделяющие классы в пространстве признаков. Может быть расширен для нелинейных границ с помощью ядровых функций [2].
- Деревья решений: эти модели и их вариации, например случайные леса, рекурсивно разбивают данные на подгруппы на основе целевой переменной.
- Метод k ближайших соседей (KNN): непараметрический метод, классифицирующий образец на основе преобладающей метки среди ближайших соседей в пространстве признаков.
- Нейронные сети: состоят из слоёв взаимосвязанных нейронов. Используют взвешенные входы, функции активации и обратное распространение ошибки для обучения и аппроксимации сложных функций для задач классификации и регрессии.
Хотя эти алгоритмы позволяют предсказывать целевую переменную по входным признакам, большинство из них не способны усваивать лежащее в основе распределение данных, необходимое для генерации новых экземпляров. Для этого обращаются к генеративным моделям.
Генеративные модели
Генеративные модели стремятся понять и воспроизвести лежащее в основе распределение данных. Формально они моделируют распределение P(X) при фокусе исключительно на входных данных (например, генерация изображений) или совместное распределение P(X, Y) при учёте и входных данных, и целевой переменной (например, генерация изображений по тексту). Это позволяет генерировать новые экземпляры данных путём выборки из усвоенных распределений.
В отличие от дискриминативных моделей, фокусирующихся на различении экземпляров данных, генеративные модели могут создавать новые экземпляры, тесно напоминающие оригинал. Например, генеративная модель, обученная на изображениях человеческих лиц, способна создавать совершенно новые лица. Такие модели применяются в разнообразных задачах: генерация текста, генерация изображений, синтез речи.
Генеративные алгоритмы делятся на два класса: классические и современные. Классические хорошо обучаются паттернам структурированных данных, но могут испытывать трудности с более сложными или неструктурированными данными. К распространённым классическим генеративным алгоритмам относятся:
- Наивный байесовский классификатор: вероятностная модель, основанная на теореме Байеса [3].
- Модели гауссовых смесей (GMM): представляют данные как смесь гауссовых распределений [4].
- Скрытые марковские модели (HMM): моделируют совместную вероятность наблюдаемых последовательностей и скрытых состояний, генерирующих эти последовательности [5].
- Машины Больцмана: энергетические модели, используемые для обучения признакам или снижения размерности [6].
Современные генеративные алгоритмы, в свою очередь, обучаются на сложных распределениях данных и хорошо подходят для задач генерации реалистичных изображений и точных текстовых ответов на запросы. К распространённым современным генеративным алгоритмам относятся:
- Вариационные автоэнкодеры (VAE): тип автоэнкодера, моделирующий распределение данных путём кодирования в скрытое пространство (latent space) и последующей реконструкции исходных данных с помощью decoder.
- Генеративно-состязательные сети (GAN): класс нейронных сетей, в котором генератор и дискриминатор обучаются одновременно. Генератор создаёт реалистичные данные, а дискриминатор пытается отличить реальные данные от сгенерированных.
- Диффузионные модели: модели, изучающие сложные распределения данных через обратный процесс диффузии. Широко используются для генерации изображений и видео.
- Авторегрессионные модели: генерируют данные, предсказывая каждый элемент последовательности на основе предшествующих элементов. Широко применяются при генерации текста и прогнозировании временных рядов.
Дискриминативные и генеративные модели используются для разных целей. Дискриминативные обычно применяются для классификации или предсказания; генеративные — для создания новых образцов. На рисунке 2 показаны популярные задачи, решаемые с помощью генеративных и дискриминативных моделей.
Что такое GenAI и почему он набирает популярность?
GenAI предполагает использование современных генеративных алгоритмов для обучения моделей, способных создавать новые образцы данных: изображения, видео, текст и аудио.
GenAI стал очень популярным по двум основным причинам. Во-первых, эти модели могут выполнять разнообразные задачи в разных областях: генерировать текст, создавать реалистичные изображения, сочинять музыку. Эта многозадачность делает их ценными в различных отраслях — от творческих искусств и развлечений до здравоохранения и разработки программного обеспечения.
Во-вторых, GenAI-приложения существенно повышают производительность. Например, в создании контента эти модели могут генерировать черновики, предлагать улучшения или даже создавать финальные материалы, экономя значительное время и ресурсы. Другой пример — использование больших языковых моделей (LLM), таких как ChatGPT [7], для ответов на сложные вопросы и ведения осмысленных диалогов. По данным недавнего отчёта McKinsey [8], ожидается, что GenAI обеспечит «рост производительности труда на 0,1–0,6% ежегодно вплоть до 2040 года».
Почему GenAI становится таким мощным?
Модели GenAI в последнее время демонстрируют впечатляющие возможности и становятся всё более мощными. Три ключевых фактора этого прогресса:
- Данные
- Ёмкость модели
- Вычислительные ресурсы
Данные
Эффективность ML-модели зависит от обучающих данных. Например, если модель не обучалась на обширных медицинских данных, она может испытывать трудности с точной диагностикой заболеваний. Улучшение модели для конкретной задачи требует больших размеченных наборов данных, но их сбор может быть сложным и дорогостоящим.
Одним из ключевых факторов успеха GenAI является самообучение без учителя. В отличие от классических моделей, хорошо работающих на размеченных данных, GenAI-модели могут обучаться на неразмеченных данных. Этот подход позволяет использовать огромные наборы данных из интернета без необходимости дорогостоящей и трудоёмкой разметки.
Благодаря лёгкому доступу к очень большим наборам данных из интернета современные GenAI-модели могут обучаться на массивных наборах данных, иногда превышающих миллиарды текстовых документов или изображений. Например, модель Llama 3 от Meta [9] обучалась на 15 триллионах токенов — порядка 50 терабайт данных; модель Flamingo от Google [10] обучалась на 1,8 миллиарда пар (изображение, текст). Обучение на этом огромном объёме данных помогает моделям усваивать сложные паттерны и нюансы, что обеспечивает высокое качество результатов.
Ёмкость модели
Ещё одним ключевым фактором эффективности ML-моделей является их способность к обучению. Ёмкость модели измеряется двумя способами:
- Количество параметров
- Число FLOP
Количество параметров
Параметры — это значения внутри модели, усваиваемые в процессе обучения. Количество параметров является ключевым показателем способности модели обучаться на данных.
Модель с большим числом параметров обычно имеет большую ёмкость для усвоения сложных паттернов и зависимостей в данных. При условии обучения на большом наборе данных это часто означает лучшую производительность. В таблице 1 показаны пять популярных моделей и их количество параметров.
| Название модели | Параметры |
|---|---|
| Google's PaLM [11] | 540B |
| OpenAI's GPT-3 [12] | 175B |
| Google's Flamingo [10] | 80B |
| Meta's Llama 3 [9] | 405B |
| Google's Imagen [13] | 2B |
Таблица 1: Популярные GenAI-модели и их количество параметров
Число FLOP
FLOP (Floating Point Operations — операции с плавающей точкой) измеряет вычислительную сложность модели путём подсчёта операций с плавающей точкой, необходимых для одного прямого прохода. Включает базовые арифметические операции — сложение, умножение и другие, выполняемые при прохождении данных через слои модели.
Для лучшего понимания числа FLOP рассмотрим простой пример.
Рассмотрим полносвязный слой с 4 входными нейронами и 3 выходными. Каждый выходной нейрон вычисляется путём умножения входных нейронов на соответствующие веса и суммирования результатов. Это даёт 4 умножения и 3 сложения для каждого выходного нейрона, как показано на рисунке 4. Следовательно, общее число FLOP равно 3×(4+3)=21.
Если количество параметров измеряет размер модели, то FLOP указывает на число арифметических вычислений и даёт представление о вычислительной сложности модели. Хотя модель с большим числом параметров часто имеет более высокое число FLOP, это не всегда так. Архитектура играет ключевую роль — плотные слои обычно требуют больше FLOP, чем разреженные связи, даже при одинаковом числе параметров. Понимание этого различия критично при оптимизации модели, так как требует баланса между параметрами и FLOP.
Вычислительные ресурсы
По мере роста ёмкости модели её производительность, как правило, улучшается, но обучение таких крупных моделей требует огромных вычислительных ресурсов. Необходимые для обучения вычисления часто измеряются в FLOP, представляющих суммарное число выполненных операций. Например, модель PaLM-2 от Google обучалась с использованием 10²² FLOP [14].
Вычислительная мощность обычно обеспечивается таким аппаратным обеспечением, как CPU, GPU (Graphics Processing Units) и TPU (Tensor Processing Units). Nvidia, например, предлагает передовые GPU: H100, A100 и A10, каждый с разной стоимостью и вычислительными возможностями. Производительность этих машин часто измеряется в FLOP/S (операций с плавающей точкой в секунду). Например, Nvidia H100 может обеспечить до 60 терафлопс в секунду (60 TFLOP/S) [15].
Обучение передовых GenAI-моделей обходится дорого: требуются тысячи GPU в течение нескольких недель. Для понимания вычислительных требований к моделям наподобие PaLM-2 рассчитаем необходимое количество GPU H100. Предполагая пиковую производительность H100 в 60 TFLOP/S, он может выполнить примерно 5,18 эксафлоп в день. Учитывая, что PaLM-2 потребовала 10²² FLOP, одному GPU H100 понадобилось бы около 5,5 лет для завершения необходимых вычислений. Поэтому стоимость обучения крупных моделей чрезвычайно высока и нередко превышает десятки миллионов долларов. Например, Сэм Альтман, генеральный директор OpenAI, заявил, что стоимость обучения GPT-4 превысила 100 миллионов долларов [16].
Обучение модели с миллиардами параметров (например, GPT-4, Llama 3) было невозможно ещё несколько лет назад. Сдвиг в возможностях произошёл главным образом благодаря развитию аппаратного обеспечения — специализированным чипам, таким как GPU и TPU, разработанным для задач глубокого обучения. Распределённое обучение также сыграло ключевую роль, позволяя распределять нагрузку между сотнями и даже тысячами машин параллельно. Это значительно ускоряет процесс, делая осуществимым обучение очень крупных моделей на огромных наборах данных. Эти улучшения инфраструктуры и методов обучения позволили обучать GenAI-модели в беспрецедентных масштабах.
Закон масштабирования
В рамках бюджета вычислений (измеряемого в FLOP) какое оптимальное сочетание размера модели и объёма обучающих данных (измеряемого числом токенов) даёт наименьшие потери? Это фундаментальный вопрос, на который исследователи стремятся ответить с помощью законов масштабирования.
В 2020 году исследователи OpenAI провели обширные эксперименты по обучению LLM, изучая различные факторы: размеры моделей (N), размеры наборов данных (D), вычислительные ресурсы (C), архитектуры моделей и длины контекста [17]. Их выводы выявили два ключевых наблюдения. Во-первых, влияние масштабирования на производительность модели значительно более выражено, чем влияние архитектурных вариаций. Во-вторых, по мере увеличения размера модели, объёма набора данных или вычислительных ресурсов наблюдается соответствующее и предсказуемое улучшение производительности, следующее степенному закону.
В 2022 году исследователи DeepMind расширили это понимание, показав, что многие существующие LLM были недообучены — модели были недостаточно велики для объёма данных, на которых они обучались [18]. Они обнаружили, что объём данных должен масштабироваться линейно с размером модели для достижения оптимальной производительности.
С недавним выходом GPT o1 [19] исследователи начали строить предположения о существовании закона масштабирования и в процессе инференса [20].
Риски и ограничения GenAI
GenAI быстро развивается, стимулируя прогресс во многих отраслях за счёт создания реалистичных текстов, изображений и видео. Однако он также несёт серьёзные риски и ограничения, требующие тщательного рассмотрения. Решение этих проблем является ключом к обеспечению ответственного и устойчивого развития. К распространённым проблемам относятся:
- Этические проблемы: вопросы, связанные с предвзятостью, интеллектуальной собственностью, дезинформацией и злоупотреблением сгенерированным контентом, могут иметь вредные социальные последствия.
- Воздействие на окружающую среду: высокие вычислительные мощности, необходимые для обучения крупных моделей, ведут к значительному энергопотреблению и выбросам углерода.
- Ограничения моделей: GenAI-моделям часто не хватает истинного понимания, что приводит к неточностям и ограничениям в сложных задачах рассуждения, а также к галлюцинациям (hallucination).
- Риски безопасности: угрозы, создаваемые использованием GenAI для создания дипфейков в целях шантажа, политических манипуляций, автоматизированных фишинговых атак и состязательных воздействий, способных манипулировать выходными данными моделей в критических системах, таких как здравоохранение и финансы.
Каждая из этих областей представляет серьёзную проблему в разработке GenAI-приложений. Для решения этих рисков требуется мультидисциплинарный подход, включающий не только технические решения, но и этические рамки, правовые нормы и осведомлённость общества.
Фреймворк для собеседований по проектированию ML-систем
Многие инженеры воспринимают ML-алгоритмы — например, авторегрессионные Transformer или диффузионные модели — как всю ML-систему целиком. Однако построение и развёртывание GenAI-систем предполагает гораздо большее, чем просто обучение модели. Эти системы сложны и включают такие компоненты, как конвейеры данных для обработки и предобработки крупных наборов данных, механизмы оценки качества и безопасности результатов, инфраструктуру для масштабной доставки контента, сгенерированного ИИ, и мониторинг для обеспечения стабильной производительности со временем.
На собеседовании по проектированию ML-систем, особенно ориентированном на GenAI, вам нередко придётся сталкиваться с открытыми вопросами. Например, вас могут попросить разработать чат-бот для обслуживания клиентов или инструмент редактирования изображений на базе ИИ для творческих людей. Единственно верного ответа не существует. Интервьюер интересуется вашим подходом к сложным проблемам, пониманием концепций GenAI, процессом проектирования системы и обоснованием принятых решений.
Для успеха на собеседовании по проектированию GenAI-системы крайне важно следовать структурированному фреймворку. Неорганизованные ответы могут затуманить процесс мышления и снизить ясность. Эта книга представляет фреймворк для решения задач проектирования GenAI-систем, включающий следующие ключевые этапы:
- Уточнение требований
- Формулировка задачи как ML-задачи
- Подготовка данных
- Разработка модели
- Оценка
- Общий дизайн ML-системы
- Развёртывание и мониторинг
Рассмотрим каждый этап подробнее, изучая ключевые соображения и темы для обсуждения при проектировании GenAI-системы.
Уточнение требований
Приступая к разработке ML-системы для решения конкретной задачи, вы часто располагаете минимальной информацией. Аналогично на собеседованиях вопросы по проектированию ML-систем зачастую расплывчаты и содержат мало деталей. Например, вас могут попросить «разработать систему генерации изображений». Первый шаг — задать уточняющие вопросы. Но какие именно?
Ваши вопросы должны помочь понять область проблемы и конкретные цели, которых система должна достичь. Они делятся на два типа:
- Функциональные требования
- Нефункциональные требования
Функциональные требования
Функциональные требования описывают, что система должна делать — её основные возможности. Например, «генерировать изображение и настраивать его стиль на основе запроса пользователя» является функциональным требованием для системы text-to-image. В контексте проектирования GenAI-системы функциональные требования имеют решающее значение, так как формируют высокоуровневую архитектуру системы. Они направляют разработку основных компонентов и функциональностей, необходимых системе для удовлетворения потребностей пользователей.
Нефункциональные требования
Нефункциональные требования фокусируются на том, как система работает, а не на том, что она делает. К ним относятся метрики производительности — задержка и пропускная способность, — а также соображения справедливости, безопасности и масштабируемости. В системе генерации изображений, например, нефункциональные требования могут определять допустимую скорость генерации изображений и стандарты их качества. Хотя эти требования обычно являются продвинутыми темами в проектировании GenAI-систем и могут не оказывать существенного влияния на начальную архитектуру, крайне важно выявить и понять их на раннем этапе, поскольку они будут формировать более поздние стадии проектирования — особенно при настройке производительности и улучшении системы.
Вот несколько вопросов для начала:
- Бизнес-цель: какова основная цель системы? Какому конкретному назначению она будет служить? Например, при проектировании системы создания подписей к изображениям необходимо знать, будет ли она использоваться для генерации подробных описаний товаров на платформе электронной торговли или для предложения коротких подписей к фотографиям в социальных сетях.
- Функции системы: какие функции должна поддерживать система, которые могут повлиять на ML-дизайн? Например, при проектировании системы генерации изображений важно знать, смогут ли пользователи оставлять отзывы или оценивать сгенерированные изображения, поскольку эти взаимодействия могут улучшить модель. Аналогично при проектировании LLM важно знать, какие языки должны поддерживаться.
- Данные: каковы источники данных? Насколько велик набор данных? Размечены ли данные? Эти вопросы критичны, поскольку качество и количество данных могут влиять на дизайн.
- Ограничения: каковы доступные вычислительные ресурсы? Система будет облачной или разработанной для работы на локальных устройствах?
- Масштаб системы: сколько пользователей ожидается? Сколько изображений нужно генерировать и какой ожидается рост спроса? Эти вопросы важно прояснить, потому что система, разработанная для генерации изображений для небольшой группы пользователей, не потребует такого же уровня масштабируемости, как система, рассчитанная на обслуживание миллионов пользователей.
- Производительность: насколько быстро должен генерироваться контент? Требуется ли генерация в реальном времени? Что имеет больший приоритет: качество контента или скорость генерации?
Этот список не исчерпывающий, но даёт хорошую отправную точку. Другие темы, такие как конфиденциальность, этика и безопасность данных, также могут быть важны.
К концу этого этапа вы должны прийти к единому пониманию с интервьюером относительно масштаба и требований системы. Как правило, разумно прояснить эти детали, чтобы убедиться, что вы отвечаете ожиданиям интервьюера.
Формулировка задачи как ML-задачи
Если интервьюер просит вас разработать функцию автоматического обобщения электронных писем, перед вами стоит задача. Но вы не можете просто попросить ИИ обобщать письма. Вместо этого вам нужно сформулировать задачу так, чтобы её могли решить ML-методы. Формулировка задачи как ML-задачи является ключевым этапом в проектировании ML-систем, поскольку определяет весь последующий дизайн.
При решении задачи сначала нужно определить, необходимо ли вообще машинное обучение. Однако для GenAI-систем можно обычно предположить, что ML потребуется, поскольку это основной инструмент разработки таких систем.
Следующие два шага полезны для формулировки задачи как ML-задачи:
- Определить входные и выходные данные системы
- Выбрать подходящий ML-подход
Определение входных и выходных данных системы
Для формулировки задачи сначала определите входные и выходные данные системы. Это включает идентификацию модальности входных данных (текст, изображение, аудио, видео) и ожидаемого выхода. Например, в системе чат-бота вход — это текстовый запрос пользователя, а выход — ответ системы.
Выбор подходящего ML-подхода
После определения входных и выходных данных системы следующий шаг — выбор наиболее подходящего ML-подхода. Это включает определение ключевых компонентов системы и выбор алгоритма, соответствующего конкретным потребностям задачи. Как показано на рисунке 8, существует множество ML-алгоритмов, каждый со своими сильными и слабыми сторонами. Крайне важно сравнивать их, обсуждать компромиссы и выбирать наиболее подходящий для вашей задачи.
Существуют разные способы выбора подходящего ML-алгоритма, и критерии выбора варьируются от приложения к приложению. Следующие шаги помогут сузить варианты при выборе наиболее подходящего алгоритма:
- Дискриминативный vs. генеративный: сначала определите, требует ли задача дискриминативной или генеративной модели. Это легко установить по выходным данным системы. Например, в задаче обнаружения объектов, где выход — класс входного изображения, задача является дискриминативной. Напротив, проектирование чат-бота, производящего текст на выходе, — это генеративная задача. Рисунок 9: Шаг первый при выборе подходящего ML-подхода
- Определите тип задачи: далее определите конкретный тип задачи для дальнейшего сужения выбора алгоритмов. Двумя наиболее распространёнными задачами дискриминативных моделей являются классификация и регрессия. Генеративные модели обычно решают задачи генерации текста, изображений, аудио и видео. Выход системы может помочь определить тип задачи. Например, система создания подписей к изображениям генерирует текст — это задача генерации текста; система генерации лиц создаёт изображения — это задача генерации изображений; система обнаружения объектов, выдающая класс объекта, является задачей классификации. Рисунок 10: Шаг второй при выборе подходящего ML-подхода
- Выберите подходящий алгоритм: наконец, выберите алгоритм, наиболее подходящий исходя из требований. Рассмотрите такие факторы, как способность обрабатывать различные модальности входных данных, эффективность и ожидания по качеству. Например, в системе text-to-image алгоритм должен обрабатывать текст на входе и генерировать изображение на выходе; поэтому VAE или GAN могут быть не идеальными, несмотря на их способность генерировать изображения. Этот шаг — идеальное время для оценки различных вариантов и обсуждения их компромиссов.
В последующих главах мы рассмотрим различные ML-подходы, используемые в популярных GenAI-приложениях.
Темы для обсуждения
- Каковы входные и выходные данные системы исходя из требований?
- Какие модальности данных (текст, изображение, аудио, видео) модель должна понимать и обрабатывать? Как модель будет работать с разными модальностями?
- Должна ли одна модель обрабатывать все модальности входных данных, или эффективнее использовать несколько моделей для разных модальностей? Каковы преимущества и недостатки единой модели по сравнению со специализированными моделями для каждой модальности?
- Какой генеративный алгоритм (например, диффузионные модели, VAE, GAN) наиболее подходит для данной задачи и почему? Каковы конкретные компромиссы между разными алгоритмами с точки зрения качества, эффективности и простоты использования?
- Каковы последствия выбора одного алгоритма вместо другого с точки зрения производительности, стабильности и ресурсов?
- Достаточно ли выбранный подход масштабируем и гибок для адаптации к будущим изменениям или дополнениям возможностей системы? Насколько легко система может адаптироваться при введении новых модальностей входных данных или выходов?
Подготовка данных
ML-модели обучаются непосредственно на данных, поэтому высококачественные данные критичны для эффективного обучения. В этом разделе рассматриваются различные типы данных и ключевые соображения при их подготовке для ML-моделей.
Типы данных
В ML данные обычно делятся на два типа: структурированные и неструктурированные.
Структурированные данные: этот тип данных можно организовать в таблицы со строками и столбцами, например базу данных или электронную таблицу. Финансовые записи и данные о клиентах — примеры структурированных данных. Структурированные данные можно дополнительно разделить на следующие категории:
- Категориальные данные: данные, представляющие отдельные группы или категории (например, пол или цвет).
- Числовые данные: данные, представляющие измеримые величины (например, количество проданных товаров, цена дома).
- Порядковые данные: данные с предопределённым порядком (например, рейтинги удовлетворённости).
Неструктурированные данные: это данные без базовой схемы или структуры, такие как текст, изображения, видео, аудиофайлы или их комбинация. Например, публикации в социальных сетях или электронные письма являются примерами неструктурированных данных.
Традиционные ML-модели обычно обучаются на структурированных данных. Напротив, модели, лежащие в основе GenAI-приложений, в основном работают с неструктурированными данными: изображениями, текстом и видео. В результате акцент в подготовке данных существенно различается между традиционными моделями, работающими со структурированными данными, и генеративными моделями, работающими с неструктурированными данными.
Подготовка данных в традиционном ML
Подготовка структурированных данных для традиционных моделей обычно включает два ключевых этапа: инженерию данных и инженерию признаков.
Инженерия данных
Инженерия данных включает построение и поддержку систем для сбора, хранения, извлечения и обработки данных. Основным компонентом является ETL (Extract, Transform, Load) [21] — процесс извлечения данных из различных источников, преобразования их в пригодный формат и загрузки в хранилище данных или другую систему хранения. Инженерия данных обеспечивает чистоту, надёжность и доступность данных.
Инженерия признаков
Инженерия признаков включает выбор и извлечение предиктивных признаков из сырых данных и их преобразование в формат, пригодный для ML-моделей. Этот процесс часто использует хранилища признаков, такие как Tecton [22] или Amazon SageMaker [23], которые предоставляют централизованную платформу для управления и обслуживания признаков в масштабе.
Выбор правильных признаков критически важен при разработке и обучении ML-моделей. Важно выбирать признаки, несущие наибольшую информацию. Процесс инженерии признаков требует экспертизы в предметной области и является узкоспециализированным. Он включает такие методы, как обработка пропущенных значений, представление категориальных признаков и разбиение на бины.
Поскольку эта книга посвящена GenAI, мы сосредоточимся в первую очередь на подготовке данных, специфичной для GenAI-моделей. Для более глубокого изучения инженерии данных и признаков обратитесь к [24].
Подготовка данных в GenAI
При подготовке неструктурированных данных для генеративных моделей акцент смещается с инженерии признаков на сбор огромных объёмов данных; обеспечение высокого качества и безопасности данных; использование инструментов для эффективного хранения и извлечения данных в масштабе.
Рассмотрим следующие ключевые этапы подготовки данных:
- Сбор данных
- Очистка данных
- Эффективность работы с данными
Сбор данных
Продвинутые GenAI-модели имеют миллиарды параметров, позволяющих им обучаться на данных и обобщать полученные знания. Из-за своего размера эти модели требуют огромного объёма обучающих данных для выявления сложных паттернов. Например, Llama 3 обучалась на 15 триллионах токенов из различных интернет-источников — эквивалент 50 терабайт данных. Для понимания масштаба: человеку, читающему непрерывно со стандартной скоростью 250 слов в минуту, потребовалось бы около 85 000 лет, чтобы прочитать такой объём текста. Процесс сбора данных включает сбор крупных наборов данных путём сканирования текста из различных источников (веб-сайты, социальные сети, форумы).
По мере роста моделей наблюдается тенденция к дополнению обучающих наборов данных контентом, сгенерированным ИИ. Это предполагает использование существующих моделей для создания синтетических данных, которые затем используются для обучения другой GenAI-модели.
Использование контента, сгенерированного ИИ, для обучения GenAI-моделей имеет ряд преимуществ и недостатков.
Преимущества:
- Повышение разнообразия данных: контент, сгенерированный ИИ, добавляет разнообразие к существующим данным, тем самым улучшая способность модели к обобщению, особенно когда исходные данные ограничены или несбалансированы.
- Масштабируемость: по мере роста потребности в данных, контент, сгенерированный ИИ, предоставляет масштабируемый способ создания крупных наборов данных, которые сложно собрать вручную.
Недостатки:
- Проблемы качества: качество синтетических данных зависит от исходной модели. Некачественные данные могут приводить к распространению предвзятостей или ошибок.
- Проблемы репрезентативности: синтетические данные могут плохо представлять исходные данные. Обеспечить разнообразие и репрезентативность синтетических данных может быть непросто.
- Расхождение с реальными распределениями: данные, сгенерированные ИИ, могут не в полной мере отражать сложность реальных сценариев, рискуя упустить важные детали.
Использование данных, сгенерированных ИИ, для обучения GenAI-моделей — быстро развивающаяся область исследований. Постоянно разрабатываются новые методы для повышения качества и релевантности синтетических данных. Для получения дополнительной информации обратитесь к [25].
Очистка данных
Очень крупные наборы данных из интернета зачастую зашумлены и могут содержать низкокачественный или неприемлемый контент. Данные необходимо тщательно очищать, чтобы избежать внесения предвзятостей, дезинформации или вредоносных материалов в модель, что может негативно сказаться на её производительности. Кроме того, крайне важно иметь репрезентативные данные — это требует удаления дубликатов и обеспечения разнообразия и сбалансированности.
На протяжении всей книги мы будем изучать ключевые методы очистки данных: фильтрацию вредоносного контента, обнаружение NSFW (Not Safe For Work), присвоение оценок качества и удаление дубликатов.
Эффективность работы с данными
Управление крупными наборами данных требует эффективных инструментов и методов хранения и извлечения. Рассмотрим каждый подробнее.
Эффективное хранение
Хранение огромных объёмов данных с традиционными инструментами может быть дорогостоящим и медленным. Распределённые системы хранения, такие как Hadoop Distributed File System (HDFS) [26] и Amazon S3 [27], созданы для хранения огромных объёмов данных на нескольких машинах. Эти системы особенно подходят для управления большими объёмами неструктурированных данных. Кроме того, колоночные форматы хранения, такие как Parquet [28] и ORC [29], идеальны для структурированных данных или неструктурированных данных, преобразованных в структурированную форму. Оптимизированные для аналитики, эти форматы обеспечивают лучшее сжатие и более высокую скорость обработки запросов.
Эффективное извлечение
Обучение ML-модели требует быстрого извлечения данных. К распространённым методам эффективного извлечения данных из крупных наборов относятся:
- Шардирование: разбиение данных на несколько устройств обеспечивает параллельный доступ и ускоряет извлечение и обработку.
- Индексирование: технологии, такие как Apache Lucene [30] или Elasticsearch [31], используются для индексирования данных, что упрощает и ускоряет поиск конкретных фрагментов информации.
- Предварительная загрузка или кэширование: часто используемые данные предварительно загружаются в память для уменьшения задержек ввода-вывода при извлечении.
Темы для обсуждения
- Источники данных: какие данные доступны и откуда их собирать? Насколько они разнообразны? Каков размер набора данных?
- Чувствительность данных: насколько чувствительны данные (например, персональные, финансовые, медицинские)? Необходима ли анонимизация для защиты конфиденциальной информации?
- Предвзятость: есть ли в данных изначально предвзятости (например, демографические, географические)? Как их обнаруживать и устранять для обеспечения справедливого представления?
- Качество данных: как фильтровать низкокачественные, нерелевантные или зашумлённые данные? Есть ли в наборе данных выбросы или аномалии? Как с ними работать?
- Неприемлемые данные: есть ли в наборе данных неприемлемый, вредоносный или NSFW-контент? Какие процессы предусмотрены для обнаружения и удаления таких данных?
- Предобработка данных: как данные представляются в формате, понятном модели? При использовании текстовых данных — как они токенизируются и преобразуются в числовой формат (например, embedding)? При работе с мультимодальными данными (изображения, текст, аудио) — как они предобрабатываются для подачи в модель?
Разработка модели
Разработка модели — критически важный этап построения ML-систем. Он включает выбор подходящей архитектуры, обучение модели и, наконец, генерацию новых данных с помощью обученной модели. Рассмотрим каждый из этих компонентов подробнее.
Архитектура модели
На этом этапе следует подробно обсудить архитектуру модели. Для разных ML-алгоритмов может существовать несколько жизнеспособных архитектур. Например, диффузионные модели можно строить как на основе архитектуры U-Net, так и DiT. На собеседовании важно изучить различные архитектурные варианты и взвесить их преимущества и недостатки.
После выбора архитектуры полезно проанализировать её конкретные слои и изучить, как входные данные преобразуются в выходные. Например, в архитектуре U-Net выходные данные должны иметь тот же размер, что и входное изображение; поэтому проверка слоёв на соответствие этому требованию полезна.
Вам могут задать уточняющие вопросы и попросить модифицировать архитектуру для поддержки новой функциональности. Например, в модели генерации изображений может понадобиться изменить архитектуру, чтобы пользователи могли управлять стилем генерируемых изображений. Аналогично в модели text-to-video вас могут попросить управлять направлением движения (например, слева направо) в процессе генерации. Для реализации этих функций может потребоваться добавление или изменение компонентов архитектуры для интеграции векторов стиля или информации о движении.
Рассмотрим реальный пример — механизм самовнимания Transformer — чтобы показать, что означает обсуждение архитектуры на собеседовании.
Архитектура самовнимания Transformer
Transformer являются краеугольным камнем современного GenAI, особенно в обработке естественного языка и генерации изображений. С момента своего появления в 2017 году [32] они стремительно завоевали сообщество ИИ, став доминирующей архитектурой для широкого спектра задач: обработка естественного языка (NLP) [33][12], компьютерное зрение [34] и мультимодальное обучение [35][36].
В основе Transformer лежит механизм внимания. Изначально введённый в контексте машинного перевода [37], он стал фундаментальным компонентом различных архитектур нейронных сетей, особенно модели Transformer. Он устраняет ограничения традиционных последовательных моделей, таких как RNN и LSTM, позволяя модели более эффективно улавливать долгосрочные зависимости и контекстную информацию.
Самовнимание, также известное как масштабированное dot-product внимание — это наиболее распространённая форма механизма внимания в современных моделях. Оно позволяет каждому элементу входной последовательности фокусироваться на каждом другом элементе. Это достигается путём преобразования входных embedding для каждого токена в три вектора: запрос (Q), ключ (K) и значение (V). Эти векторы вычисляются с использованием обучаемых матриц весов WQ, WK и WV: где X представляет входную последовательность embedding.
Оценки внимания вычисляются путём взятия скалярного произведения вектора запроса Q со всеми векторами ключей K, с последующей операцией масштабирования и функцией softmax. Это можно представить как:
Здесь dK — размерность векторов ключей, а масштабирующий коэффициент 1/√dK используется для предотвращения слишком больших значений скалярного произведения, что могло бы привести к чрезвычайно малым градиентам при обратном распространении ошибки. Функция softmax применяется для нормализации оценок внимания, обеспечивая их суммирование в 1. В результате получается взвешенная сумма векторов значений V, где веса определяются релевантностью каждого входного токена в соответствии с оценками внимания.
Многоголовое внимание
Для улавливания различных типов зависимостей и контекстных связей механизм самовнимания часто расширяется до многоголового внимания. Вместо вычисления единственного набора векторов Q, K и V входные данные проецируются во множество наборов, или «голов» (heads), каждый со своими обучаемыми матрицами весов: где каждая голова внимания вычисляется независимо как:
Результаты разных голов конкатенируются и затем линейно преобразуются с помощью выходной матрицы весов WO. Это позволяет модели одновременно фокусироваться на информации из разных подпространств представлений и улавливать более богатые зависимости.
Не существует универсальной архитектуры для каждой задачи. Интервьюеры хотят оценить ваше понимание различных ML-архитектур, их сильных и слабых сторон, а также вашу способность выбирать подходящую на основе конкретных требований и ограничений. В этой книге мы рассмотрим различные модели на основе Transformer через призму GenAI-приложений и объясним необходимость их применения.
Обучение модели
Обучение модели — это процесс корректировки параметров модели (весов) для получения желаемых выходных данных. Ключевые аспекты для обсуждения в процессе обучения модели включают:
- Методология обучения
- Обучающие данные
- ML-цель и функция потерь
- Специфические задачи и их решения
Рассмотрим каждый подробнее.
Методология обучения
Каждая модель следует особому процессу обучения, соответствующему её архитектуре и назначению. Например, диффузионные модели постепенно удаляют шум из данных для генерации высококачественных образцов из шума. Напротив, GAN опираются на состязательное обучение, в котором генератор и дискриминатор конкурируют, совершенствуясь со временем.
Многие модели также проходят многоэтапное обучение для улучшения производительности. LLM, например, обычно проходят три этапа: предобучение на крупных наборах данных для усвоения общих паттернов; контролируемое fine-tuning для адаптации к конкретной задаче; этап согласования для обеспечения соответствия выходных данных человеческим ценностям или предполагаемым поведениям. Такой подход помогает моделям хорошо работать в различных приложениях.
Глубокое понимание методологий обучения различных GenAI-приложений крайне важно, особенно при их обсуждении на собеседовании.
Обучающие данные
Понимание обучающих данных необходимо для успешной разработки модели. Используемые данные могут различаться в зависимости от GenAI-приложения и также могут отличаться при многоэтапных подходах к обучению. Например, при обучении LLM публично доступные наборы данных, такие как Common Crawl [38], могут использоваться на этапе предобучения, тогда как аннотированные экспертами, тщательно отобранные данные — на этапе согласования.
Важно обсудить наборы данных: как они получены, почему ценны для обучения модели, и оценку их размера для эффективного обучения.
ML-цель и функция потерь
ML-цель — это цель ML-задачи в процессе обучения. Например, для LLM ею может быть точное предсказание следующего токена (предсказание следующего токена). Напротив, для VAE ML-цель — реконструкция исходного изображения.
Функция потерь измеряет, насколько точно предсказания модели соответствуют желаемым результатам. Она направляет процесс оптимизации, цель которого — минимизировать эти потери. Выбор правильной функции потерь критичен в процессе обучения, поскольку она количественно оценивает ошибки предсказания и помогает алгоритму оптимизации корректировать параметры модели для улучшения производительности.
Проектирование новой функции потерь может быть сложным. В большинстве случаев вы будете выбирать из существующих вариантов, основываясь на том, как вы сформулировали задачу. Иногда требуются незначительные корректировки функции потерь для адаптации её к конкретной задаче. Мы рассмотрим несколько функций потерь в последующих главах.
Специфические задачи и их решения
Разные задачи имеют свои специфические трудности, требующие особых решений. Например, обучение крупных моделей генерации видео очень ресурсоёмко, поскольку требует большой вычислительной мощности и огромных объёмов данных. Это означает, что обучение модели генерации видео может быть невозможным без надлежащих методов оптимизации — таких как методы распараллеливания [39][40][41], обучение со смешанной точностью [42] и латентные диффузионные модели [43]. Эти подходы помогают масштабировать модели генерации видео, сохраняя использование ресурсов и затраты на приемлемом уровне. Хотя специфические задачи и решения будут рассматриваться в будущих главах, кратко изучим методы повышения эффективности и оптимизации, применимые ко всему крупномасштабному обучению моделей.
Три наиболее распространённых метода обучения крупномасштабных моделей:
- Gradient checkpointing
- Обучение со смешанной точностью
- Распределённое обучение
Gradient checkpointing
Gradient checkpointing [44] — метод сокращения использования памяти при обучении модели путём сохранения только выбранного подмножества активаций. При обратном проходе недостающие активации пересчитываются. Это значительно снижает использование памяти, что особенно полезно при обучении крупных моделей с ограниченной памятью GPU.
Обучение со смешанной точностью
Обучение со смешанной точностью — метод, использующий как 16-битные (полуточные), так и 32-битные (одинарной точности) числа с плавающей точкой для ускорения обучения модели и сокращения использования памяти. Оно сохраняет точность обучения, повышая эффективность за счёт выполнения большинства вычислений с меньшей точностью; критические операции выполняются с более высокой точностью по мере необходимости.
Автоматическое смешанное обучение (AMP) [45] — конкретная реализация обучения со смешанной точностью, предоставляемая такими фреймворками, как PyTorch и TensorFlow. AMP автоматически управляет переходом между полуточным и одинарным представлением, оптимизируя использование каждого типа точности и применяя методы масштабирования для поддержания числовой стабильности при обучении.
Распределённое обучение
По мере роста размера и сложности моделей обучение на одной машине становится нецелесообразным. Методы распределённого обучения позволяют эффективно обучать крупные модели, задействуя несколько машин или устройств параллельно.
Распространённые методы параллелизма:
- Параллелизм данных
- Параллелизм модели
- Гибридный параллелизм
Параллелизм данных
При параллелизме данных набор данных распределяется по нескольким устройствам (например, GPU), каждое из которых хранит полную копию модели и обрабатывает часть данных параллельно. Каждое устройство обучается на своём подмножестве данных, а сервер параметров координирует обновление и распределение параметров модели по всем устройствам. Этот подход полезен при работе с крупными наборами данных, поскольку параллельная обработка данных эффективна и ускоряет обучение.
Существуют два основных метода обновления параметров модели на всех устройствах:
- Синхронный: при этом подходе все устройства завершают вычисления и отправляют градиенты на сервер параметров. Сервер ждёт получения градиентов от всех устройств, затем агрегирует и обновляет модель перед отправкой обновлённых параметров обратно всем устройствам. Это обеспечивает согласованность, поскольку устройства всегда работают с одной и той же версией модели. Однако может быть медленнее, так как обновление ждёт самое медленное устройство.
- Асинхронный: при асинхронном обновлении каждое устройство отправляет градиенты на сервер параметров сразу по завершении обработки своей части данных, а сервер обновляет модель немедленно при получении градиентов от любого устройства и отправляет новые параметры всем устройствам. Этот подход может быть быстрее, поскольку устройства работают независимо, но может приводить к несогласованности, так как устройства могут работать с несколько разными версиями модели в любой момент времени.
Для получения дополнительной информации о параллелизме данных обратитесь к [39].
Параллелизм модели
При параллелизме модели одна модель разделяется по нескольким устройствам, и каждое устройство отвечает за вычисление только части операций модели. Этот подход полезен, когда модель слишком велика для размещения в памяти одного устройства.
Параллелизм модели можно дополнительно разделить на типы:
- Конвейерный параллелизм (межслойный)
- Тензорный параллелизм (внутрислойный)
Конвейерный параллелизм (PP): при PP слои модели разделяются по нескольким устройствам, а вычисления выполняются конвейерным образом. При прямом проходе каждое устройство передаёт промежуточную активацию следующему устройству в конвейере, а при обратном проходе — отправляет градиент входного тензора предшествующему устройству.
PP особенно полезен при работе с очень глубокими моделями, поскольку позволяет нескольким устройствам работать одновременно, сокращая время простоя и повышая эффективность обучения. Для получения дополнительной информации о PP обратитесь к [40][41].
Тензорный параллелизм (TP): при TP операции внутри одного слоя модели распределяются по нескольким устройствам. Каждое устройство обрабатывает часть вычислений этого слоя, и выходные данные объединяются перед переходом к следующему слою. Например, при крупной операции матричного умножения разные части матрицы могут обрабатываться параллельно на нескольких устройствах — как путём поколончатого, так и построчного разбиения. Для получения дополнительной информации обратитесь к [46].
TP особенно полезен, когда один слой слишком велик для размещения в памяти. Он помогает снижать использование памяти, распределяя вычислительную нагрузку по нескольким устройствам. Если вас интересует дополнительная информация о TP и его вариантах (например, параллелизм последовательностей), обратитесь к [47][48].
Гибридный параллелизм
Гибридный параллелизм объединяет параллелизм данных и модели для более эффективного обучения крупных моделей на нескольких устройствах. Этот подход сокращает использование памяти за счёт распределения как модели, так и данных по устройствам. Он также обеспечивает масштабирование на большее количество устройств, делая возможным обучение очень крупных моделей, с которыми традиционные методы параллелизма не справляются.
В дополнение к гибридному параллелизму, такие методы, как ZeRO (Zero Redundancy Optimizer) [49] от Microsoft и FSDP (Fully Sharded Data Parallel) [50] от Meta, дополнительно оптимизируют использование ресурсов и эффективность коммуникации. Эти методы сокращают избыточность в памяти и вычислениях между устройствами, обеспечивая более эффективное обучение массивных моделей.
Рассмотренные в этом разделе методы являются неотъемлемыми компонентами современного проектирования ML-систем и широко применяются на практике. Сочетая методы параллелизма, такие как FSDP, методы экономии памяти, такие как gradient checkpointing, и оптимизации, такие как AMP, можно эффективно масштабировать обучение крупных сложных моделей. Твёрдое понимание этих методов и знание того, когда их применять, критически важно для построения масштабируемых и эффективных ИИ-систем.
Сэмплирование модели
После обучения модели следующий шаг — сэмплирование, предполагающее генерацию новых данных или выходов из обученной генеративной модели. Существуют различные методы сэмплирования для разных GenAI-приложений. Например, для LLM такие методы, как жадный поиск, лучевой поиск (beam search) [51] и сэмплирование top-k [52], имеют свои достоинства и недостатки. Beam search, например, склонен производить связный и релевантный текст, но может ограничивать разнообразие.
На собеседовании крайне важно подробно обсудить различные методы сэмплирования, их преимущества и недостатки, а также выбрать тот, который подходит для проектируемой системы.
Темы для обсуждения
- Архитектуры моделей: каковы возможные архитектуры для выбранного ML-алгоритма? Каковы плюсы и минусы каждой? Каковы конкретные слои архитектуры и почему они именно такие?
- Методология обучения: какова методология обучения? Как работает процесс обучения (например, процесс диффузии, состязательное обучение)?
- Обучающие данные: откуда берутся обучающие данные? Насколько велик набор данных? Используются ли разные данные для различных этапов обучения (например, предобучение vs fine-tuning)?
- ML-цели: каковы возможные ML-цели для данной задачи? Каковы плюсы и минусы каждой цели и как они влияют на производительность модели?
- Функции потерь: какова функция потерь, соответствующая выбранной ML-цели? Используется одна функция потерь или несколько? Если несколько — как они комбинируются для оптимизации процесса обучения? В чём назначение каждой функции потерь?
- Трудности обучения и их решения: каковы типичные трудности обучения, специфичные для выбранного ML-алгоритма? Как их можно преодолеть для обеспечения эффективного обучения?
- Эффективность обучения: каковы основные методы повышения эффективности обучения? Как работает распределённое обучение и какие преимущества оно даёт? Как AMP повышает скорость и эффективность обучения? Как работают параллелизм данных, тензоров и конвейеров?
- Сэмплирование: как работают различные методы сэмплирования (например, top-k, top-p)? Каковы плюсы и минусы каждого? Как влияют на качество и творческий потенциал выходных данных модели? Какие методы позволяют ускорить процесс сэмплирования без потери качества?
Оценка
После разработки модели следующим критическим этапом является оценка. Она включает использование различных метрик для оценки производительности ML-модели. В этом разделе рассмотрим два метода оценки: офлайн и онлайн.
Офлайн-оценка
Офлайн-оценка — это процесс оценки производительности модели или системы с использованием заранее собранных данных без её развёртывания в среде реального времени. Этот подход критичен для обеспечения эффективности модели до её использования реальными пользователями.
Офлайн-оценка различается для дискриминативных и генеративных моделей. Цель дискриминативных моделей — делать предсказания на оценочном наборе; оценка сравнивает эти предсказания с реальными метками. Традиционные метрики, такие как точность (precision), полнота (recall) и F1-мера, используются для количественной оценки производительности модели на основе этих сравнений. В таблице 2 перечислены распространённые метрики для различных дискриминативных задач, подробно рассмотренных в [1].
| Задача | Метрики |
|---|---|
| Классификация | Precision, Recall, F1 score, Accuracy, Confusion matrix |
| Регрессия | MSE, MAE, RMSE |
| Ранжирование | Precision@k, Recall@k, MRR, mAP, nDCG |
Таблица 2: Популярные метрики в дискриминативных задачах
В генеративных моделях оценка более сложна. Вместо сравнения предсказаний с фиксированными реальными метками эти модели генерируют новый контент: текст или изображения. Оценка такого контента часто требует субъективных методов или участия человека для измерения соответствия ожиданиям и эталонам. В таблице 3 показаны часто используемые метрики для различных генеративных задач.
| Задача | Метрики |
|---|---|
| Генерация текста | Perplexity, BLEU, METEOR, ROUGE, CIDEr |
| Генерация изображений | FID, IS, KID, SWD, PPL, LPIPS |
| Text-to-Video | FVD, CLIPScore, FID, LPIPS, KID |
Таблица 3: Популярные метрики в генеративных задачах
На собеседованиях необходимо оценивать сгенерированный контент с нескольких сторон. Например, в модели text-to-image важно убедиться, что сгенерированное изображение как высококачественно, так и соответствует заданному текстовому запросу. Аналогично для чат-бота возможности модели следует оценивать по различным задачам: математика, здравый смысл и генерация кода. На протяжении книги мы подробно рассматриваем офлайн-оценку различных GenAI-приложений и изучаем большинство метрик из таблицы 3.
Онлайн-оценка
Онлайн-оценка определяет производительность модели в продакшне (то есть после развёртывания). Для оценки влияния модели используются различные метрики, согласованные с бизнес-целями.
На практике компании обычно отслеживают несколько онлайн-метрик. На собеседовании следует сосредоточиться на выборе наиболее критичных для оценки воздействия системы. В отличие от офлайн-метрик, выбор онлайн-метрик более субъективен и часто предполагает участие владельцев продуктов и заинтересованных сторон. Этот шаг помогает интервьюеру оценить ваше бизнес-мышление. Полезно чётко формулировать обоснование и ход мыслей при выборе конкретных метрик. В таблице 4 перечислены некоторые метрики, широко используемые в онлайн-оценке.
| Метрика | Описание |
|---|---|
| Click-Through Rate (CTR) | Процент пользователей, кликнувших на контент или предложения. |
| Conversion Rate | Процент пользователей, выполнивших желаемое действие (покупка, подписка) после взаимодействия с системой. |
| Latency (время инференса) | Время, затрачиваемое моделью на генерацию контента. |
| Engagement Rate | Показатель взаимодействия пользователей, например время, проведённое в системе. |
| Revenue Per User | Средний доход, генерируемый на одного пользователя. |
| Churn Rate | Процент пользователей, прекративших использование системы за определённый период. |
| User Satisfaction | Прямая обратная связь от пользователей об их опыте взаимодействия с контентом, сгенерированным ИИ. |
| User Retention | Процент пользователей, продолжающих использовать систему в течение определённого периода. |
| Completion Rate | Процент задач (генерация текста, изображений), успешно завершённых моделью. |
Таблица 4: Распространённые метрики для онлайн-оценки
Темы для обсуждения
Ключевые темы для обсуждения на этапе оценки:
- Офлайн-метрики: какие офлайн-метрики лучше всего оценивают качество и точность генеративной модели? Как эти метрики измеряют разнообразие, реалистичность и связность генерируемых выходных данных?
- Онлайн-метрики: какие метрики критически важны для оценки эффективности генеративной модели в продакшне? Как эти метрики согласуются с бизнес-целями, такими как повышение творческого потенциала пользователей, увеличение вовлечённости или стимулирование инноваций в продукте?
- Предвзятость: генеративные модели могут непреднамеренно отражать социальные предвзятости в отношении чувствительных атрибутов, таких как пол или раса. Как можно оценить предвзятость модели?
- Надёжность и безопасность: насколько устойчива генеративная модель к состязательным атакам — намеренно вводящим в заблуждение входным данным, разработанным для эксплуатации слабых мест модели?
- Оценка людьми: для генеративных моделей, особенно в творческих областях (генерация текста, синтез изображений), обратная связь от людей жизненно важна. Как рецензенты-люди могут дополнить автоматическую оценку? Какие методы (опросы, A/B-тестирование, экспертные отзывы) лучше всего подходят для оценки производительности модели? Как нивелировать влияние субъективности разных рецензентов?
Общий дизайн ML-системы
Следующий шаг фреймворка — предложить общий дизайн GenAI-системы. Такие системы включают гораздо больше, чем просто обучение модели; они требуют слаженной совместной работы нескольких конвейеров и компонентов. Например, в чат-боте помимо основной модели необходимы другие компоненты для обеспечения безопасности, например фильтрация вредоносного контента. Аналогично для системы генерации видео могут потребоваться дополнительные модели для повышения разрешения видео до желаемого качества. На этом этапе крайне важно интегрировать все компоненты, необходимые для функционирования системы по назначению. На протяжении книги мы рассмотрим несколько компонентов, часто используемых совместно с генеративной моделью.
Темы для обсуждения
- Компоненты системы: каковы различные компоненты системы? Какова роль каждого компонента — основной модели, предобработки, фильтрации контента, постобработки и любых необходимых моделей масштабирования или повышения качества?
- Механизмы безопасности: как в систему интегрированы безопасность и модерация контента? Например, как система обеспечивает безопасность и приемлемость генерируемого контента для пользователей? Опишите компоненты, такие как фильтры NSFW и вредоносного контента.
- Обратная связь от пользователей и непрерывное обучение: как система использует обратную связь от пользователей для постоянного улучшения производительности модели? Обсудите механизмы обратной связи, позволяющие выполнять fine-tuning модели. Какие системы предусмотрены для переобучения моделей на обновлённых данных для повышения точности и релевантности со временем?
- Масштабируемость: как система масштабируется с ростом спроса? Какие облачные или аппаратные ресурсы используются и как ими эффективно управлять? Как компоненты, такие как балансировщики нагрузки, распределённый инференс и параллелизм модели, способствуют масштабируемости системы?
- Соображения безопасности: как система обеспечивает конфиденциальность пользователей, особенно при работе с чувствительными данными или генерации персонализированного контента? Какие протоколы безопасности реализованы для защиты от состязательных атак, фальсификации модели или утечки данных?
- Предвзятость: генеративные модели могут непреднамеренно отражать социальные предвзятости в отношении чувствительных атрибутов, таких как пол или раса. Обсудите стратегии, включая алгоритмы обнаружения предвзятостей, аудит справедливости и фильтрацию предвзятых результатов. Как вы будете решать этические проблемы, если пользователи попытаются использовать генеративную модель для создания вредоносного, предвзятого или неприемлемого контента?
- Надёжность и безопасность: насколько устойчива генеративная модель к состязательным атакам, например к намеренно вводящим в заблуждение входным данным, разработанным для эксплуатации слабых мест модели? Могут ли злоумышленники генерировать вредоносные или бессмысленные выходные данные? В продакшне как убедиться, что модель не используется в злонамеренных целях, таких как создание дипфейков, дезинформации или неприемлемого контента?
Развёртывание и мониторинг
Последний шаг — развернуть систему в продакшне и обслуживать миллионы пользователей. После развёртывания система может давать сбои по многим причинам. Мониторинг — это задача отслеживания, измерения и протоколирования различных метрик для обнаружения сбоев системы при их возникновении, чтобы устранять их как можно быстрее. Однако поскольку эта тема широка и не специфична для GenAI или какой-либо конкретной задачи, мы не будем углубляться в неё в этой книге. Мы рекомендуем обратиться к [53] или книге «ML System Design Interview» [1] для более глубокого изучения.
Резюме
В этой главе мы дали обзор GenAI и представили фреймворк для подхода к собеседованию по проектированию GenAI-системы. Хотя некоторые детали специфичны для GenAI, многие концепции применимы к проектированию ИИ-систем в целом. Мы сосредоточимся на аспектах, уникальных для GenAI, избегая общих тем, присущих всем ИИ-системам, таких как развёртывание, инфраструктура и мониторинг.
Наконец, не от каждого инженера ожидается экспертиза во всех областях жизненного цикла GenAI. Разные роли и компании могут делать упор на различных аспектах: инфраструктуре, мониторинге или разработке LLM. Этот фреймворк помогает кандидатам понять ожидания и адаптировать ответы в соответствии с акцентами конкретного собеседования.
Теперь, когда вы понимаете эти основы, мы готовы приступить к рассмотрению наиболее распространённых вопросов для собеседований по проектированию GenAI-систем.
Ссылки
[1] Machine Learning System Design Interview. https://www.aliaminian.com/books. [2] Support Vector Machines. https://scikit-learn.org/stable/modules/svm.html. [3] Bayes' theorem. https://en.wikipedia.org/wiki/Bayes%27_theorem. [4] Gaussian mixture models. https://scikit-learn.org/1.5/modules/mixture.html. [5] Hidden Markov model. https://en.wikipedia.org/wiki/Hidden_Markov_model. [6] Boltzmann machine. https://en.wikipedia.org/wiki/Boltzmann_machine. [7] OpenAI's ChatGPT. https://openai.com/index/chatgpt/. [8] Economic Potential of Generative AI. https://www.mckinsey.com/capabilities/mckinsey-digital/our-insights/the-economic-potential-of-generative-ai-the-next-productivity-frontier. [9] The Llama 3 Herd of Models. https://arxiv.org/abs/2407.21783. [10] Flamingo: a Visual Language Model for Few-Shot Learning. https://arxiv.org/abs/2204.14198. [11] PaLM: Scaling Language Modeling with Pathways. https://arxiv.org/abs/2204.02311. [12] Language Models are Few-Shot Learners. https://arxiv.org/abs/2005.14165. [13] Photorealistic Text-to-Image Diffusion Models with Deep Language Understanding. https://arxiv.org/abs/2205.11487. [14] PaLM2 Technical Report. https://arxiv.org/abs/2305.10403. [15] H100 Tensor Core GPU. https://www.nvidia.com/en-us/data-center/h100/. [16] GPT-4 training cost. www.wired.com/story/openai-ceo-sam-altman-the-age-of-giant-ai-models-is-already-over/. [17] Scaling Laws for Neural Language Models. https://arxiv.org/abs/2001.08361. [18] Training Compute-Optimal Large Language Models. https://arxiv.org/abs/2203.15556. [19] Introducing OpenAI o1. https://openai.com/index/introducing-openai-o1-preview/. [20] Large Language Monkeys: Scaling Inference Compute with Repeated Sampling. https://arxiv.org/abs/2407.21787. [21] ETL. https://aws.amazon.com/what-is/etl/. [22] Tecton. https://www.tecton.ai/feature-store/. [23] Amazon SageMaker. https://aws.amazon.com/sagemaker/. [24] ML System Design Interview. https://www.amazon.com/gp/product/1736049127/. [25] Comprehensive Exploration of Synthetic Data Generation: A Survey. https://arxiv.org/abs/2401.02524. [26] HDFS Architecture Guide. https://hadoop.apache.org/docs/r1.2.1/hdfs_design.html. [27] Amazon S3. https://aws.amazon.com/s3/. [28] Apache Parquet. https://parquet.apache.org/. [29] Apache ORC. https://orc.apache.org/docs/. [30] Apache Lucene. https://lucene.apache.org/. [31] Elasticsearch. https://www.elastic.co/elasticsearch. [32] Attention Is All You Need. https://arxiv.org/abs/1706.03762. [33] BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. https://arxiv.org/abs/1810.04805. [34] An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. https://arxiv.org/abs/2010.11929. [35] Learning Transferable Visual Models From Natural Language Supervision. https://arxiv.org/abs/2103.00020. [36] Zero-Shot Text-to-Image Generation. https://arxiv.org/abs/2102.12092. [37] Neural Machine Translation by Jointly Learning to Align and Translate. https://arxiv.org/abs/1409.0473. [38] Common Crawl. https://commoncrawl.org/. [39] Data Parallelism. https://en.wikipedia.org/wiki/Data_parallelism. [40] Model Parallelism. https://huggingface.co/docs/transformers/v4.15.0/en/parallelism. [41] Pipeline Parallelism. https://pytorch.org/docs/stable/distributed.pipelining.html. [42] Mixed Precision Training. https://arxiv.org/abs/1710.03740. [43] High-Resolution Image Synthesis with Latent Diffusion Models. https://arxiv.org/abs/2112.10752. [44] Training Deep Nets with Sublinear Memory Cost. https://arxiv.org/abs/1604.06174. [45] Automatic Mixed Precision. https://pytorch.org/tutorials/recipes/recipes/amp_recipe.html. [46] Model parallelism. https://huggingface.co/docs/transformers/v4.17.0/en/parallelism. [47] Paradigms of Parallelism. https://colossalai.org/docs/concepts/paradigms_of_parallelism/. [48] Tensor Parallelism tutorial. https://pytorch.org/tutorials/intermediate/TP_tutorial.html. [49] ZeRO: Memory Optimizations Toward Training Trillion Parameter Models. https://arxiv.org/abs/1910.02054. [50] Introducing PyTorch Fully Sharded Data Parallel (FSDP) API. https://pytorch.org/blog/introducing-pytorch-fully-sharded-data-parallel-api/. [51] Beam search. https://en.wikipedia.org/wiki/Beam_search. [52] Top-k sampling. https://docs.cohere.com/docs/controlling-generation-with-top-k-top-p. [53] Model monitoring for ML in production. https://www.evidentlyai.com/ml-in-production/model-monitoring.
Gmail Smart Compose
Введение
Функция Smart Compose в Gmail [1] помогает пользователям, предлагая следующие несколько слов по мере написания письма. В этой главе рассматривается данная функция, а также архитектура Transformer, лежащая в основе большинства генеративных систем.
Уточнение требований
Ниже представлен типичный диалог между кандидатом и интервьюером:
Кандидат: У разных пользователей могут быть разные стили письма. Должна ли система предоставлять персонализированные подсказки? Интервьюер: Для простоты давайте не будем включать персонализацию.
Кандидат: Должна ли система предлагать следующие несколько слов только тогда, когда она уверена в своём прогнозе? Интервьюер: Да.
Кандидат: Набор данных по электронным письмам должен быть достаточно большим для обучения модели. Известен ли нам приблизительный объём данных? Интервьюер: Предположим, что наш набор данных состоит примерно из одного миллиарда сообщений электронной почты.
Кандидат: При формировании подсказок можно использовать разные части данных. Например, прошлые письма пользователя или тему текущего письма. Чтобы упростить задачу, могу ли я использовать в качестве контекста только тело письма? Интервьюер: Хорошее замечание. На практике мы используем больше, чем то, что пользователь уже написал в текущем письме. Давайте начнём с тела письма. Если останется время, можно расширить контекст, включив другую релевантную информацию.
Кандидат: Какие языки должна поддерживать система? Интервьюер: Начнём с английского.
Кандидат: Нужно ли убедиться, что система не проявляет предвзятости? Интервьюер: Это важное требование для данной системы. Система не должна делать предвзятых предположений при формировании подсказок.
Кандидат: Сколько активных пользователей у Gmail? Является ли стоимость вычислений проблемой для этой функции? Интервьюер: У Gmail около 1,8 миллиарда пользователей, и один пользователь может отправлять до 500 писем в день. Мы заботимся о вычислительных затратах, но давайте сначала сосредоточимся на разработке системы. Оптимизацию эффективности можно провести в последующих итерациях.
Кандидат: Должна ли система формировать подсказки в режиме реального времени? Интервьюер: Да. Ожидаемая задержка должна быть незаметной; около 100 миллисекунд будет достаточно.
Формулировка задачи как задачи ML
В этом разделе мы формулируем функцию Smart Compose как задачу ML. Для этого необходимо понять входные и выходные данные системы и выбрать подходящий подход ML для решения задачи.
Определение входных и выходных данных системы
Входные данные модели — это последовательность слов, введённых пользователем. Выходные данные — продолжение этой последовательности. Модель генерирует слова, которые пользователь, вероятнее всего, напишет следующими.
Выбор подходящего подхода ML
Smart Compose генерирует текстовый контент, поэтому мы относим её к задачам генерации текста. Различные архитектуры ML предназначены для обработки последовательных данных, что необходимо для генерации текста. Две популярные архитектуры — рекуррентные нейронные сети (RNN) [2] и Transformer'ы [3].
Transformer'ы имеют ряд преимуществ перед RNN, из которых можно выделить два основных:
- Параллелизм: в RNN результаты вычислений одного временного шага передаются на следующий, образуя цепочку операций, зависящих от времени. Transformer'ы, напротив, могут обрабатывать все входные токены одновременно благодаря механизму самовнимания (self-attention).
- Лучшая обработка длинных последовательностей: Transformer'ы используют механизмы самовнимания, чтобы сосредоточиться на любой части последовательности, независимо от расстояния. RNN, напротив, плохо справляются с долгосрочными зависимостями из-за своей последовательной структуры и проблемы затухающего градиента.
Благодаря этим преимуществам Transformer'ы демонстрируют выдающиеся результаты в задачах генерации текста и поэтому используются в большинстве современных генеративных систем. Именно поэтому мы выбрали Transformer'ы для реализации функции Smart Compose.
| Feature | RNN (GRU [4], LSTM [5]) | Transformer |
|---|---|---|
| Architecture | Simple | Complex |
| Training efficiency | Inefficient due to sequential processing | Efficient due to parallel processing |
| Effectiveness | Low as it struggles with long sequences | High as it handles long sequences |
| Scalability | Limited scalability | Highly scalable |
| Applications | Simple tasks such as time series modeling | Complex tasks such as language completion or translation |
Таблица 1: Сравнение архитектур RNN и Transformer
Хотя Transformer'ы лучше поддаются распараллеливанию благодаря отсутствию строгих последовательных зависимостей, их механизм самовнимания имеет вычислительную сложность O(n2)O(n^2)O(n2), где nnn — длина последовательности. Эта сложность обусловлена тем, что механизм самовнимания требует вычисления оценок внимания между каждой парой токенов в последовательности. Для снижения сложности внимания предложены различные методы. Подробнее см. Group Attention [6] и Flash Attention [7].
Подготовка данных
На этапе подготовки данных мы преобразуем исходные данные в формат, ожидаемый моделью ML. Сначала кратко рассмотрим доступные данные.
Для обучения модели доступны два источника данных: общие данные и данные электронной почты. Общие данные включают общедоступные тексты из таких источников, как книги, веб-сайты и публикации в социальных сетях. Эти данные важны для обучения языковых моделей, поскольку содержат разнообразный словарный запас, синтаксис и контексты.
Данные электронной почты, как указано в требованиях, состоят из одного миллиарда сообщений. Эти данные необходимы для того, чтобы модель усвоила стили написания писем и распространённые фразы. Таблица 2 показывает упрощённый пример данных электронной почты. На практике для каждого сообщения хранится больше метаданных.
| Email ID | Sender | Recipient | Subject | Body |
|---|---|---|---|---|
| 4953 | [email protected] | [email protected] | Catchup? | Hey Mike, let's catch up this Sat. … |
| 9356 | [email protected] | [email protected] | Project Deadline | Hi TA, I hope you are well. I am writing to you to … |
Таблица 2: Пример данных электронной почты
Исходный текст как в общих данных, так и в данных электронной почты зачастую содержит шум и несоответствия, которые могут снизить производительность модели. Кроме того, модели ML требуют данные в числовом формате. По этим причинам исходный текст необходимо подготовить с помощью двух ключевых шагов:
- Очистка и нормализация текста
- Токенизация текста и индексация токенов
Очистка и нормализация текста
Очистка текста
Очистка текста удаляет ненужную или нерелевантную информацию. Распространённые методы включают:
- Удаление неанглийского текста: используйте методы идентификации языка [8], такие как [9], для выявления и удаления неанглийского текста из общих данных и данных электронной почты.
- Удаление конфиденциальной информации: письма могут содержать конфиденциальные данные, такие как номера телефонов и кредитных карт. Эти сведения необходимо удалить, чтобы модель не обучилась на них и не раскрыла их впоследствии. Мы заменяем личные имена, URL, адреса электронной почты и номера телефонов символами-заполнителями. Например, «[email protected]» заменяется на «##@gmail.com».
- Удаление нерелевантных символов: удаление ненужных или нерелевантных символов и знаков, не несущих смысловой нагрузки. Например, такие символы, как «©», «™» или эмодзи, удаляются, поскольку они обычно не влияют на смысл текста.
- Удаление дублирующихся данных: дублирующиеся данные — это одинаковые тексты из разных источников, которые встречаются в наборе данных несколько раз. Дубликаты удаляются, чтобы предотвратить смещение модели и искажение процесса её обучения.
Нормализация текста
Нормализация текста преобразует текст в единый формат. Например, она приводит различные способы записи номера телефона — «(123) 456-7890», «123.456.7890» и «123-456-7890» — к стандартному формату, например «1234567890». Нормализация текста обеспечивает единообразие и снижает сложность текстовых данных.
Далее мы преобразуем исходный текст в последовательность чисел с помощью токенизации текста и индексации токенов.
Токенизация текста и индексация токенов
Токенизация текста с последующей индексацией токенов преобразует исходный текст в формат, ожидаемый моделью Transformer: последовательность чисел.
Рассмотрим каждый шаг подробнее.
Токенизация текста
Токенизация текста — это процесс разбиения текста на более мелкие единицы, называемые токенами. На рисунке 5 показано, как GPT-4 от OpenAI токенизирует предложение «Let's go to NYC».1
Токенизация может выполняться на разных уровнях. Например, «Hello world» можно разбить на [«Hello», «world»] или на [«H», «e», «l», «l», «o», « », «w», «o», «r», «l», «d»]. В целом алгоритмы токенизации делятся на три категории:
- Токенизация на уровне символов
- Токенизация на уровне слов
- Токенизация на уровне подслов
Понимание каждой категории токенизации и её достоинств и недостатков крайне важно для большинства интервью по ML. Рассмотрим их подробнее.
Токенизация на уровне символов
Токенизация на уровне символов разбивает текст на набор символов. Она проста в реализации, но модели сложно обучиться осмысленным представлениям для каждого токена. Например, труднее научиться осмысленному представлению буквы «g», чем слова «go», поскольку «go» имеет чёткое значение, тогда как «g» — нет. Из-за этого токенизация на уровне символов нередко приводит к снижению производительности.
Токенизация на уровне слов
Токенизация на уровне слов разбивает текст на отдельные слова. Существуют разные алгоритмы для этого, однако простейший — разбиение текста по пробелам.
Преимущество токенизации на уровне слов состоит в том, что модели проще обучиться осмысленным представлениям для каждого токена. Однако главный недостаток — как правило, очень большой размер словаря. Например, Transformer-XL [10] использует токенизатор на уровне слов, что даёт словарь из 267 735 токенов. Большой размер словаря проблематичен, поскольку модели приходится обучать представления для сотен тысяч токенов. Это делает обучение трудоёмким и, следовательно, более затратным по сравнению с токенизацией на уровне символов.
Рассмотрим токенизацию на уровне подслов, которая обеспечивает баланс между токенизацией на уровне слов и на уровне символов.
Токенизация на уровне подслов
Токенизация на уровне подслов разбивает текст на более мелкие единицы, называемые подсловами. Она основана на принципе: часто используемое слово не следует разбивать на подслова, тогда как редкое слово следует разбивать на более мелкие значимые подслова. Например, «unhappily» может считаться редким словом и быть разбито на «unhappy» и «ly». Оба подслова встречаются в текстовых данных чаще, что упрощает обучение модели осмысленному представлению для каждого из них.
Хотя токенизация на уровне подслов может быть сложна в реализации, она обладает рядом преимуществ. Во-первых, она обеспечивает управляемый размер словаря, снижая стоимость обучения представлений для каждого подслова. Во-вторых, токенизация на уровне подслов позволяет модели представлять незнакомые слова, разбивая их на известные подслова.
Таблица 3 ниже сравнивает характеристики трёх категорий токенизации.
| Characteristics | Character-level | Word-level | Subword-level |
|---|---|---|---|
| Granularity | Individual characters | Individual words | Subwords |
| Vocabulary size | Small | Large | Moderate |
| Algorithm complexity | Simple | Simple | Complex |
| Handling unseen words | Decomposes unseen words into characters | Cannot easily handle unseen words | Decomposes unseen words into known subwords |
| Vocabulary size | ~100 | ~300,000+ | ~50,000–150,000 |
| Performance | Poor performance | High performance but not practical | High performance and practical |
Таблица 3: Сравнение различных категорий токенизации
Какая токенизация подходит для функции Smart Compose?
Большинство передовых языковых моделей используют алгоритмы токенизации на уровне подслов, такие как Byte-Pair Encoding (BPE) [11] и SentencePiece [12]. Эти алгоритмы более эффективны и способны работать с несколькими языками. Например, GPT-4 от OpenAI использует вариант BPE [13], а Gemini от Google — SentencePiece [14].
Учитывая эффективность токенизации на уровне подслов, мы используем её в качестве текстового токенизатора для функции Smart Compose. Для токенизации текста мы опираемся на популярные Python-библиотеки: Tiktoken [13] от OpenAI или SentencePiece [15] от Google. Эти библиотеки надёжно реализованы и поддерживают различные алгоритмы токенизации.
В главе 3 мы подробно рассмотрим BPE и его алгоритмы. Для более глубокого изучения алгоритмов токенизации на уровне подслов см. [16].
Индексация токенов
Индексация токенов — это процесс преобразования текстовых токенов в целые числа.
В рамках подготовки к индексации токенов алгоритм токенизации сначала строит словарь — набор всех уникальных токенов — из обучающих текстовых данных и сохраняет его в таблице. На рисунке 9 показаны примеры словарей для различных категорий токенизации. Порядок и значения идентификаторов выбраны произвольно в демонстрационных целях.
После того как алгоритм токенизации построил словарь, можно преобразовать любой токен в число и любое число обратно в токен. На рисунке 10 показана индексация токенов с использованием словаря GPT-4 [17].
Подводя итог этапу подготовки данных: сначала мы очищаем и нормализуем текстовые данные, чтобы обеспечить высокое качество и единообразие обучающих данных. Затем мы используем алгоритм токенизации на уровне подслов, такой как BPE, чтобы разбить текст на текстовые токены (подслова), после чего заменяем каждый токен его числовым индексом. Эти шаги обеспечивают представление обучающих данных в числовом формате, пригодном для использования моделью ML.
Разработка модели
Функция Smart Compose представляет собой задачу генерации текста, в которой модель Transformer предсказывает наиболее вероятное продолжение предложений в письме. В этом разделе мы рассматриваем детали архитектуры Transformer, стратегии обучения и методы сэмплирования для разработки модели генерации текста.
Архитектура
Архитектура Transformer, представленная в статье «Attention Is All You Need» [3], предназначена для обработки последовательностей. Это делает её идеальной для задач, требующих понимания текста и взаимосвязей между его словами. Например, в функции Smart Compose модель обрабатывает последовательность слов, уже введённых пользователем, чтобы предложить следующие слова.
Transformer'ы имеют три основных варианта:
- Только энкодер (encoder-only)
- Только декодер (decoder-only)
- Энкодер-декодер (encoder-decoder)
Каждый вариант имеет незначительные архитектурные различия, благодаря которым он подходит для определённых задач. Кратко рассмотрим каждый вариант и его применение.
Только энкодер (encoder-only)
Transformer типа encoder-only используется для задач, требующих понимания общего смысла текста. Он обрабатывает входную последовательность целиком и делает прогнозы на её основе. Например, в задаче анализа тональности encoder-only Transformer предсказывает тональность входного предложения.
Transformer'ы типа encoder-only широко применяются в таких задачах, как классификация предложений и распознавание именованных сущностей, которые сосредоточены на понимании входных данных, а не на генерации нового контента. BERT от Google [18] — известный пример Transformer'а типа encoder-only. Однако эти модели, как правило, не используются для генерации новых последовательностей. Transformer'ы типа decoder-only, напротив, специально разработаны именно для этой цели.
Только декодер (decoder-only)
Transformer типа decoder-only обрабатывает входную последовательность и итеративно генерирует новую последовательность.
Transformer'ы типа decoder-only широко используются в генеративных задачах, включая генерацию текста, где модель генерирует последовательность по одному токену за раз на основе ранее сгенерированных токенов. Большинство больших языковых моделей (LLM), таких как GPT-4 от OpenAI [19], LLaMA от Meta [20] и Gemini от Google [14], основаны на Transformer'е типа decoder-only.
Энкодер-декодер (encoder-decoder)
Архитектура encoder-decoder использует как encoder-only, так и decoder-only Transformer'ы. Компонент-энкодер обрабатывает входную последовательность, а декодер использует эту обработанную информацию для генерации выходной последовательности.
Transformer типа encoder-decoder особенно подходит для задач, в которых выходные данные являются преобразованием входных. Например, в задаче перевода текста входное предложение на одном языке преобразуется в эквивалентное предложение на другом. Мы рассмотрим эту архитектуру в главе 3.
На рисунке 14 ниже показаны широко используемые модели, применяющие различные варианты Transformer'ов.
Какой вариант Transformer подходит для функции Smart Compose?
Выбор между Transformer'ами типа encoder-only, decoder-only и encoder-decoder зависит от того, является ли задача генеративной или направленной на понимание. Smart Compose — это задача генерации текста, цель которой — завершить частично написанный текст. Поэтому Transformer типа decoder-only идеально подходит для этой задачи благодаря способности генерировать текст на основе заданной последовательности.
На собеседованиях по проектированию систем ML обычно акцент делается на концепциях высокого уровня и взаимодействии компонентов, а не на архитектурных деталях. Мы дадим краткий обзор архитектуры Transformer, не углубляясь в детали. Для более глубокого понимания архитектур Transformer см. [21] и [22].
Transformer типа decoder-only состоит из следующих компонентов:
- Текстовые эмбеддинги
- Позиционное кодирование
- Transformer
- Голова прогнозирования (prediction head)
Текстовые эмбеддинги
Компонент текстовых эмбеддингов преобразует каждый идентификатор токена в вектор фиксированной длины, называемый «эмбеддингом». Эмбеддинги обычно хранятся в таблице, как показано на рисунке 15, и обучаются в процессе тренировки модели.
Текстовые эмбеддинги играют ключевую роль в Transformer'е типа decoder-only. Разберёмся, почему.
В процессе подготовки данных мы токенизировали текст и преобразовали токены в идентификаторы. Однако в том, как представлен текст, есть два существенных ограничения:
- Разреженность: словарь, как правило, включает десятки тысяч идентификаторов токенов. Представление этих идентификаторов с помощью унарного кодирования (one-hot encoding) порождает разреженные данные высокой размерности, что неэффективно.
- Отсутствие семантической информации: идентификаторы токенов произвольны и не отражают никаких отношений между словами. Например, слова «happy» и «joyful» могут быть близки по смыслу, однако их идентификаторы не отражают этого сходства.
Компонент текстовых эмбеддингов устраняет оба этих ограничения, преобразуя идентификаторы токенов в обучаемые эмбеддинги. Поскольку эмбеддинги представляют собой плотные векторы в пространстве меньшей размерности, разреженность перестаёт быть проблемой.
Кроме того, поскольку эмбеддинги обучаются в процессе тренировки модели, они фиксируют семантические значения. Например, эмбеддинги слов «happy» и «joyful» окажутся ближе друг к другу в пространстве эмбеддингов, чем эмбеддинги «happy» и «sad», как показано на рисунке 16.
Позиционное кодирование
Transformer-ы изначально не учитывают порядок входных токенов. Если обратиться к формуле вычисления внимания, am, n=exp(qmknd)j=1Nexp(qmknd),
видно, что она инвариантна к перестановкам, то есть механизм внимания не принимает во внимание позиции токенов в последовательности. Например, Transformer не может различить фразы «инициализируй переменную, затем используй её» и «используй переменную, затем инициализируй её». Это снижает способность модели понимать или генерировать связный текст.
Чтобы преодолеть это ограничение, позиционное кодирование предоставляет Transformer-у информацию о положении каждого токена во входной последовательности. Без позиционного кодирования модель воспринимает входную последовательность как набор слов (bag of words), что проблематично. С позиционным кодированием позиция каждого токена кодируется с помощью функции позиционного кодирования, pi=f(i), где f(⋅)f(⋅)f(⋅) — функция позиционного кодирования, а iii — позиция токена. Это позволяет модели различать фразы «используй переменную, затем инициализируй её» и «инициализируй переменную, затем используй её».
Позиционное кодирование реализуется двумя распространёнными способами:
- Фиксированное позиционное кодирование
- Обучаемое позиционное кодирование
Фиксированное позиционное кодирование
В этом методе для отображения позиции (целого числа) в вектор фиксированного размера используется фиксированная функция. В оригинальной статье о Transformer в качестве функции позиционного кодирования были предложены синусно-косинусные функции различных частот.
Рисунок 19 иллюстрирует пример синусно-косинусного позиционного кодирования, показывая векторные представления для четырёх различных позиций. Для простоты в примере используется размерность вектора, равная четырём. На практике эта размерность обычно совпадает с размерностью эмбеддингов токенов, чтобы их можно было складывать (см. рисунок 17).
Рассмотрим плюсы и минусы фиксированного позиционного кодирования.
Плюсы:
- Эффективность: фиксированные кодировки не добавляют в модель дополнительных обучаемых параметров, что делает их вычислительно более эффективными.
- Поддержка длинных последовательностей: фиксированные методы способны отобразить любую позицию в представление. Эта гибкость позволяет модели обрабатывать последовательности, превышающие по длине те, что встречались в обучающих данных.
Минусы:
- Предопределённые ограничения: некоторые методы фиксированного кодирования требуют заранее заданной максимальной позиции, что ограничивает их применимость последовательностями ниже этого максимума.
- Неоптимальная производительность: в ряде задач фиксированные кодировки могут уступать обучаемым методам по качеству захвата позиционных зависимостей, что приводит к снижению производительности.
Обучаемое позиционное кодирование
В этом методе позиционные представления обучаются в ходе тренировки. Конкретно: инициализируется матрица весов P∈RN×dP \in \mathbb{R}^{N \times d}P∈RN×d, где NNN — максимальная длина последовательности, а ddd — размерность эмбеддингов. Эта матрица PPP рассматривается как обучаемый параметр и оптимизируется совместно с остальными параметрами модели.
Обучаемое позиционное кодирование имеет следующие плюсы и минусы.
Плюсы:
- Оптимальная производительность: поскольку эмбеддинги обучаются на основе тренировочных данных, обучаемое позиционное кодирование может обеспечить оптимальное представление позиций для конкретной задачи.
Минусы:
- Неэффективность: требует обучения дополнительных параметров, что увеличивает время обучения и вычислительные затраты.
- Недостаточная обобщаемость: обученные эмбеддинги могут переобучиться под конкретные длины последовательностей, встречавшиеся при тренировке. Если модель в основном видела последовательности определённой длины, она может неэффективно представлять другие позиции, что снижает её способность к обобщению на разнообразные позиции.
Подводя итог: выбор между обучаемым и фиксированным позиционным кодированием зависит от ограничений задачи, в том числе от ожидаемого разнообразия длин последовательностей. Ряд работ, включая оригинальную статью о Transformer, применяет фиксированное позиционное кодирование ввиду его эффективности и лучшей обобщаемости. Следуя этой практике, для обучения функции Smart Compose мы используем фиксированное позиционное кодирование, в частности синусно-косинусное.
Transformer
Компонент Transformer принимает на вход последовательность эмбеддингов и преобразует её в обновлённую последовательность эмбеддингов.
Архитектура Transformer состоит из стека блоков. Каждый блок содержит следующие элементы:
- Многоголовое внимание (multi-head attention): этот слой обновляет каждый эмбеддинг с помощью механизма внимания. Механизм внимания фиксирует зависимости в последовательности, позволяя каждому эмбеддингу обращаться к предшествующим эмбеддингам. В силу природы своего механизма многоголовое внимание широко известно как само-внимание (self-attention) — именно этот термин будет использоваться далее по тексту книги.
- Прямой проход (feed forward): этот слой независимо применяет к каждому эмбеддингу в последовательности два линейных преобразования с активацией ReLU между ними.
Архитектура Transformer включает такие детали, как остаточные соединения, нормализацию слоёв и слои дропаута. Для глубокого понимания этих компонентов рекомендуем обратиться к статье «Attention Is All You Need» [3] и [21].
Голова предсказания
Голова предсказания — финальный компонент decoder-only Transformer — преобразует выход Transformer в вероятности для каждого токена в словаре (рисунок 22). Эти вероятности используются для выбора наиболее вероятного следующего токена.
Обучение
Обучение корректирует параметры decoder-only Transformer с использованием данных электронной почты. По завершении процесса обучения модель способна предлагать вероятные варианты завершения текста.
Однако непосредственное обучение модели на задаче-специфическом наборе данных, например на данных электронной почты, не является хорошей стратегией. Такой прямой подход сопряжён с рядом трудностей:
- Нехватка большого объёма обучающих данных: задаче-специфические наборы данных, как правило, невелики по размеру, что может мешать модели эффективно обучаться.
- Риск переобучения: при обучении модели на задаче-специфическом наборе данных существует высокий риск переобучения. Переобучение происходит, когда модель запоминает обучающие данные настолько, что теряет способность обобщаться на новые данные.
- Дорогостоящее и длительное обучение: обучение большой модели с нуля требует значительных вычислительных ресурсов и времени, поскольку модель должна освоить различные аспекты языка — сложный и ресурсоёмкий процесс.
Для решения перечисленных проблем широко применяется двухэтапная стратегия обучения: предобучение (pretraining) с последующей дообучением (finetuning). На этапе предобучения модель тренируется на большом объёме общих данных для изучения структуры языка. На этапе дообучения предобученная модель дополнительно обучается на данных, специфичных для конкретной задачи (например, на завершении электронных писем).
Эта двухэтапная стратегия использует форму переноса обучения (transfer learning): общие знания, полученные на этапе предобучения, передаются на этап дообучения. Перенос выгоден тем, что модели не нужно начинать с нуля при освоении новой задачи — вместо этого она адаптирует уже предобученные веса, что значительно эффективнее.
Рассмотрим подробнее каждый этап и разберём необходимые обучающие данные, цели ML и функции потерь для каждого из них.
. Предобучение
Предобучение предполагает тренировку модели на большом объёме общих текстовых данных. Эти данные, как правило, разнообразны и охватывают широкий спектр тем и языковых структур. Цель предобучения — сформировать модель, способную понимать естественный язык, включая синтаксис, общие знания и языковые структуры.
Данные для предобучения
Данные для предобучения на этом этапе обычно представляют собой большой объём общих текстов из различных источников в интернете: веб-страницы, книги, социальные сети. Например, Common Crawl [23] — общедоступный набор данных, собранный путём обхода большого числа веб-страниц в интернете. Он содержит петабайты данных, регулярно собираемых с 2008 года.
Цель ML и функция потерь
Цель ML — это формализованная задача, которую стремится решить процесс обучения. В случае генерации текста наиболее распространённой целью ML является «предсказание следующего токена» (next-token prediction). При этой цели модель должна предсказать следующий токен по заданной последовательности предыдущих токенов. Например, в предложении «I hope you are __» модель должна присвоить высокую вероятность слову «well» как следующему токену.
Предсказание следующего токена хорошо подходит для задач генерации текста, поскольку после обучения модель способна строить предложения постепенно. Например, получив на вход «I ordered food because I», модель может предсказать «was» как следующее слово. Затем процесс повторяется с новой последовательностью «I ordered food because I was», что приводит к следующему предсказанию — возможно, «hungry». Этот итеративный процесс продолжается до тех пор, пока модель не предскажет «⟨\langle⟨EOS⟩\rangle⟩» — специальный токен, обозначающий конец последовательности. На рисунке 25 показан пошаговый процесс генерации текста с помощью предсказания следующего токена.
Для оптимизации модели с целью корректного предсказания следующего токена определяется функция потерь, направляющая процесс обучения. Кросс-энтропийные потери [24] — широко используемая функция потерь для задачи предсказания следующего токена. Эта функция измеряет расхождение между предсказанными вероятностями и правильным токеном, позволяя оптимизатору обновлять параметры модели для получения более точных вероятностей в будущем.
На практике модель обрабатывает токены всех длин в последовательности параллельно, что позволяет одновременно вычислять потери для каждой позиции токена. Параллелизация этого шага ускоряет обучение, обрабатывая несколько токенов сразу, а не последовательно.
. Дообучение
Дообучение предполагает адаптацию базовой модели, полученной на этапе предобучения, к конкретной задаче, например к завершению электронных писем. На этом этапе модель обучается мастерству в конкретной задаче на меньшем, задаче-специфическом наборе данных. В ходе дообучения модель сохраняет языковое понимание, приобретённое на этапе предобучения, и адаптируется к особенностям задачи.
Данные для дообучения
Мы используем набор данных, содержащий около одного миллиарда переписок по электронной почте, как указано в разделе требований. Эти данные включают различные форматы писем, как формальный, так и неформальный стили, а также специфическую лексику, более характерную для переписки по электронной почте.
Цель ML и функция потерь
На этапе дообучения и цель ML, и функция потерь остаются неизменными. Целью ML по-прежнему является предсказание следующего токена, а обучение направляется функцией кросс-энтропийных потерь. Единственное отличие от этапа предобучения состоит в том, что потери вычисляются на основе данных электронной почты, ориентируясь на предсказание следующего токена в контексте письма.
Однако опираться исключительно на тело письма как на единственный входной сигнал не слишком эффективно, поскольку таким образом предсказать следующий токен удаётся не всегда. Представьте пользователя, который хочет ответить на письмо от Джона. Когда пользователь вводит «Дорогой», модель должна в идеале предложить «Джон». Если же эта информация не предоставлена на вход, модель не сможет предсказать «Джон» как вероятный следующий токен.
Чтобы устранить это ограничение, мы включаем во входные данные дополнительную информацию. Например, можно использовать тему письма, получателя и предыдущие письма, если они доступны. Это обогащает контекст и помогает модели делать более релевантные предсказания.
Объединение различных входных данных
В традиционном ML архитектура модели, как правило, зависит от типа обрабатываемых данных. Это требует настройки предобработки и разработки признаков под разные типы данных: текст, изображения или таблицы.
В эпоху GenAI архитектура модели зачастую отделена от структуры входных данных. Это разделение повышает гибкость, позволяя одной и той же архитектуре модели обрабатывать разнообразные входные данные в рамках единой архитектуры, упрощая разработку и расширяя универсальность систем GenAI. Разделение достигается с помощью таких техник, как проектирование промптов (prompt engineering) [25]. В главе 6 мы подробно рассмотрим проектирование промптов.
Для объединения различных входных данных в Gmail Smart Compose, как показано на рисунке 31, мы объединяем несколько текстовых входных данных в одну последовательность с тегами, используя шаблон промпта. Нам не нужно беспокоиться об отсутствующих необязательных полях, если наш обучающий набор содержит подобные примеры. Модель обрабатывает различные комбинации входных данных вне зависимости от того, содержат ли они все детали или лишь частичную информацию. Эта гибкость демонстрирует продуманный дизайн модели, позволяющий генерировать контекстуально уместные выходные данные даже при неполных входных данных. Включая разнообразные сценарии в обучающие данные, мы обеспечиваем хорошую обобщаемость модели на различные структуры входных данных с получением надёжных результатов.
Преимущества двухэтапного обучения
Двухэтапная стратегия обучения имеет ряд преимуществ, в том числе:
- Адаптируемость: одна и та же базовая модель, полученная на этапе предобучения, может быть адаптирована для различных задач.
- Улучшенная обобщаемость: предобучение на больших и разнообразных текстовых данных позволяет модели выработать широкое понимание языка, что способствует лучшей обобщаемости на различные задачи.
- Быстрое дообучение: модель усваивает общие знания на этапе предобучения, что ускоряет последующий процесс дообучения.
- Работа с дефицитом данных: для задач, где большие наборы данных недоступны, знания, полученные в ходе предобучения, могут компенсировать их нехватку. Это позволяет модели хорошо работать даже при ограниченном объёме задаче-специфических данных.
- Снижение риска переобучения: при обучении модели с нуля на небольшом задаче-специфическом наборе данных существует риск переобучения. В двухэтапном обучении предобучение выступает в роли регуляризации: модель сначала учится широко понимать язык, и лишь затем сосредотачивается на особенностях конкретной задачи.
- Оптимизация ресурсов: разбивая процесс обучения на два этапа, мы выполняем вычислительно затратное предобучение один раз и можем повторно использовать одну и ту же модель для адаптации к разным задачам. Это снижает вычислительные затраты, поскольку нет необходимости повторять этап предобучения для каждой задачи.
Сэмплирование
Генеративные модели обучаются воспроизводить лежащее в основе распределение обучающих данных. После обучения эти модели способны генерировать новые образцы, схожие с теми, на которых обучались. Сэмплирование — это процесс использования обученной генеративной модели для создания новых данных.
В контексте Smart Compose сэмплирование предполагает генерацию вероятного завершения письма на основе частичного текста письма пользователя и другой релевантной информации. Как показано на рисунке 32, сэмплирование реализуется путём генерации токенов по одному. Например, когда модели передаётся «Hi Alex, does today», токен «work» выбирается как следующий на основе предсказанных вероятностей. Затем на вход модели подаётся «Hi Alex, does today work», и токен «for» выбирается как следующий. Этот процесс продолжается до тех пор, пока модель не предскажет токен ⟨\langle⟨EOS⟩\rangle⟩.
Существуют два основных типа стратегий генерации нового текста в генеративных моделях: детерминированные и стохастические. Рассмотрим каждый из них.
Детерминированный
Детерминированные методы генерируют текст детерминированным образом, то есть без случайности или вариативности в выходных данных. Например, на каждом шаге генерации токенов модель выбирает токен с наибольшей вероятностью из предсказанного распределения. Этот метод гарантирует, что при одинаковом входе сгенерированный текст всегда будет одинаковым, обеспечивая согласованность и воспроизводимость. На рисунке 33 показан «жадный поиск» (greedy search) — простой детерминированный метод генерации текста путём итеративного выбора следующего токена на основе наивысшей предсказанной вероятности.
Плюсы:
- Согласованность: для одинакового ввода сгенерированный текст всегда одинаков — желательное свойство для систем, требующих предсказуемых результатов.
- Предсказуемые выходные данные: неожиданные выходные данные встречаются реже, поскольку на каждой итерации всегда выбирается наиболее вероятный токен.
Минусы:
- Отсутствие разнообразия: модель может пропустить менее вероятные, но более интересные токены, что ведёт к снижению креативности генерируемого текста. Например, при генерации истории модель всегда выбирает наиболее распространённые фразы, получая предсказуемое, но менее интересное повествование.
- Повторяющийся текст: текст может стать повторяющимся, поскольку один и тот же высоковероятный токен выбирается постоянно. Например, при генерации длинной статьи модель может многократно использовать определённые фразы. Реальный пример этого показан на рисунке 34.
Стохастическое сэмплирование
Методы стохастического сэмплирования вносят случайность в процесс генерации. Например, на каждом шаге генерации токенов модель выбирает образец из предсказанного распределения на основе вероятностей, присвоенных каждому токену. Это означает, что каждый раз при генерации текста, даже при одинаковых начальных входных данных, сгенерированный текст может отличаться.
На рисунке 35 показаны два примера сэмплирования с одинаковым начальным токеном «How». В первом случае последовательность сгенерированных токенов приводит к «How are you»; во втором — из того же начального токена генерируется иная последовательность вследствие случайности, присущей сэмплированию.
Плюсы:
- Разнообразие: наличие случайности обеспечивает более вариативные выходные данные, что особенно полезно в таких приложениях, как генерация диалогов.
- Новизна: выбирая образцы из распределения, модель может обращаться к менее вероятным, но потенциально более интересным токенам, что стимулирует креативность и получение оригинальных результатов.
Минусы:
- Непоследовательность: результат может различаться при каждой генерации текста, что менее пригодно для приложений, требующих точных и воспроизводимых результатов.
- Непредсказуемые выходные данные: случайность может приводить к неожиданным вариациям в сгенерированном тексте, которые могут оказаться неуместными.
Какой метод генерации подходит для функции Smart Compose?
Для Smart Compose предпочтительны детерминированные методы по ряду причин:
- Согласованность: согласованность генерируемого текста критически важна для таких приложений, как завершение электронных писем, где пользователи ожидают предсказуемых и надёжных подсказок. Использование детерминированного метода означает, что пользователи не будут видеть кардинально различающихся предложений каждый раз, когда начинают вводить одно и то же.
- Лучшая обработка распространённых фраз: детерминированные методы, как правило, предпочтительнее в контексте электронной почты, поскольку в них приоритет отдаётся наиболее вероятным вариантам завершения, а не новизне, которую предлагают стохастические методы.
- Снижение риска неуместных предложений: стохастические методы иногда могут генерировать неуместные предложения из-за присущей им случайности. Такое поведение нежелательно в функции автодополнения электронной почты.
Эти причины подчёркивают, почему детерминированные методы предпочтительны в приложениях, требующих согласованности, — например, в автодополнении электронной почты. Теперь, выбрав детерминированную генерацию текста, рассмотрим два основных алгоритма:
- Жадный поиск
- Лучевой поиск
Жадный поиск
Жадный поиск — простейший детерминированный алгоритм. Он всегда выбирает токен с наибольшей вероятностью в качестве следующего. Как показано на Рисунке 34, жадный поиск может приводить к повторяющимся паттернам в генерируемом тексте. Это происходит потому, что алгоритм следует узкому пути, основанному на токенах с наибольшей вероятностью, не рассматривая альтернативные пути, которые могли бы привести к более связным предложениям. Из-за этого ограничения жадный поиск редко используется на практике.
Лучевой поиск
Лучевой поиск [26] — популярный детерминированный алгоритм генерации текста из обученной модели. Основная идея состоит в одновременном отслеживании нескольких потенциальных последовательностей токенов. На каждом шаге модель вычисляет вероятности следующих возможных токенов для каждой последовательности и выбирает «top-k» наиболее вероятных последовательностей. Значение k, известное как ширина луча, является настраиваемым параметром.
Ниже приведён краткий пошаговый процесс генерации текста с использованием лучевого поиска при ширине луча, равной 3:
- Инициализация: начало работы с частично написанным письмом пользователя в качестве входных данных для обученной модели. Модель предсказывает распределение вероятностей для следующего токена. Лучевой поиск выбирает три токена с наибольшими вероятностями.
- Расширение: для каждой из трёх лучших последовательностей она передаётся в модель, которая возвращает вероятности следующего токена.
- Отсечение: выбираются три лучшие последовательности на основе их накопленных вероятностей.
Шаги расширения и отсечения повторяются до тех пор, пока все три потенциальных предложения не достигнут токена ⟨\langle⟨EOS⟩\rangle⟩ или максимальной длины. После остановки алгоритма лучевого поиска в качестве выходных данных выбирается последовательность с наибольшей накопленной вероятностью.
Лучевой поиск эффективен на практике, поскольку он отслеживает несколько потенциальных последовательностей одновременно, а не только наиболее вероятную. Однако у лучевого поиска есть два основных недостатка:
- Ограниченное разнообразие: лучевой поиск нередко приводит к похожим результатам, что не идеально для приложений, требующих разнообразных ответов.
- Сложности с длинными последовательностями: лучевой поиск плохо справляется с длинными последовательностями, поскольку одновременное отслеживание слишком большого числа последовательностей может стать вычислительно затратным.
Предложения функции Smart Compose, как правило, короткие, поэтому учёт дальних зависимостей менее критичен. Кроме того, разнообразие в вариантах дополнения электронных писем нежелательно. По этим причинам лучевой поиск выбран в качестве основного алгоритма выборки для генерации предложений.
Оценка
Оценка является неотъемлемой частью собеседований по проектированию ML-систем. Интервьюеры проверяют, умеют ли кандидаты эффективно тестировать и верифицировать проектируемую ML-систему. Идеальный ответ должен охватывать как онлайн-, так и офлайн-оценку и включать обсуждение популярных метрик для измерения производительности модели в каждой из этих настроек.
Рассмотрим некоторые распространённые метрики для оценки функции Smart Compose.
Метрики офлайн-оценки
Офлайн-оценка использует заранее собранные и исторические данные для оценки производительности модели. Её цель — убедиться в приемлемости производительности модели до её развёртывания в производственной среде. Например, рекомендательная система тестируется на исторических данных о взаимодействиях пользователей, чтобы проверить, насколько точно она предсказывает предпочтения пользователей. Аналогично, производительность обученной модели для функции Smart Compose оценивается на исторических данных электронной почты. Двумя широко используемыми метриками являются:
- Перплексия
- ExactMatch@N
Перплексия
Перплексия [27] — стандартная метрика, широко используемая при офлайн-оценке языковых моделей. Эта метрика измеряет, насколько точно модель предсказывает точную последовательность токенов, присутствующих в текстовых данных. В математическом выражении перплексия определяется как экспонента среднего «отрицательного логарифмического правдоподобия» предсказанной вероятности с учётом предшествующих токенов в последовательности:
В этом уравнении:
- XXX — токенизированная последовательность (x1,x2,⋯ ,xN)\left(x_1, x_2, \cdots, x_N\right)(x1,x2,⋯,xN) в текстовых данных, используемая для оценки точности предсказания последовательности моделью.
- NNN — количество токенов в последовательности.
- P(xi∣x1:i−1)P\left(x_i \mid x_{1: i-1}\right)P(xi∣x1:i−1) — условная вероятность i-го токена при наличии предшествующих токенов x1:i−1x_{1: i-1}x1:i−1, то есть вероятность того, что модель предскажет i-й токен, зная предыдущие токены.
Рисунок 38 иллюстрирует конкретный пример для лучшего понимания перплексии.
Более низкое значение перплексии означает, что модель в среднем присваивала более высокие вероятности токенам, встречающимся в текстовых данных. Таким образом, более низкая перплексия свидетельствует о том, что модель лучше предсказывает следующие токены.
ExactMatch@N
ExactMatch@N измеряет процент сгенерированных фраз длиной ровно N слов, которые совпадают с первыми N словами эталонного текста. На Рисунке 39 показаны вычисления ExactMatch@3 для трёх сгенерированных последовательностей. На практике для оценки обычно используется значительно больше трёх последовательностей.
Вычисление ExactMatch@N для различных значений N позволяет оценить производительность модели при разных длинах предложений. Для измерения общей производительности модели вычисляется ExactMatch для всех длин вплоть до определённой максимальной длины, после чего берётся среднее значение.
Хотя перплексия и ExactMatch@N традиционно использовались для оценки Gmail Smart Compose, другие метрики, такие как оценка BLEU и ROUGE-N, появившиеся позднее, также оказались полезными. Подробнее эти метрики рассматриваются в Главе 3.
Метрики онлайн-оценки
Онлайн-оценка измеряет производительность модели в реальном времени по мере взаимодействия пользователей с системой. Для оценки функции Smart Compose в онлайн-среде используются дополнительные метрики помимо перплексии и ExactMatch@N. Эти онлайн-метрики измеряют вовлечённость пользователей, задержку модели и общее влияние на пользовательский опыт.
В отличие от офлайн-метрик, которые, как правило, стандартизированы, метрики онлайн-оценки определяются исходя из конкретных требований и потребностей. Компании часто используют сотни метрик для онлайн-оценки. Однако в контексте собеседования обычно обсуждаются наиболее распространённые из них. В данном разделе мы сосредоточимся на следующих метриках:
- Метрики вовлечённости пользователей
- Метрики эффективности
- Метрики задержки
- Метрики качества
Метрики вовлечённости пользователей
- Процент принятия: доля предложений функции Smart Compose, принятых пользователями. Более высокий процент принятия свидетельствует о том, что предложения актуальны и полезны для пользователей.
- Процент использования: доля всех написанных писем, при создании которых была задействована функция Smart Compose. Высокий процент использования, как правило, указывает на то, что пользователи доверяют этой функции.
Метрики эффективности
- Среднее время написания: отслеживает среднее время, затрачиваемое пользователями на написание писем с использованием Smart Compose и без него. Сокращение среднего времени написания при использовании Smart Compose будет свидетельствовать о том, что функция ускоряет процесс написания электронной почты.
Метрики задержки
- Время отклика системы: измеряет время, необходимое для появления предложений Smart Compose после начала набора текста пользователем. Важно обеспечить, чтобы эта метрика оставалась ниже определённого порогового значения, — тогда предложения будут появляться прежде, чем пользователь успевает ввести соответствующий текст.
Метрики качества
- Частота обратной связи: измеряет, как часто пользователи оставляют отзывы о предложениях. Обратная связь полезна для непрерывного улучшения системы.
- Экспертная оценка: для оценки полезности предложений проводятся качественные исследования с участием пользователей. Эта метрика отражает удовлетворённость пользователей функцией Smart Compose.
Эти онлайн-метрики необходимы для оценки эффективности функции Smart Compose в производственной среде. Мониторинг этих метрик позволяет заинтересованным сторонам получить целостное представление о производительности функции.
Общий дизайн ML-системы
В данном разделе предлагается дизайн упрощённой версии функции Smart Compose.
При проектировании такой функции необходимо учитывать не только базовую модель, предсказывающую следующий токен. Эффективность системы зависит от слаженной работы различных компонентов, обеспечивающих оперативность отклика, генерацию актуальных предложений и соблюдение этических стандартов. Для функции Smart Compose рассматриваются следующие ключевые компоненты:
- Сервис триггера
- Генератор фраз
- Сервис постобработки
Рассмотрим каждый из них подробнее.
Сервис триггера
Сервис триггера активирует функцию Smart Compose, отслеживая действия пользователя, такие как нажатия клавиш. Он определяет момент активации функции на основе таких критериев, как количество введённых символов или ввод определённых ключевых слов в тексте. Например, если пользователь вводит «Я,» сервис может не активировать Smart Compose, поскольку контекста недостаточно для предсказания намерений пользователя. Однако если пользователь вводит «Я надеюсь,» сервис активирует Smart Compose, поскольку дополнительный контекст позволяет формировать более полезные предложения.
Сервис триггера обеспечивает, чтобы предложения не появлялись слишком часто. Как только сервис определяет, что активация функции Smart Compose будет полезной, он запускает компонент генератора фраз, который рассматривается далее.
Генератор фраз
Генератор фраз является основным компонентом функции Smart Compose. Он генерирует наиболее вероятное дополнение на основе частично введённого пользователем текста.
Для этого генератор фраз взаимодействует с обученной моделью и использует лучевой поиск для получения top-k наиболее вероятных вариантов дополнения. Каждый вариант дополнения заканчивается токеном ⟨\langle⟨EOS⟩\rangle⟩ и сопровождается оценкой уверенности, указывающей на степень достоверности предложения модели.
Исходя из возможных вариантов дополнения, необходимо учесть два ключевых момента:
- Удаление длинных предложений
- Удаление предложений с низкой уверенностью
Удаление длинных предложений
Поскольку короткие предложения легче читаются автором в процессе набора текста, предлагаемые фразы, которые слишком длинны, исключаются. Например, если пользователь вводит «Не могли бы вы,» генератор фраз может предложить «помочь мне с этим?» Более длинные предложения, такие как «помочь мне с этим проектом, срок сдачи которого — следующая неделя,» будут слишком конкретными и, следовательно, с меньшей вероятностью предугадают намерения автора.
Удаление предложений с низкой уверенностью
Предложения с оценкой уверенности ниже определённого порогового значения удаляются. Это гарантирует, что предложения не будут отображаться, если модель недостаточно уверена в их корректности.
Наконец, если итоговый список предложений не пуст, генератор фраз передаёт предложение с наибольшей оценкой уверенности в сервис постобработки.
Сервис постобработки
Сервис постобработки устраняет потенциальные предубеждения до того, как предложения будут представлены пользователю. Этот компонент следует заранее определённым правилам для эффективного обнаружения и исправления предубеждений. Среди распространённых стратегий для достижения этой цели можно выделить следующие:
- Замена местоимений: замена гендерно-специфичных местоимений для обеспечения нейтральности. Например, «он» или «она» может быть заменено на «они» в контекстах, где пол не указан.
- Замена гендерно-маркированных слов: замена гендерно-маркированных слов гендерно-нейтральными альтернативами там, где это уместно. Это включает замену таких слов, как «chairman» на «chairperson» или «policeman» на «police officer».
- Лексический анализ на предмет чувствительных терминов: использование заранее составленного списка помеченных терминов, которые в случае обнаружения могут быть заменены нейтральными альтернативами. Например, термины, способные указывать на предубеждения, связанные с возрастом, расой или инвалидностью, корректируются, чтобы предложения воспринимались как уважительные и нейтральные.
- Фильтрация контента NSFW (неприемлемого на рабочем месте): внедрение автоматических фильтров, сканирующих и помечающих нецензурную лексику. Эти фильтры используют заранее составленные списки ключевых слов, фраз и паттернов NSFW для обнаружения и удаления проблемного контента.
Применяя эти правила, сервис постобработки поддерживает этические стандарты в функции Smart Compose, тем самым обеспечивая актуальность, уважительность и инклюзивность предлагаемых вариантов.
Ниже приведён краткий пошаговый рабочий процесс общей ML-системы, используемой функцией Smart Compose:
- Мониторинг: сервис триггера отслеживает действия пользователя в процессе набора текста.
- Запуск: сервис активирует генератор фраз после обнаружения определённых паттернов.
- Лучевой поиск: генератор фраз использует лучевой поиск для получения top-k потенциальных вариантов дополнения из обученной модели.
- Фильтрация: генератор фраз взаимодействует с компонентом фильтрации для удаления длинных предложений и предложений с низкими оценками уверенности.
- Постобработка: выбирается вариант дополнения с наибольшей оценкой и передаётся в сервис постобработки. Сервис заменяет гендерно-специфичные местоимения и корректирует чувствительные термины.
- Отображение предложения: предложение отображается пользователю для его рассмотрения.
Дополнительные темы для обсуждения
Если в конце собеседования останется дополнительное время, вам могут задать уточняющие вопросы или предложить обсудить более сложные темы. Это зависит от таких факторов, как предпочтения интервьюера, ваш опыт и требования к должности. Для старших позиций рекомендуется подготовиться к следующим темам:
- Поддержка Smart Compose на нескольких языках [28].
- Персонализация предложений [28].
- Включение дополнительного контекста для улучшения предсказаний [28].
- Понимание принципов работы различных алгоритмов токенизации, таких как BPE [11], SentencePiece [12] и WordPiece [29].
- Понимание различных ML-целей, таких как маскированное языковое моделирование (MLM) и его вариации [18].
- Цель предсказания нескольких токенов и её преимущества и недостатки [30].
- Балансирование качества и задержки при инференсе [28].
Резюме
Справочные материалы
[1] Gmail's Smart Compose feature. https://research.google/pubs/gmail-smart-compose-real-time-assisted-writing/. [2] Fundamentals of Recurrent Neural Network. https://arxiv.org/abs/1808.03314. [3] Attention Is All You Need. https://arxiv.org/abs/1706.03762. [4] Gated recurrent unit. https://en.wikipedia.org/wiki/Gated_recurrent_unit. [5] Long Short-Term Memory. https://deeplearning.cs.cmu.edu/F23/document/readings/LSTM.pdf. [6] RITA: Group Attention is All You Need for Timeseries Analytics. https://arxiv.org/abs/2306.01926. [7] FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. https://arxiv.org/abs/2205.14135. [8] Language identification. https://en.wikipedia.org/wiki/Language_identification. [9] FastText model for language identification. https://huggingface.co/facebook/fasttext-language-identification. [10] Transformer-XL. https://arxiv.org/abs/1901.02860. [11] Byte-Pair Encoding tokenization. https://huggingface.co/learn/nlp-course/en/chapter6/5. [12] SentencePiece tokenization. https://arxiv.org/abs/1808.06226. [13] Tiktoken library. https://github.com/openai/tiktoken. [14] Google's Gemini. https://gemini.google.com/. [15] SentencePiece library. https://github.com/google/sentencepiece. [16] Summary of tokenizers. https://huggingface.co/docs/transformers/en/tokenizer_summary. [17] OpenAI's tokenizers. https://tiktokenizer.vercel.app/?model=gpt-4-1106-preview. [18] BERT. https://arxiv.org/abs/1810.04805. [19] OpenAI's models. https://platform.openai.com/docs/models. [20] Meta's LLaMA. https://llama.meta.com/. [21] Introduction to Transformers by Andrej Karpathy. https://www.youtube.com/watch?v=XfpMkf4rD6E. [22] Transformer visualized. https://jalammar.github.io/illustrated-transformer/. [23] Common Crawl. https://commoncrawl.org/. [24] Cross-entropy. https://en.wikipedia.org/wiki/Cross-entropy. [25] Prompt engineering. https://platform.openai.com/docs/guides/prompt-engineering. [26] Beam search. https://en.wikipedia.org/wiki/Beam_search. [27] Perplexity. https://en.wikipedia.org/wiki/Perplexity. [28] Gmail Smart Compose: Real-Time Assisted Writing. https://arxiv.org/abs/1906.00080. [29] WordPiece tokenization. https://huggingface.co/learn/nlp-course/en/chapter6/6. [30] Better & Faster Large Language Models via Multi-token Prediction. https://arxiv.org/abs/2404.19737.
Сноски
- Посетите https://platform.openai.com/tokenizer, чтобы ознакомиться с примерами работы различных токенизаторов. ↩
Google Переводчик
Введение
Google Переводчик — широко используемый сервис перевода, предоставляемый компанией Google. Сервис использует модели машинного обучения (ML) для понимания и перевода текстов между языками. По состоянию на 2024 год сервис поддерживает более 130 различных языков и насчитывает более миллиарда пользователей [1]. В этой главе рассматривается системный дизайн сервиса перевода.
Уточнение требований
Вот типичный диалог между кандидатом и интервьюером:
Candidate: Есть ли конкретные языки, которые система должна поддерживать изначально? Interviewer: Давайте сосредоточимся на четырёх языках: английском, испанском, корейском и французском. В будущем можно расширить поддержку до большего числа языков.
Candidate: Учитывая разнообразие языков, есть ли у нас доступ к достаточно большому и разнообразному набору данных для обучения? Interviewer: Да. У нас есть доступ к обширному многоязычному корпусу, включающему официальные документы, веб-контент и разговорные тексты на всех четырёх языках. Набор данных содержит 300 миллионов примеров, где каждый пример определяется как пара предложений на исходном и целевом языках.
Candidate: Есть ли у нас доступ к общим текстовым данным? Это важно, поскольку позволило бы нам выполнить pre-training модели на общих текстовых данных и тем самым дать ей возможность приобрести общие знания. Interviewer: Предположим, у нас есть доступ к терабайтам общих текстовых данных на каждом из языков, полученных из различных источников.
Candidate: Будут ли пользователи указывать язык вводимого текста, или система должна определять его автоматически? Interviewer: Пользователи не всегда могут определить язык текста. Представьте название книги на языке, с которым пользователь не знаком. Наша система должна автоматически определять язык ввода.
Candidate: Есть ли ограничение на длину вводимого текста? Interviewer: Давайте создадим систему, которая поддерживает входные данные объёмом до 1000 слов.
Candidate: Должна ли система поддерживать перевод без подключения к интернету? Иными словами, должна ли модель работать на устройстве? Interviewer: Фокус этого интервью не на эффективности и оптимизации модели для развёртывания на устройстве. Предположим, что требуется подключение к интернету и модель будет развёрнута в облаке.
Candidate: Должна ли система поддерживать перевод в реальном времени? Interviewer: Изначально нет.
Постановка задачи как ML-задачи
В этом разделе мы формулируем задачу создания системы перевода как ML-задачу. Это включает понимание входных и выходных данных системы, а также выбор подходящего ML-подхода.
Определение входных и выходных данных системы
Входными данными для системы перевода является последовательность слов на исходном языке и целевой язык, указанный пользователем. Выходными данными является последовательность слов на целевом языке.
Выбор подходящего ML-подхода
При языковом переводе последовательность слов на одном языке преобразуется в последовательность слов на другом языке. Такая структура «последовательность в последовательность» (seq2seq) встречается и в других задачах, например в суммаризации текста и распознавании речи.
Модели seq2seq — это класс ML-моделей, специально разработанных для решения подобных задач. Они преобразуют входную последовательность в выходную, которая может отличаться по длине от входной. Модели seq2seq следуют архитектуре encoder-decoder, имеющей два основных компонента:
- Encoder: обрабатывает входную последовательность и преобразует её в последовательность контекстных векторов, тем самым кодируя информацию из входной последовательности.
- Decoder: использует контекстные векторы encoder'а для пошагового генерирования выходной последовательности — по одному token'у за раз.
Для компонентов encoder и decoder существует несколько архитектур. В частности, для обработки последовательных данных можно использовать архитектуры LSTM, GRU и Transformer. Среди них Transformer'ы продемонстрировали превосходную производительность в задачах перевода, превзойдя более ранние модели, такие как GRU и LSTM, особенно в обработке долгосрочных зависимостей. Примечательно, что механизм attention был изначально предложен именно в контексте языкового перевода [2].
Как описано в Главе 2, Transformer'ы имеют три вариации: архитектура только encoder, только decoder и encoder-decoder. Модели только с encoder, такие как BERT [3], хорошо справляются с пониманием и обработкой входной последовательности, но, как правило, требуют дополнительных механизмов для генерации выходных данных. Модели только с decoder, такие как GPT от OpenAI [4] и Claude от Anthropic [5], высокоэффективны в генеративных задачах.
Хотя все три архитектуры демонстрируют высокую производительность и могут быть адаптированы для задач языкового перевода с помощью таких техник, как prompt engineering, модели encoder-decoder обычно предпочтительны по трём основным причинам. Во-первых, архитектура encoder-decoder разделяет понимание входных данных и генерацию выходных, что идеально для seq2seq-задач, таких как языковой перевод. Это позволяет encoder специализироваться на исходном языке и полностью понять входную последовательность, прежде чем decoder генерирует выходные данные. Например, encoder'ы часто используют двунаправленные механизмы, такие как двунаправленные LSTM [6] или Transformer'ы, которые позволяют им понимать контекст с обоих направлений.
Во-вторых, эта архитектура естественным образом обрабатывает последовательности переменной длины. Модели encoder-decoder разработаны для обработки входных/выходных последовательностей различной длины, что делает их очень универсальными для различных приложений. Такая гибкость критически важна для задач, в которых входные и выходные данные не имеют фиксированного соотношения длин.
Наконец, механизм cross-attention, присутствующий в Transformer'ах encoder-decoder, позволяет decoder'у динамически фокусироваться на релевантных частях входной последовательности в процессе генерации. Это целенаправленное внимание обеспечивает тесное соответствие выходной последовательности важным элементам исходной последовательности, тем самым повышая точность и качество перевода. Более подробно cross-attention будет рассмотрен в разделе об архитектуре.
Подготовка данных
В этом разделе мы подготавливаем сырые текстовые данные для Transformer'а encoder-decoder. У нас есть два типа данных для обучения: общие данные и данные для перевода. Общие данные включают общедоступные тексты из интернета. Данные для перевода включают 300 миллионов пар предложений, каждая из которых содержит предложение на исходном языке и соответствующий перевод на целевом языке.
Сырой текст как в общих данных, так и в данных для перевода часто зашумлён и не соответствует формату, ожидаемому ML-моделью. Поскольку подготовка общих данных была рассмотрена в Главе 2, мы сосредоточимся на подготовке данных для перевода. В частности, мы сосредоточимся на следующих двух шагах:
- Предобработка текста
- Tokenization текста
Предобработка текста
К сырому тексту в данных для перевода применяются следующие методы предобработки:
- Удаление отсутствующих данных: удаление пар, в которых отсутствует исходный или целевой текст.
- Удаление зашумлённых данных: удаление пар с HTML-тегами или неправильными языковыми парами.
- Дедупликация: удаление дублирующихся пар предложений из набора данных, чтобы предотвратить переобучение модели на определённых примерах.
- Обработка именованных сущностей: модели языкового перевода часто испытывают затруднения с именованными сущностями. Мы определяем эти сущности в тексте и заменяем их placeholder-token'ами. После перевода token'ы заменяются исходными сущностями. Например, рассмотрим предложение: «Калифорнийский город Бёрлингейм назван в честь дипломата Энсона Бёрлингейма.» Сначала мы определяем именованные сущности: «Калифорния (название местности)», «Бёрлингейм (название местности)» и «Энсон Бёрлингейм (имя человека)». Затем заменяем эти сущности placeholder-token'ами: «Город ENTITY_1, ENTITY_2, назван в честь дипломата ENTITY_3.» Такой подход помогает модели сосредоточиться на контексте предложения в процессе обучения, не отвлекаясь на редкие термины.
В современном языковом переводе, особенно с такими моделями, как Transformer'ы, некоторые традиционные шаги предобработки стали менее значимыми или выполняются иначе. Вот несколько шагов предобработки, которые были необходимы в традиционных моделях перевода, но теперь устарели или менее актуальны:
- Приведение к нижнему регистру: современные модели языкового перевода могут обрабатывать чувствительность к регистру как часть своего обучения. Они способны различать разные формы слов в зависимости от регистра (например, «Apple» как компания и «apple» как фрукт) без необходимости преобразовывать всё в нижний регистр. Поэтому приведение к нижнему регистру часто пропускается для сохранения исходной информации о регистре.
- Удаление стоп-слов: стоп-слова (например, «the», «and», «in») необходимы для грамматической структуры предложений. Их удаление может нарушить беглость и смысл переводов. Современные модели языкового перевода выигрывают от наличия полных предложений, включая стоп-слова, для полного понимания контекста и получения более естественных переводов.
- Стемминг и лемматизация: стемминг (сведение слов к базовой или корневой форме) и лемматизация (сведение слов к словарной форме) обычно не требуются в современном языковом переводе, поскольку эти модели спроектированы для обработки морфологических вариаций слов. Модели учатся переводить слова в их правильные формы на основе контекста; поэтому сведение к базовой форме могло бы фактически удалить ценную информацию.
- Удаление знаков препинания: знаки препинания важны для понимания структуры и смысла предложений. Современные модели языкового перевода обучены естественной обработке знаков препинания; поэтому их удаление может ухудшить качество перевода. Знаки препинания обычно сохраняются, чтобы помочь модели поддерживать грамматическую целостность предложений.
Tokenization текста
В контексте языкового перевода, в котором мы работаем с несколькими языками, выбор алгоритма tokenization текста имеет большое значение. Например, если бы мы выбрали токенизатор на уровне слов, наш словарь содержал бы сотни тысяч уникальных слов на всех языках, что было бы огромным и неэффективным.
При языковом переводе обработка разнообразия слов разных языков является ключевой задачей. Традиционные модели tokenization на уровне слов часто испытывают затруднения со словами, выходящими за пределы словаря (OOV), тогда как алгоритмы tokenization на уровне подслов более эффективны и могут эффективно решать проблему OOV. Учитывая их важность и широкое применение, полезно подробно рассмотреть Byte-Pair Encoding (BPE) [7] — широко используемый алгоритм tokenization на уровне подслов.
Byte-Pair Encoding (BPE)
BPE строит словарь на уровне подслов посредством итеративного слияния. Он начинает с отдельных символов и итеративно объединяет наиболее частотные комбинации в новые подслова. Это позволяет модели разбивать слова, в том числе редкие или незнакомые, на известные компоненты, обеспечивая тем самым точное понимание и перевод. Давайте разберём конкретный пример для лучшего понимания BPE.
Начальная настройка
Предположим, у нас есть корпус со следующим набором слов: «cat», «cats», «dog» и «dogs».
На начальном этапе наша цель — инициализировать словарь, состоящий из различных символов и их частоты встречаемости в корпусе. Для этого выполняются следующие шаги:
- Добавить специальный токен конца слова «</w>» в конец каждого слова для обозначения его границы. Этот специальный token помогает модели знать, когда слово закончилось.
- Выполнить tokenization корпуса, разбив каждое слово на отдельные символы.
- Инициализировать словарь отдельными символами и их частотой встречаемости.
Итеративное слияние
После создания начального словаря BPE итеративно объединяет наиболее частотные пары символов в подслова. Это продолжается до тех пор, пока словарь не достигнет заданного размера или не будут выполнены критерии остановки.
Ниже приведены первые пять итераций BPE:
- Итерация 1: сначала мы определяем наиболее частотную пару символов — «d» и «o», встречающиеся вместе 10 раз (в словах «dog» и «dogs»). Мы объединяем их, чтобы создать новый token «do». Хотя «og» также встречается 10 раз, «do» стоит раньше по алфавиту. Token «do» добавляется в словарь, а счётчики частот обновляются. «do» теперь встречается 10 раз, а индивидуальные счётчики «d» и «o» соответственно уменьшаются. Рисунок 7: Итерация BPE 1
- Итерация 2: теперь мы ищем следующую наиболее частотную пару — «do» и «g» (также из слов «dog» и «dogs»). Эти символы встречаются вместе 10 раз, поэтому мы объединяем их, создавая token «dog».
- Итерация 3: двигаясь дальше, мы замечаем, что «c» и «a» (из слов «cat» и «cats») встречаются вместе 8 раз. Мы объединяем их, создавая token «ca». После слияния «cat» может быть представлен как «ca» и «t». Token «ca» теперь имеет частоту 8.
- Итерация 4: продолжаем, объединяя «ca» и «t», которые встречаются вместе 8 раз (из слов «cat» и «cats»). Объединяем их, создавая token «cat». Теперь «cats» может быть представлен как «cat» и «s», а «dogs» как «dog» и «s». Счётчик частоты для «cat» обновляется до 8.
- Итерация 5: наконец, следующая наиболее частотная пара — «s» и «</w>» (из слов «dogs» и «cats»), встречающаяся 7 раз. Мы объединяем «s» и «</w>», создавая token «s</w>».
BPE итеративно объединяет наиболее частотные пары символов, приводя к более компактному представлению корпуса. Слияние продолжается до достижения желаемого количества token'ов или итераций.
Обратите внимание, что специальный token «</w>» играет ключевую роль в различении разных форм слов. Например, token «cat», за которым следует «</w>», указывает на конец слова «cat», тогда как token «cat» без «</w>» может быть частью другого слова. Это различие помогает BPE точно представлять и интерпретировать слова при переводе, позволяя эффективно обрабатывать как знакомые, так и незнакомые слова.
После создания словаря мы формируем обучающие данные, заменяя каждое токенизированное предложение последовательностью целых чисел. В результате получается несколько таблиц, каждая для конкретной языковой пары. На рисунке 11 показаны подготовленные данные для перевода для языковых пар английский–французский и английский–корейский.
Разработка модели
Мы использовали Transformer encoder-decoder для обучения языковому переводу. В этом разделе мы рассматриваем архитектуру encoder и decoder, стратегии обучения и методы сэмплирования.
Архитектура
Ключевые компоненты архитектуры Transformer encoder-decoder очень схожи с компонентами Transformer'а только с decoder, описанного в Главе 2. Рассмотрим encoder и decoder по отдельности и выделим их ключевые различия.
Encoder
Encoder обрабатывает входную последовательность и выдаёт последовательность embedding для каждого входного token'а.
Encoder состоит из следующих компонентов:
- Text embedding
- Positional encoding
- Transformer
Text embedding: этот компонент преобразует каждый входной token в вектор embedding. Эти embedding'и фиксируют семантическую информацию каждого token'а.
Positional encoding: компонент positional encoding вводит информацию о позиции каждого token'а во входной последовательности. Как обсуждалось в предыдущей главе, фиксированные и обучаемые методы эффективны на практике. Для простоты мы выбираем фиксированный метод positional encoding, например кодирование синус–косинус.
Transformer: Transformer обрабатывает последовательность token embedding'ов через стек блоков Transformer. Каждый блок содержит слой self-attention, использующий механизм multi-head attention (MHA) на входной последовательности, и feed-forward слой, с нормализующими слоями между ними для обеспечения стабильности в процессе обучения. Поскольку требования предполагают поддержку входной последовательности в 1000 слов, нам не нужно применять оптимизированные механизмы attention для эффективности, так как стандартный механизм attention достаточен для обработки последовательностей такой длины без существенных проблем с производительностью.
Decoder
Decoder генерирует выходную последовательность по одному token'у за раз, используя выходные данные encoder'а и ранее сгенерированные token'ы. Decoder состоит из следующих компонентов:
- Text embedding: преобразует каждый token целевой последовательности в embedding
- Positional encoding: вводит информацию о позиции каждого token'а
- Transformer: обрабатывает целевую последовательность и выдаёт обновлённую последовательность embedding'ов
- Prediction head: использует обновлённые embedding'и для предсказания следующего token'а.
Каковы ключевые различия между encoder и decoder?
Между encoder и decoder есть три ключевых различия:
- Слой cross-attention
- Механизм self-attention
- Prediction head
Слой cross-attention
Компонент Transformer в decoder'е включает слой cross-attention. Этот слой выполняет механизм MHA над выходными данными encoder'а. Он позволяет каждому token'у в decoder'е обращаться ко всем embedding'ам в encoder'е. Это позволяет cross-attention эффективно интегрировать информацию из входной последовательности в процессе генерации выходной последовательности.
Механизм self-attention
Слой self-attention работает по-разному в encoder и decoder. В encoder'е каждый token обращается ко всем остальным token'ам в последовательности. Это помогает encoder'у всесторонне понять всю последовательность. В decoder'е, напротив, каждый token ограничен обращением только к тем token'ам, которые стоят перед ним, путём маскировки будущих token'ов в последовательности. Это различие важно для задач генерации, поскольку модель должна использовать только ранее сгенерированные token'ы, а не будущие, для предсказания следующего token'а.
Prediction head
Decoder имеет prediction head поверх компонента Transformer. Prediction head обычно включает линейный слой, за которым следует слой softmax для преобразования выходных данных Transformer'а в вероятности по всему словарю. Эти вероятности используются для определения наиболее вероятного следующего token'а.
Обучение
Мы применяем двухэтапную стратегию для обучения модели языкового перевода:
- Неконтролируемый pre-training
- Контролируемый fine-tuning
. Неконтролируемый pre-training
На этом этапе мы обучаем базовую модель на большом корпусе общих данных. Это создаёт базовую модель, способную понимать язык, грамматику и контекст.
Давайте рассмотрим данные для pre-training, ML-цель и функцию потерь для этапа pre-training.
Данные для pre-training
Мы используем популярные наборы данных для pre-training, такие как C4 [8], Wikipedia [9] и StackExchange [10]. В отличие от Главы 2, где мы сосредоточились на pre-training языковой модели только для английского языка, для языкового перевода нам нужна базовая модель с общим пониманием нескольких языков. Поэтому мы не удаляем неанглийские текстовые данные из этих наборов данных. Вместо этого мы оставляем набор языков, которые ожидаем от модели переводить, и удаляем все текстовые данные, относящиеся к языкам за пределами этого набора.
ML-цель и функция потерь
В Главе 2 мы рассмотрели предсказание следующего token'а как основную ML-цель для генерации текста. Предсказание следующего token'а не является идеальным выбором при pre-training encoder-decoder, поскольку обучение является неконтролируемым. Если мы передадим всё предложение в encoder, он закодирует информацию таким образом, что decoder всегда сможет точно предсказать следующее слово, фактически «мошенничая». Вместо этого мы используем «маскированное языковое моделирование» — распространённую ML-цель для pre-training Transformer'а encoder-decoder. Рассмотрим её подробнее.
Маскированное языковое моделирование (MLM)
В MLM, также известном как предсказание замаскированных token'ов, часть входных token'ов маскируется, и модель обучается предсказывать эти замаскированные token'ы.
MLM позволяет encoder'у обработать входное предложение и закодировать его так, чтобы decoder мог предсказать замаскированные слова. Поскольку замаскированные слова никогда не видны в процессе кодирования, это предотвращает «мошенничество» модели.
Для измерения производительности модели при предсказании замаскированных token'ов мы используем cross-entropy loss. Эта широко используемая функция потерь измеряет расхождения между предсказанными вероятностями и эталонными token'ами, направляя тем самым процесс обучения. Вот пошаговое объяснение того, как вычисляется потеря с использованием цели MLM:
- Случайным образом выбрать подмножество token'ов во входной последовательности и заменить их токеном маски («[MASK]»). Например, входное предложение «Thank you for inviting me» может стать «Thank [MASK] for inviting [MASK]».
- Передать замаскированную последовательность в encoder, чтобы он мог понять контекст, несмотря на отсутствующие token'ы. Encoder выдаёт последовательность новых embedding'ов для каждого token'а.
- Передать в decoder ту же входную последовательность, но на этот раз ни один из token'ов не замаскирован, и последовательность сдвинута на одну позицию вправо путём вставки начального token'а («<BOS>»). Обратитесь к Главе 2 или [11], чтобы понять, почему мы сдвигаем входную последовательность в процессе обучения.
- Decoder предсказывает следующий token для каждой позиции в последовательности. Каждое предсказание использует все предыдущие входные token'ы и закодированную информацию от encoder'а.
- Вычислить cross-entropy loss по предсказанным вероятностям и эталонным значениям только для замаскированных token'ов.
Таким образом, мы в основном используем цель MLM при pre-training Transformer'ов encoder-decoder, поскольку она задействует и encoder, и decoder. Encoder развивает понимание языка, кодируя замаскированный входной текст. Decoder учится обрабатывать эту закодированную информацию и предсказывать замаскированные token'ы. Эта цель готовит как encoder, так и decoder к этапу контролируемого fine-tuning.
Прежде чем перейти к этапу контролируемого fine-tuning, отметим, что pre-training базовой модели требует значительных ресурсов и, следовательно, является дорогостоящим. На практике мы часто используем общедоступные модели encoder-decoder, такие как T5 от Google [12] или BART от Meta [13], которые прошли pre-training на обширных наборах данных. Такой подход значительно снижает затраты и ресурсы, необходимые для pre-training.
. Контролируемый fine-tuning
Контролируемый fine-tuning, второй этап нашего процесса обучения, адаптирует базовую модель к конкретной задаче языкового перевода. Это достигается путём fine-tuning базовой модели на данных для перевода. Для адаптации базовой модели к языковому переводу у нас есть два варианта:
- Двуязычный подход
- Многоязычный подход
Двуязычный подход
При этом подходе мы обучаем модели, специфичные для каждой языковой пары. Обучение языкоспецифичных моделей имеет ряд преимуществ. Во-первых, они улавливают уникальные лингвистические нюансы каждой языковой пары. Во-вторых, они обычно демонстрируют более высокую точность перевода благодаря своей специализированности. Наконец, улучшение производительности проще с языкоспецифичными моделями, поскольку мы можем легко изолировать и устранить конкретные проблемы, которые могут возникнуть для каждой языковой пары. Однако обучение, развёртывание и поддержка нескольких моделей требуют значительных ресурсов и затрат.
Многоязычный подход
Здесь единая модель обучается переводить между несколькими языками. Многоязычные модели проще, менее затратны и легче в развёртывании и обслуживании, чем двуязычные модели. Недавние исследования, такие как mT5 [14] и mBART [15], подчеркнули тенденцию к многоязычным моделям перевода, которые часто соответствуют или превосходят производительность двуязычных моделей.
В этой главе мы отдаём приоритет точности перевода над простотой и, следовательно, выбираем двуязычный подход.
Обучающие данные
На рисунке 20 показан пример подготовленных обучающих данных, где каждая таблица представляет языковую пару. В каждой таблице строка представляет один пример, содержащий последовательность идентификаторов token'ов для предложения на исходном языке и последовательность идентификаторов token'ов для перевода на целевом языке.
ML-цель и функция потерь
В то время как этап pre-training был неконтролируемым, этап fine-tuning является контролируемым. Encoder обрабатывает token'ы исходного предложения для каждого обучающего примера, а decoder генерирует token'ы целевого предложения. Поскольку decoder должен генерировать token'ы последовательно после обучения, мы используем предсказание следующего token'а в качестве нашей ML-цели. Мы используем cross-entropy в качестве функции потерь для измерения точности предсказанного следующего token'а.
На рисунке 21 показано вычисление потерь на этапе fine-tuning. Для простоты на нём визуализировано единственное предсказание. На практике, как мы видели в Главе 2, decoder предсказывает следующий token для всех позиций одновременно, и потери вычисляются для всех предсказаний.
Сэмплирование
В процессе сэмплирования обученная модель генерирует потенциальный перевод, предсказывая каждый последующий token на основе ранее сгенерированных token'ов и контекста входной последовательности.
Как обсуждалось в Главе 2, существуют две основные стратегии сэмплирования текста в генеративных моделях: детерминированные методы (например, beam search) и стохастическое сэмплирование. Здесь мы выбираем beam search по двум основным причинам:
- Точность перевода: beam search обычно приводит к более точным переводам. Это объясняется тем, что алгоритм оценивает несколько возможных последовательностей и выбирает наиболее вероятную.
- Согласованность: beam search детерминирован, то есть всегда выдаёт один и тот же результат при одинаковых входных данных. Эта согласованность гарантирует, что переводы будут давать мало неожиданностей, что критически важно в большинстве систем перевода. Хотя разнообразие может быть полезным, оно ни необходимо, ни желательно для систем языкового перевода.
Обратите внимание, что в приложениях, где разнообразие и творчество ценятся выше, например в творческом письме, обычно предпочтительны стохастические методы сэмплирования. В Главе 4 мы подробно рассмотрим стохастические методы, такие как top-k и top-p сэмплирование.
| Characteristic | Deterministic methods | Stochastic methods |
|---|---|---|
| Approach | Follow a predictable process to generate output | Generate output based on probability distribution |
| Efficiency | Typically less efficient due to tracking multiple paths | More efficient since randomness allows for quicker selections |
| Quality | Coherent and predictable | Diverse and creative |
| Risk | Usually lead to repetitive output for longer sequences | Might produce inappropriate output due to their creativeness |
| Use case | Suitable for tasks requiring consistency, such as language translation | Suitable for tasks requiring creativity, such as open-ended text generation |
| Methods | Greedy search, beam search | Multinomial, top-k, top-p |
Таблица 1: Сравнение детерминированных и стохастических методов
Оценка
Метрики оффлайн-оценки
Для всестороннего оценивания модели языкового перевода метрики должны измерять как точность перевода, так и контекстную уместность. Исследовательское сообщество предложило ряд метрик, которые со временем получили широкое признание в качестве стандартов. Среди наиболее распространённых метрик:
- BLEU
- ROUGE
- METEOR
BLEU
BLEU (BiLingual Evaluation Understudy) [16] — метрика, основанная на точности, которая сравнивает n-граммы (последовательность «n» слов) кандидата на перевод с n-граммами эталонных переводов и подсчитывает долю совпадений. Значение варьируется от 0 до 1, где более высокое значение указывает на более точный перевод.
Оценка BLEU вычисляется по следующей формуле:
где:
- N — максимальная длина n-граммы, учитываемая при оценке
- BP — штраф за краткость (brevity penalty)
- pn — точность n-граммы
- wn — вес для различных точностей n-граммы
Рассмотрим каждый из этих терминов подробнее.
Штраф за краткость (BP) BP — константный член, который штрафует переводы, более короткие, чем эталонный перевод. Формула:
где:
- c — длина перевода
- r — длина эталонного перевода
Если длина перевода кандидата c больше длины эталонного перевода r, штраф за краткость равен 1 (т. е. штрафа нет). Если длина перевода кандидата меньше или равна длине эталонного перевода, штраф за краткость представляет собой экспоненциальный спад, основанный на соотношении длин.
Точность (pn): Точность измеряет, сколько n-граммов перевода кандидата присутствует в эталонных переводах. Она вычисляется путём деления числа совпадающих n-граммов на общее число n-граммов в переводе кандидата. На рисунке 23 приведён пример вычисления p2 для предложения кандидата и одного эталонного предложения.
Веса (wn) Эти веса соответствуют точности каждого размера n-граммы. Обычно они распределяются равномерно, придавая одинаковую важность каждой точности n-граммы. Например, для n-граммов до 4-грамма каждый wn равен 1/4.
Главное преимущество BLEU — простота и лёгкость вычисления. Однако у него есть существенный недостаток: он может несправедливо штрафовать переводы, которые правильны, но отличаются от эталонного перевода. Например, если эталонный перевод — «The engineer discovered a new algorithm», а сгенерированный перевод — «The engineer found a new method», BLEU может оштрафовать сгенерированный перевод, несмотря на то что он передаёт тот же смысл. Несмотря на это ограничение, BLEU остаётся информативным и широко используемым на практике для оценки моделей языкового перевода.
ROUGE
ROUGE (Recall-Oriented Understudy for Gisting Evaluation) [17] — популярная метрика, которая дополняет BLEU, фокусируясь на полноте вместо точности. Она измеряет долю перекрывающихся n-граммов между текстами кандидата и эталона. Например, полнота ROUGE-N определяется следующим образом:
Если вы хотите узнать больше о ROUGE и его формуле, обратитесь к [17].
Как и BLEU, ROUGE прост в реализации и эффективен при вычислении. Однако его главный недостаток — отсутствие контекстного понимания. Перевод с разными, но семантически схожими словами может получить низкую оценку ROUGE.
METEOR
METEOR (Metric for Evaluation of Translation with Explicit ORdering) [18] — популярная метрика для оценки моделей языкового перевода. Она вычисляет точность и полноту, а затем объединяет эти измерения с помощью взвешенного гармонического среднего.
В отличие от BLEU и ROUGE, которые опираются на точные совпадения n-граммов, METEOR учитывает синонимы и морфологию слов. Например, если эталонный перевод использует «run», а сгенерированный перевод — «running», METEOR распознаёт их как связанные термины. Синонимы находятся с помощью лингвистических ресурсов, таких как словари синонимов или лексические базы данных. Одним из широко используемых ресурсов является WordNet [19], который организует слова в синонимические множества различных типов и показывает отношения между этими синсетами.
Хотя METEOR является более всесторонней метрикой, у неё есть ряд недостатков. Рассмотрим её преимущества и недостатки.
Преимущества:
- Семантическое понимание: METEOR более точно оценивает качество перевода, когда разные формулировки передают одинаковый смысл. Это объясняется тем, что при оценке переводов учитываются синонимы и стемминг.
- Сбалансированная оценка: METEOR обеспечивает сбалансированную оценку, поскольку объединяет точность и полноту. Это помогает выявлять переводы, которые являются как точными, так и полными.
- Корреляция с человеческими оценками: METEOR лучше коррелирует с человеческими оценками, чем BLEU и ROUGE.
Недостатки:
- Вычислительная сложность: METEOR сложнее в реализации и требует больше времени для вычисления, чем BLEU и ROUGE. Это объясняется тем, что он требует дополнительных шагов, таких как сопоставление синонимов и стемминг.
- Зависимость от ресурсов: METEOR опирается на лингвистические ресурсы, такие как словари синонимов и алгоритмы стемминга, которые могут быть недоступны для всех языков.
Подводя итог, все три метрики дают представление о производительности модели и широко используются на практике. Перейдём к онлайн-оценке, чтобы понять, как наша модель работает в реальных сценариях.
Метрики онлайн-оценки
В процессе онлайн-оценки мы оцениваем, насколько хорошо наша система языкового перевода работает в production. Мы используем следующие две метрики для измерения удовлетворённости и вовлечённости пользователей:
- Обратная связь от пользователей: сбор оценок или отзывов пользователей о качестве переводов. Эта метрика информативна, поскольку напрямую отражает удовлетворённость пользователей. Рисунок 25: Сбор обратной связи от пользователей
- Вовлечённость пользователей: измерение вовлечённости пользователей путём мониторинга того, как часто они используют функцию перевода, как долго взаимодействуют с ней и как часто возвращаются. Это помогает нам понять, насколько ценным и эффективным является инструмент перевода в реальном использовании.
Сочетание метрик оффлайн- и онлайн-оценки даёт нам более полное представление о производительности языкового перевода. Такая всесторонняя оценка гарантирует, что модели соответствуют техническим стандартам и удовлетворяют ожиданиям пользователей.
Общий дизайн ML-системы
В этом разделе мы рассматриваем ML-дизайн системы языкового перевода. В частности, рассмотрим два ключевых компонента:
- Детектор языка
- Сервис перевода
Детектор языка
Детектор языка определяет язык данного текста, позволяя нам использовать модель, специально обученную для этого языка. Эту задачу можно сформулировать как задачу классификации последовательностей, и архитектура только с encoder является хорошим кандидатом для такой задачи. Мы можем модифицировать Transformer только с encoder двумя способами (рисунок 27) для классификации входных предложений:
- Average pooling: передать выходные данные Transformer'а в слой average pooling, а затем в prediction head для получения вероятностей языкового класса.
- Представление последнего token'а: использовать представление последнего token'а из выходных данных Transformer'а и передать его в prediction head для предсказания вероятностей.
Сервис перевода
Сервис перевода взаимодействует с конкретной моделью на основе определённого и желаемого языков. Затем он применяет beam search для генерации последовательности token'ов на целевом языке и преобразует token'ы обратно в текст. Итоговый перевод затем отображается пользователю.
Дополнительные темы для обсуждения
Если в конце интервью останется время, рассмотрите возможность обсуждения следующих дополнительных тем:
- Поддержка перевода для языков с ограниченными обучающими данными с использованием transfer learning и многоязычных моделей [20].
- Подход к языковому переводу с использованием Transformer'а только с decoder [21].
- Непрерывное улучшение моделей перевода на основе обратной связи от пользователей [22].
- Методы оптимизации для эффективного инференса и перевода на устройстве [23].
- Разработка единой многоязычной модели [24].
- Другие автоматические метрики, такие как WER, и способы их вычисления [25][26].
- Как построить модель определения языка [27].
Резюме
Справочные материалы
[1] Google Translate service. https://blog.google/products/translate/google-translate-new-languages-2024/. [2] Neural Machine Translation by Jointly Learning to Align and Translate. https://arxiv.org/abs/1409.0473. [3] BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. https://arxiv.org/abs/1810.04805. [4] GPT models. https://platform.openai.com/docs/models. [5] Claude models. https://www.anthropic.com/claude. [6] Bidirectional Long Short-Term Memory (BLSTM) neural networks for reconstruction of top-quark pair decay kinematics. https://arxiv.org/abs/1909.01144. [7] BPE tokenization. https://huggingface.co/learn/nlp-course/en/chapter6/5. [8] C4 dataset. https://www.tensorflow.org/datasets/catalog/c4. [9] Wikipedia dataset. https://www.tensorflow.org/datasets/catalog/wikipedia. [10] Stack Exchange dataset. https://huggingface.co/datasets/HuggingFaceH4/stack-exchange-preferences. [11] How Transformers work. https://huggingface.co/learn/nlp-course/en/chapter1/4. [12] Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer. https://arxiv.org/pdf/1910.10683.pdf. [13] BART: Denoising Sequence-to-Sequence Pre-training for Natural Language Generation, Translation, and Comprehension. https://arxiv.org/abs/1910.13461. [14] mT5: A massively multilingual pre-trained text-to-text transformer. https://arxiv.org/abs/2010.11934. [15] Multilingual denoising pre-training for neural machine translation. https://arxiv.org/abs/2001.08210. [16] BLEU metric. https://en.wikipedia.org/wiki/BLEU. [17] ROUGE metric. https://en.wikipedia.org/wiki/ROUGE_(metric). [18] METEOR metric. https://www.cs.cmu.edu/~alavie/METEOR/pdf/Banerjee-Lavie-2005-METEOR.pdf. [19] WordNet. https://wordnet.princeton.edu/. [20] No Language Left Behind: Scaling Human-Centered Machine Translation. https://research.facebook.com/publications/no-language-left-behind/. [21] Decoder-Only or Encoder-Decoder? Interpreting Language Model as a Regularized Encoder-Decoder. https://arxiv.org/abs/2304.04052. [22] Towards Continual Learning for Multilingual Machine Translation via Vocabulary Substitution. https://arxiv.org/abs/2103.06799. [23] Efficient Inference For Neural Machine Translation. https://arxiv.org/abs/2010.02416. [24] Meta's multilingual model. https://ai.meta.com/blog/nllb-200-high-quality-machine-translation/. [25] Machine translation evaluation. https://en.wikipedia.org/wiki/Evaluation_of_machine_translation. [26] Word error rate (WER) metric. https://en.wikipedia.org/wiki/Word_error_rate. [27] Automatic Language Identification using Deep Neural Networks. https://research.google.com/pubs/archive/42538.pdf.
ChatGPT: персональный ассистент-чатбот
Введение
ChatGPT [1] — это чатбот, разработанный компанией OpenAI и запущенный в 2022 году. Он генерирует текст, похожий на человеческий, на основе получаемого ввода. Чатбот может помогать с различными задачами, включая ответы на вопросы, предоставление объяснений и создание творческого контента.
ChatGPT быстро стал одним из самых быстро набравших популярность приложений в истории. Он привлёк более 100 миллионов пользователей менее чем за три месяца после запуска [2]. Этот стремительный рост подчёркивает возможности генеративного ИИ и его потенциал для помощи в повседневных задачах и повышения продуктивности. В этой главе мы рассмотрим ключевые компоненты создания чатбота, аналогичного ChatGPT.
Уточнение требований
Вот типичное взаимодействие между кандидатом и интервьюером:
Кандидат: Какие языки должен поддерживать чатбот? Интервьюер: Для начала давайте сосредоточимся на английском.
Кандидат: Нам необходимо убедиться, что чатбот генерирует непредвзятые и безопасные ответы, используя строгую модерацию контента и соответствующие алгоритмы. Это корректное предположение? Интервьюер: Безусловно.
Кандидат: Можете ли вы уточнить спектр задач, которые должен выполнять чатбот? Интервьюер: Чатбот должен уметь справляться с такими задачами, как предоставление информации и ответы на вопросы.
Кандидат: Принимает ли чатбот на вход или выдаёт на выход нетекстовые модальности, такие как изображения, аудио или видео? Интервьюер: Давайте пока сосредоточимся на текстовом чатботе. И ввод, и вывод — это текст.
Кандидат: Должен ли чатбот обрабатывать уточняющие вопросы? Как долго чатбот должен сохранять контекст диалога? Интервьюер: Хороший вопрос. Чатбот должен уметь обрабатывать уточняющие вопросы в рамках одной сессии диалога. Допустим, ожидается, что контекстное окно будет не менее 4096 токенов.
Кандидат: Должен ли чатбот уметь просматривать веб-сайты, вызывать внешние API или искать информацию в интернете? Интервьюер: Давайте не будем на этом фокусироваться в этом раунде.
Кандидат: Должен ли чатбот персонализировать взаимодействие с пользователями? Интервьюер: Давайте не будем сосредотачиваться на персонализации.
Кандидат: Есть ли у нас обучающие данные на основе инструкций? Интервьюер: Да, у нас есть набор данных с 80 000 примеров инструкций и ответов.
Формулировка задачи как задачи ML
Определение входа и выхода системы
Вход чатбота — это текстовый запрос (промпт), предоставленный пользователем. Промпт может быть вопросом, командой или любой другой формой текстового запроса. Выход — это релевантный и контекстуально уместный ответ, сгенерированный чатботом.
Выбор подходящего подхода ML
Разработка чатбота — это задача генерации текста, в которой языковая модель обрабатывает входной промпт и генерирует ответ. Такая языковая модель обычно требует миллиардов параметров для эффективного обучения; поэтому их часто называют большими языковыми моделями (LLM).
Как мы видели в предыдущих главах, decoder-only Transformer является стандартным архитектурным выбором для языковых моделей. Большинство современных LLM, таких как GPT от OpenAI [3], Gemini от Google [4] и Llama от Meta [5], основаны на архитектуре decoder-only Transformer. В соответствии с этими моделями мы используем decoder-only
Подготовка данных
Эффективность LLM зависит от качества обучающих данных, которые в основном поступают из веб-источников. Эти данные, часто автоматически собранные с веб-сайтов, форумов и блогов, требуют особого внимания и тщательной подготовки. Наиболее распространённые этапы включают:
- Извлечение и парсинг контента: Данные, собранные из веба, часто содержат посторонние элементы, например, HTML-теги, рекламу и навигационные ссылки. Этот этап включает парсинг сырого HTML-контента с использованием библиотек, таких как Beautiful Soup [6] или lxml [7], и извлечение основного текста с отбрасыванием нерелевантных разделов. Для выделения и сохранения ключевого контента, значимого для языкового моделирования, применяются такие техники, как анализ DOM [8] и обнаружение шаблонного контента [9].
- Фильтрация по URL и доменам: Не все веб-домены предоставляют качественный или релевантный контент. Фильтрация URL использует предопределённые правила или классификаторы машинного обучения (ML) для исключения нежелательных источников, например, низкокачественных блогов, контентных ферм или спам-сайтов. Также применяются техники белых и чёрных списков доменов для курирования данных из надёжных и релевантных источников, обеспечивая качество и надёжность набора данных.
- Определение языка: Собранные данные часто включают многоязычный контент, который необходимо отфильтровать в соответствии с целевым языком (языками) для обучения. Для классификации и фильтрации документов используются инструменты определения языка, такие как fastText [10] или langid.py [11].
- Фильтрация по качеству контента: Не весь веб-контент одинаково ценен для обучения. Для оценки и фильтрации низкокачественного текста используются техники оценки качества, включая оценку читаемости, алгоритмы обнаружения спама и эвристические проверки (например, длина контента, структура предложений). ML-модели также могут применяться для прогнозирования качества веб-контента на основе извлечённых из текста признаков. Этот этап критически важен для обеспечения использования только высококачественных данных для обучения.
Наряду с этими техниками, следующие методы обычно применяются как к данным, собранным из веба, так и к другим источникам данных, таким как книги, статьи и посты в социальных сетях:
- Удаление неприемлемого контента: Мы используем ML-модели для удаления оскорбительного, вредоносного или NSFW-контента из обучающих данных. Это гарантирует, что наша модель обучается только на уместном и безопасном контенте.
- Анонимизация конфиденциальной информации: Мы анонимизируем любую персонально идентифицируемую информацию (PII) в наборе данных. Этот этап критически важен для соблюдения законов о конфиденциальности и этических норм.
- Удаление низкокачественных данных: Мы используем ML-модели для удаления низкокачественных текстовых данных. Модели оценивают связность, релевантность, грамматику и читаемость. Этот этап обеспечивает высокое качество обучающих данных и их пригодность для обучения.
- Удаление дублированных данных (дедупликация): Мы удаляем похожие тексты, которые могут присутствовать в различных источниках. Например, если одна и та же новостная статья собрана с нескольких веб-сайтов, мы идентифицируем и сохраняем только одну копию. Этот этап уменьшает избыточность в обучающем наборе данных и гарантирует, что модель не подвергается чрезмерному воздействию определённых данных.
- Удаление нерелевантных данных: Мы используем эвристики и методы на основе правил для удаления нерелевантных данных. Например, мы удаляем тексты с нестандартными символами или на языках, которые чатбот не должен поддерживать.
- Токенизация текста: Мы используем алгоритм подсловной токенизации, такой как Byte-Pair Encoding (BPE), для токенизации текстовых данных. Для обзора BPE обратитесь к Главе 3.
Разработка модели
Архитектура
Архитектура LLM основана на decoder-only Transformer. Хотя текстовый embedding, блоки Transformer и голова предсказания аналогичны decoder-only Transformer, рассмотренному в Главе 2, LLM обычно используют более продвинутые методы позиционного кодирования.
Давайте подробнее рассмотрим позиционное кодирование LLM.
Позиционное кодирование
В контексте чатбота входная последовательность обычно значительно длиннее одного предложения или письма. Согласно требованиям интервьюера, наша цель — построить систему с контекстным окном не менее 4096 токенов. Это требует метода позиционного кодирования, который позволит модели понимать позиции всех токенов и взаимосвязи между ними.
В этом разделе мы начнём с краткого обзора абсолютного позиционного кодирования, а затем рассмотрим относительное позиционное кодирование. Наконец, мы углубимся в rotary positional embedding (RoPE) [12], надёжный метод позиционного кодирования, используемый популярными LLM, такими как Llama 3 [13].
Абсолютное позиционное кодирование
Абсолютное позиционное кодирование относится к традиционным методам, таким как синусоидальные или обучаемые кодировки, при которых каждая позиция в последовательности представлена уникальным вектором.
В этом подходе закодированные позиции затем добавляются к embedding токенов, предоставляя модели информацию о том, где каждый токен появляется в последовательности. Формально ключи и запросы attention вычисляются с помощью следующих уравнений:
Где:
- qmq_mqm — это вектор запроса на позиции mmm,
- knk_nkn — это вектор ключа на позиции nnn,
- WqW_qWq и WkW_kWk — обучаемые весовые матрицы,
- eme_mem и ene_nen — embedding токенов на позициях mmm и nnn,
- pmp_mpm и pnp_npn — позиционные векторы (обучаемые или фиксированные) на позициях mmm и nnn.
Оценка attention вычисляется как скалярное произведение векторов запроса и ключа:
Обратите внимание, что позиционные кодировки pmp_mpm и pnp_npn зависят только от абсолютных позиций. Следовательно, этот подход захватывает только информацию об абсолютной позиции, ограничивая способность модели улавливать относительные расстояния между токенами и обобщать на последовательности различной длины или с невиданными ранее позициями токенов. Например, модель, обученная на последовательностях длиной до 512 токенов, может испытывать трудности при применении к последовательностям из 4096 токенов. Синусоидальные паттерны имеют тенденцию становиться повторяющимися на больших расстояниях, что приводит к потере информации о взаимосвязях токенов. Этот недостаток устраняется относительным позиционным кодированием.
Относительное позиционное кодирование
В относительном позиционном кодировании вместо кодирования абсолютных позиций токенов мы кодируем разницу в позициях двух токенов. Таким образом, модель может сосредоточиться на относительных расстояниях между токенами, что зачастую важнее их абсолютных позиций. Например, в предложении знание того, что слово «car» следует за словом «chased», информативнее, чем знание позиций токенов как чисел 5 и 10.
Вычисление attention в относительном позиционном кодировании может быть выражено различными способами. В статье T5 [14] предлагается, что второй и третий члены взаимодействия в исходном выражении абсолютного позиционного кодирования можно отбросить, а четвёртый член заменить обучаемым смещением:
Напротив, в статье DeBERTa [15] отбрасывается последний член и заменяются второй и третий члены, состоящие из абсолютных позиционных векторов pmp_mpm и pnp_npn соответственно, на вектор относительной позиции Rn−mR_{n-m}Rn−m:
Относительное позиционное кодирование позволяет модели понимать взаимосвязи между токенами независимо от их абсолютных позиций. Однако оно вносит дополнительную сложность, поскольку qm⋅knq_m \cdot k_nqm⋅kn в механизме attention больше не может быть сведено к простому скалярному произведению. Это ограничивает нашу способность использовать эффективные техники, такие как linear attention [16]. RoPE, который мы рассмотрим далее, решает это ограничение путём кодирования как абсолютной, так и относительной позиционной информации через вращение в пространстве embedding.
Вращательное позиционное кодирование (RoPE)
RoPE представляет позиционную информацию как матрицу вращения, применяемую к embedding токенов. Это можно описать математически следующим образом: Для заданной входной последовательности RoPE применяет матрицу вращения к каждому embedding. Это преобразование может быть выражено как:
где qmq_mqm — embedding токена на позиции mmm, а R(θm)R\left(\theta_m\right)R(θm) — матрица вращения, параметризованная позиционным углом θm\theta_mθm. Этот угол обычно выводится из индекса позиции mmm и конструируется таким образом, чтобы вращение захватывало как абсолютную, так и относительную позиционную информацию. Эта матрица вращения, построенная с использованием тригонометрических функций, вращает embedding в комплексной плоскости, захватывая как абсолютную, так и относительную позиционную информацию.
Рисунок 4 показывает, как вращательное позиционное кодирование работает путём вращения embedding слов в двумерном пространстве. Слова «cat» и «dog» представлены как векторы, и угол между ними, обозначенный θ\thetaθ, кодирует их позиционное соотношение. Слева показано предложение «The cat chased the dog.» Позиция «cat» показана красным цветом, а позиция «dog» — синим. Угол между этими векторами отражает относительное расположение этих двух слов в предложении.
Справа показано другое предложение «Once upon a time, the cat chased the dog». Обратите внимание, что относительный угол между векторами «cat» и «dog» остаётся тем же, но их абсолютные позиции отличаются. Это демонстрирует, как RoPE захватывает как абсолютные, так и относительные позиции слов, позволяя модели понимать порядок и расстояние между словами в предложении.
В пространстве более высокой размерности матрица вращения может быть расширена для ddd измерений:
Преимущества:
- Трансляционная инвариантность: RoPE кодирует позиционную информацию таким образом, что она остаётся согласованной даже при смещении позиций токенов. Это помогает модели лучше справляться с изменениями позиций, чем другие методы.
- Представление относительной позиции: Вращения RoPE кодируют позиционную информацию геометрически в пространстве embedding. Это позволяет модели по своей природе понимать относительные расстояния между токенами, в отличие от традиционных синусоидальных кодировок, которые кодируют позицию аддитивно без использования этого геометрического понимания.
- Обобщение на невиданные позиции: Поскольку RoPE кодирует позицию через вращения, результирующие embedding сохраняют согласованные взаимосвязи независимо от абсолютной позиции. Это обеспечивает лучшее обобщение для последовательностей различной длины.
Недостатки:
- Математическая сложность: RoPE вводит дополнительные математические операции, включающие вращения в пространстве embedding. Хотя они не чрезмерно сложны, они более замысловаты, чем традиционные методы позиционного кодирования, такие как синусоидальные или обучаемые позиционные embedding.
Обучение
В предыдущих главах мы рассмотрели двухэтапную стратегию обучения языковых моделей. Однако этих этапов недостаточно при обучении продвинутых чатботов. Большинство чатботов, включая ChatGPT, используют трёхэтапную стратегию обучения:
- Pre-training
- Supervised finetuning (SFT)
- Reinforcement learning from human feedback (RLHF)
Давайте обсудим каждый этап подробнее, чтобы понять его назначение.
. Pre-training
Pre-training — это начальный этап процесса обучения. На этом этапе модель обучается на огромном объёме текстовых данных из интернета. Цель pre-training — создать базовую модель с широким пониманием языка и знаниями о мире.
Этап pre-training требует значительных вычислительных ресурсов. Обычно он требует тысяч GPU, стоит миллионы долларов и занимает месяцы обучения.
Данные для pre-training
Данные для pre-training обычно состоят из большого корпуса общих текстовых данных из различных источников в интернете, например, веб-страниц, книг и постов в социальных сетях.
При pre-training LLM обычно используется несколько наборов данных. Каждый из них служит уникальной цели — от расширения знакомства модели с различными стилями языка до углубления её понимания конкретных областей. Часто используемые наборы данных:
- Common Crawl: Common Crawl [17] — это общедоступный набор данных, собранный с большого количества веб-страниц в интернете. Он содержит петабайты данных, регулярно собираемых с 2008 года. Эти данные часто содержат нерелевантную информацию и вредоносный контент; следовательно, для подготовки к обучению LLM требуется значительная очистка данных.
- C4: C4 [18], созданный Google, — это очищенная версия набора данных Common Crawl, специально предназначенная для обучения LLM.
- GitHub: Набор данных GitHub включает обширную коллекцию репозиториев с открытым исходным кодом. Его назначение — помочь модели понять языки программирования и структуры кода.
- Wikipedia: Набор данных Wikipedia включает широкий спектр фактической информации, извлечённой из Wikipedia. Этот набор данных обычно является более надёжным источником, так как он написан и отредактирован более тщательно.
- Books: Набор данных Books — это коллекция книг различных жанров. Книги содержат длинный текстовый контент и обладают хорошим качеством данных, что способствует улучшению производительности LLM.
- ArXiv: Набор данных ArXiv содержит академические материалы и опубликованные статьи. Этот набор данных помогает модели понять терминологию и знания в академической области.
- Stack Exchange: Stack Exchange [19] — это веб-сайт с высококачественными вопросами и ответами, преимущественно в формате диалога между пользователями. Большинство популярных LLM обучаются на всех или некоторых из перечисленных наборов данных. Например, модель Llama-1 от Meta использует все вышеуказанные наборы данных, содержащие около 1,4 триллиона токенов. Таблица 1 показывает долю каждого набора данных, использованного Llama-1 во время обучения.
| Dataset | Sampling proportion | Disk size |
|---|---|---|
| Common Crawl | 67.0% | 3.3 TB |
| C4 | 15.0% | 783 GB |
| Github | 4.5% | 328 GB |
| Books | 4.5% | 85 GB |
| Wikipedia | 4.5% | 83 GB |
| ArXiv | 2.5% | 92 GB |
| Stack Exchange | 2.0% | 78 GB |
Таблица 1: Набор данных для pre-training Llama 1
Целевая функция ML и функция потерь
Поскольку мы обучаем decoder-only Transformer для генерации текста, мы используем стандартное предсказание следующего токена в качестве целевой функции ML. В качестве функции потерь мы применяем кросс-энтропию для измерения разницы между предсказанными вероятностями токенов и правильными токенами.
Результат этапа pre-training
Этап pre-training создаёт модель, которая хорошо понимает язык. Эта модель, обычно называемая базовой моделью, предсказывает текст, продолжающий данный входной промпт, генерируя релевантный и осмысленный текст.
Хотя базовая модель хорошо понимает язык, она способна лишь продолжать текстовый промпт. Чтобы сделать модель полезным чатботом, отвечающим на вопросы, необходимо дополнительно обучить базовую модель. Это приводит нас к следующему этапу: supervised finetuning.
. Supervised finetuning (SFT)
SFT, также называемый instruction finetuning, — это второй этап процесса обучения. На этом этапе мы дообучаем базовую модель на меньшем высококачественном наборе данных в формате (промпт, ответ). Цель этого этапа — сохранить языковое понимание базовой модели и знания о мире, адаптируя её поведение к ответам на промпты вместо простого их продолжения.
Обучающие данные
Обучающие данные для этапа SFT следуют формату (промпт, ответ). Эти данные обычно называются демонстрационными данными, поскольку они демонстрируют модели, как отвечать на промпты.
Основные различия между демонстрационными данными и данными pre-training, помимо формата, — это размер и качество.
Размер: Демонстрационные данные значительно меньше данных pre-training. Обычно они составляют от 10 000 до 100 000 пар (промпт, ответ). Таблица 2 показывает размеры данных популярных демонстрационных наборов.
| Dataset | Size | Notes |
|---|---|---|
| InstructGPT [20] | ~14,500 | OpenAI’s GPT-3 instruction datasets |
| Alpaca [21] | 52,000 | Developed by Stanford researchers |
| Dolly-15K [22] | ~15,000 | Created by Databricks |
| FLAN 2022 [23] | ~104,000 | Developed by Google Research |
Таблица 2: Распространённые демонстрационные наборы данных
Качество: Демонстрационные данные имеют более высокое качество по сравнению с данными pre-training. Данные обычно создаются квалифицированными подрядчиками. В специализированных отраслях, таких как здравоохранение или финансы, необходимо привлекать экспертов в предметной области для обеспечения точности и релевантности данных. Например, как показано в Таблице 3, более трети разметчиков OpenAI для демонстрационного набора данных GPT имели степень магистра [20]. Хотя это требование затратно, оно критически важно для создания надёжных, отраслевых ответов.
| Education | Percentage |
|---|---|
| Less than a high school degree | 0% |
| High school degree | 10.5% |
| Undergraduate degree | 52.6% |
| Master’s degree | 36.8% |
Таблица 3: Уровни образования разметчиков OpenAI
Целевая функция ML и функция потерь
Хотя обучающие данные отличаются от этапа pre-training, модель по-прежнему обучается аналогичной задаче: генерации текста по одному токену за раз на основе входного промпта. Поэтому целевая функция ML и функция потерь остаются аналогичными тем, что на этапе pre-training: целевая функция предсказания следующего токена и функция потерь кросс-энтропии.
Результат этапа SFT
Результатом этого этапа является модель SFT — дообученная версия базовой модели. Вместо простого продолжения текстового промпта модель SFT генерирует подробные и полезные ответы, поскольку она была обучена на демонстрационных данных в формате (промпт, ответ).
Модель SFT обычно генерирует грамматически корректный и разумный ответ. Однако она не всегда генерирует лучший ответ; её ответы могут быть бесполезными или даже небезопасными. Рисунок 11 показывает четыре правдоподобных ответа на вопрос. Только второй ответ является одновременно безопасным и полезным. Первый и четвёртый ответы грамматически и контекстуально корректны, но не дают точных рекомендаций. Третий ответ полезен, но невежлив.
Чтобы гарантировать, что модель генерирует релевантные, безопасные и полезные ответы, необходимо дополнительно дообучить модель. Это дополнительное дообучение является основным фокусом следующего этапа: RLHF.
. RLHF
RLHF, также известный как этап выравнивания (alignment), является финальным этапом процесса обучения. Этот этап выравнивает модель с человеческими предпочтениями, то есть адаптирует модель для генерации ответов, предпочитаемых людьми.
Чтобы понять RLHF, давайте кратко вернёмся к этапу SFT. На этапе SFT модель обучается на демонстрационных данных генерировать правдоподобный ответ на заданный промпт. Однако демонстрационные данные предоставляют модели один правдоподобный ответ на промпт, который не обязательно является наиболее полезным или релевантным. Обычно возможны несколько правдоподобных ответов, и одни будут более релевантными, чем другие, как показано на Рисунке 11.
Если у нас есть отдельная модель вознаграждения, которая может оценить, насколько релевантен ответ модели промпту, мы можем дополнительно дообучить модель SFT, чтобы она генерировала не просто любой правдоподобный ответ, а ответ с высокой оценкой. Это ключевая идея RLHF. RLHF состоит из двух последовательных шагов:
- Обучение модели вознаграждения
- Оптимизация модели SFT
.1 Обучение модели вознаграждения
Первый шаг в RLHF — обучение модели вознаграждения, которая оценивает релевантность ответа промпту. Эта модель принимает пару (промпт, ответ) на вход и выдаёт оценку, предсказывающую полезность ответа. Чем выше оценка, тем более полезным ожидается ответ. Рисунок 12 иллюстрирует предсказанные оценки для различных пар (промпт, ответ).
Архитектура модели вознаграждения
Обучение модели для выдачи оценки — очень распространённая задача в ML. Существуют различные архитектуры, которые мы можем использовать для моделирования вознаграждения: это может быть decoder-only, encoder-only или encoder-decoder Transformer, если он выдаёт скалярное значение.
На основе публичных исследований нет устойчивой закономерности в том, что модели вознаграждения больше или меньше языковых моделей, для обучения которых они используются. Например, OpenAI использует модель вознаграждения на 6B для языковой модели на 175B [20]. Anthropic использует языковые модели и модели вознаграждения от 10B до 52B параметров [24]. Типичный вариант — создать копию модели SFT и добавить голову предсказания для получения оценки релевантности для данной пары (промпт, ответ).
Обучающие данные
Для сбора обучающих данных для моделирования вознаграждения мы следуем этим шагам:
- Сбор промптов: Вручную создать список промптов.
- Генерация нескольких ответов: Использовать модель SFT для генерации нескольких ответов на каждый промпт.
- Ранжирование ответов: Попросить подрядчиков оценить эти ответы и ранжировать их по релевантности. Причина, по которой обычно используется ранжирование, а не оценка каждого ответа, состоит в том, что ранжирование снижает субъективность и несогласованность. Аннотаторам проще и интуитивнее сравнивать ответы напрямую, чем присваивать числовые оценки, которые могут варьироваться между аннотаторами. Этот подход упрощает процесс оценки и обеспечивает более надёжные данные для обучения.
- Создание пар предпочтений: Сформировать обучающий набор данных, создавая пары в формате (промпт, выигрышный ответ, проигрышный ответ). В каждой паре выигрышный ответ предпочтительнее проигрышного на основе ранжирования с предыдущего шага.
Рисунок 14 показывает процесс сбора обучающих данных для обучения модели вознаграждения.
После сбора обучающих данных, где каждый пример имеет формат (промпт, выигрышный ответ, проигрышный ответ), мы определяем целевую функцию ML и соответствующую функцию потерь для обучения нашей модели вознаграждения.
Целевая функция ML и функция потерь
Модель вознаграждения стремится предсказать более высокую оценку для выигрышного ответа по сравнению с проигрышным. Более формально, для данного (промпт, выигрышный ответ, проигрышный ответ) целевая функция ML — максимизировать Swin −Slose S_{\text {win }}-S_{\text {lose }}Swin −Slose , где:
- Swin S_{\text {win }}Swin — предсказанная оценка для пары (промпт, выигрышный ответ)
- Slose S_{\text {lose }}Slose — предсказанная оценка для пары (промпт, проигрышный ответ)
Для достижения этой целевой функции ML нам нужна функция потерь, которая штрафует модель, когда разница между оценками выигрышного и проигрышного ответов слишком мала. Часто используемая функция потерь для этой цели — маржинальная ранговая функция потерь (margin ranking loss). Функция потерь определяется как:
Где mmm — это гиперпараметр, определяющий маржу. Эта маржа указывает минимальную желаемую разницу между оценками выигрышного и проигрышного ответов. Если разница между SwinS_{\text {win}}Swin и SloseS_{\text {lose}}Slose меньше mmm, оптимизатор обновит параметры модели так, чтобы либо SwinS_{\text {win}}Swin увеличилась, либо SloseS_{\text {lose}}Slose уменьшилась.
Результат моделирования вознаграждения
Результатом этого шага является модель вознаграждения, которая предсказывает оценки релевантности для пар (промпт, ответ). Эти оценки отражают человеческие суждения и критически важны для второго шага в RLHF.
.2. Оптимизация модели SFT
На втором шаге RLHF модель SFT оптимизируется с помощью модели вознаграждения. Цель этого шага — адаптировать модель SFT для генерации ответов, которые не только правдоподобны, но и полезны, на основе оценок модели вознаграждения.
Распространённый подход к оптимизации модели SFT — использование алгоритма обучения с подкреплением (RL), такого как proximal policy optimization (PPO) [25], при котором модель SFT дообучается для максимизации оценок, предсказанных моделью вознаграждения. Этот процесс дообучения итеративно выполняет следующие шаги:
- Генерация ответов модели: Модель генерирует несколько возможных ответов на заданный промпт.
- Вычисление вознаграждений: Модель вознаграждения оценивает эти ответы.
- Обновление весов модели: Алгоритм RL обновляет веса модели для максимизации ожидаемого вознаграждения. Этот шаг подкрепляет ответы, получившие более высокие оценки от модели вознаграждения.
Рисунок 17 показывает этот процесс для одного ответа. На практике несколько ответов генерируются и оцениваются одновременно.
Обучающие данные
Для этого шага обучающие данные обычно включают список промптов, созданных подрядчиками, обычно в количестве от 10 000 до 100 000.
Целевая функция ML и функция потерь
Известные LLM, такие как ChatGPT и Llama, используют алгоритмы RL, такие как PPO и direct policy optimization (DPO) [26]. Однако детали этих алгоритмов обычно выходят за рамки большинства интервью по проектированию ML-систем. Для получения дополнительной информации обратитесь к [27] и [28].
Результат RLHF
Результатом этапа RLHF обычно является финальная модель, которая может быть развёрнута как чатбот. Таблица 4 перечисляет некоторые из наиболее популярных LLM.
| LLM name | Developer | Release date | Access | Parameters |
|---|---|---|---|---|
| o1 | OpenAI | September 12, 2024 | Preview only | Unknown |
| GPT-4o | OpenAI | May 13, 2024 | API | Unknown |
| Claude 3 | Anthropic | March 14, 2024 | API | Unknown |
| Gemini 1.5 | DeepMind | February 2, 2024 | API | Unknown |
| Llama 3 | Meta AI | April 18, 2024 | Open-Source | 8 and 70 billion |
| Grok-1 | xAI | November 4, 2023 | Open-Source | 314 billion |
| Mixtral 8x22B | Mistral AI | April 10, 2024 | Open-Source | 141 billion |
| Gemma | DeepMind | February 21, 2024 | Open-Source | 2 and 7 billion |
| Phi-3 | Microsoft | April 23, 2024 | Open-Source | 3.8 billion |
| DBRX | Databricks | March 27, 2024 | Open-Source | 132 billion |
Таблица 4: Популярные LLM
Подводя итог раздела об обучении, мы применяем трёхэтапную стратегию обучения, включающую pre-training, SFT и RLHF. Pre-training включает обучение модели на большом корпусе текста для получения широкого языкового понимания. SFT дообучает модель для адаптации её вывода к формату (промпт, ответ). RLHF дополнительно совершенствует ответы модели, делая их полезными, безопасными и согласованными с человеческими предпочтениями.
Сэмплирование
В LLM сэмплирование означает то, как мы выбираем токены из предсказанного моделью распределения вероятностей для генерации связных и полезных ответов.
Как обсуждалось в Главе 2, существуют различные методы генерации текста. Некоторые из них детерминированные, а другие — стохастические. В этом разделе мы рассмотрим эти методы, чтобы определить, какой лучше подходит для открытой генерации текста.
Детерминированные методы
Детерминированные методы, такие как beam search, хорошо работают для задач с коротким предсказуемым объёмом текста. Однако они менее эффективны для открытой генерации, такой как диалог, где длина выхода варьируется. Давайте рассмотрим типичные проблемы, возникающие при использовании детерминированных методов, таких как жадный поиск или beam search, для генерации текста LLM.
Жадный поиск
Жадный поиск выбирает токен с наивысшей вероятностью на каждом шаге процесса генерации.
Хотя этот метод прост и часто создаёт связный текст, у него есть два основных недостатка:
- Повторение
- Субоптимальная генерация
Повторение: Когда мы используем жадный поиск для выбора следующих токенов, текст быстро начинает повторяться. Это происходит потому, что модель иногда «застревает» в циклах, повторно используя одну и ту же последовательность токенов. Это случается, когда модель определяет, что определённые слова следуют друг за другом с высокой вероятностью.
Субоптимальная генерация: Жадный поиск игнорирует альтернативные пути в процессе генерации текста. Он может пропустить последовательность токенов с высокой вероятностью, скрытую за токеном с низкой вероятностью.
Beam search
Beam search улучшает жадный поиск за счёт рассмотрения нескольких последовательностей одновременно. На каждом шаге он отслеживает топ-k последовательностей, где k — настраиваемый параметр.
Beam search позволяет более широкий поиск и создаёт текст более высокого качества, чем жадный поиск. Однако он может испытывать трудности с открытой генерацией. Две типичные проблемы beam search:
- Неэффективность
- Повторение
Неэффективность: Beam search может быть вычислительно неэффективным, поскольку требует одновременной оценки нескольких последовательностей, что может замедлить процесс генерации.
Повторение: Beam search может приводить к повторяющимся и шаблонным ответам. Иногда он застревает в цикле и повторяет распространённые фразы.
До сих пор мы видели, что детерминированные методы испытывают трудности с повторением и плохо подходят для генерации текста. Давайте рассмотрим стохастические методы, которые чаще используются для генерации текста в LLM.
Стохастические методы
Стохастические методы генерируют текст, вводя случайность. Эта случайность делает их более подходящими для открытой генерации текста. Три популярных стохастических метода:
- Мультиномиальное сэмплирование
- Top-k сэмплирование
- Top-p (nucleus) сэмплирование
Мультиномиальное сэмплирование
Мультиномиальное сэмплирование выбирает следующий токен на основе распределения вероятностей предсказаний модели. Каждому токену соответствует определённая вероятность, и токен выбирается на основе этих вероятностей.
Этот подход обеспечивает широкое разнообразие возможных выходов. Однако он вносит значительную долю случайности, особенно когда распределение вероятностей плоское. Эта случайность часто приводит к несвязным результатам генерации. Например, сгенерированный текст, показанный на Рисунке 25, является выходом модели GPT-2 при использовании мультиномиального сэмплирования.
Из-за проблем со связностью мультиномиальное сэмплирование редко используется в LLM для генерации текста.
Top-k сэмплирование
Top-k сэмплирование [30] — это более продвинутый метод, который выбирает из k наиболее вероятных токенов, а не из всего распределения.
Вот пошаговый процесс выбора следующего токена при top-k сэмплировании:
- Модель предсказывает распределение вероятностей для следующего токена, предоставляя вероятность для каждого токена в словаре.
- Токены сортируются в порядке убывания по предсказанным вероятностям.
- Топ-k токенов с наивысшими вероятностями рассматриваются для сэмплирования.
- Вероятности топ-k токенов нормализуются, чтобы их сумма равнялась 1.
- Токен выбирается из этого нормализованного распределения.
Top-k сэмплирование балансирует между связностью и разнообразием, выбирая из топ-k токенов. Это снижает вероятность выбора нерелевантных токенов, при этом допуская некоторую случайность. GPT-2 изначально использовал top-k сэмплирование, что было ключевым для его успеха и популярности.
Главным ограничением top-k сэмплирования является то, что оно всегда выбирает из фиксированного количества топ-токенов. Это проблематично в зависимости от того, как распределены предсказанные вероятности. Давайте разберёмся почему.
Предсказанные вероятности токенов могут быть распределены остро или равномерно. При остром распределении ограничение выбора фиксированным числом топ-токенов может привести к бессмысленным результатам, поскольку модель может пропустить лучший выбор. Напротив, при плоском распределении это фиксированное ограничение сдерживает креативность модели, не рассматривая достаточно вариантов слов. Например, как показано на Рисунке 28, модель на 89% уверена, что следующим токеном должно быть «lot», но top-k сэмплирование по-прежнему рассматривает «much» и «high» как возможные следующие токены для выборки.
Это ограничение решается в top-p сэмплировании, которое мы рассмотрим далее.
Top-p (nucleus) сэмплирование
Top-p сэмплирование [31], также известное как nucleus сэмплирование, было разработано в 2019 году. Этот метод динамически регулирует количество рассматриваемых токенов на основе их совокупных вероятностей. Вместо выборки только из наиболее вероятных k токенов, он выбирает из наименьшего возможного набора токенов, совокупная вероятность которых превышает порог p. Это обеспечивает более гибкий и адаптивный подход по сравнению с top-k сэмплированием.
Вот пошаговый процесс выбора следующего токена при top-p сэмплировании:
- Модель предсказывает распределение вероятностей для следующего токена, предоставляя вероятность для каждого токена в словаре.
- Токены сортируются в порядке убывания по предсказанным вероятностям.
- Вместо выбора фиксированного количества токенов top-p сэмплирование выбирает наименьший возможный набор токенов, совокупная вероятность которых превышает порог p.
- Вероятности выбранных токенов нормализуются, чтобы их сумма равнялась 1.
- Токен выбирается из этого нормализованного распределения.
Top-p сэмплирование широко используется в продвинутых LLM для генерации текста, похожего на человеческий. Этот метод обеспечивает связность и контекстуальную релевантность текста, фокусируясь на наиболее вероятных токенах, допуская при этом некоторую случайность.
Хотя мы рассмотрели ключевые аспекты различных методов сэмплирования, каждый метод имеет больше деталей. Две популярные техники, часто используемые в продвинутых методах сэмплирования:
- Температура
- Штраф за повторение
Температура
Температура — это параметр в методах сэмплирования, который контролирует случайность предсказаний во время сэмплирования. Математически параметр температуры T масштабирует логиты (необработанные оценки) выхода модели перед применением функции softmax для генерации вероятностей. Скорректированная формула softmax с температурой задаётся как:
где:
- xix_ixi — логиты (необработанные оценки) для каждого возможного выхода
- TTT — параметр температуры
- pip_ipi представляет вероятность выхода iii после применения функции softmax
При T=1T=1T=1 функция softmax работает нормально. При T>1T>1T>1 модель генерирует более равномерное распределение вероятностей, делая предсказания более случайными и разнообразными. Более высокие температуры увеличивают случайность, что помогает модели генерировать более креативные выходы. Если модель начинает отклоняться от темы или создавать бессмысленные выходы, это указывает на то, что температура установлена слишком высоко.
Напротив, при T<1T<1T<1 выход модели становится более детерминированным, при этом наивысшие значения логитов оказывают большее влияние на финальное предсказание. Более низкие температуры снижают случайность, что более подходит для задач, требующих точных ответов, таких как суммаризация или перевод. Если модель начинает повторяться, это указывает на то, что температура установлена слишком низко. Температура 0 стабильно приводит к одному и тому же выходу, делая сэмплирование детерминированным.
Каковы типичные значения температуры?
Большинство провайдеров моделей устанавливают допустимый диапазон температуры от 0 до 2. Рисунок 32 иллюстрирует справочник API OpenAI по настройке температуры.
В современных LLM параметр температуры обычно варьируется от 0.1 до 1.5. При значениях выше 1.5 выходы могут становиться всё более хаотичными и менее связными, что нежелательно. Оптимальное значение зависит от желаемого поведения и часто определяется эмпирически. Следующая таблица, созданная [33], предлагает возможные значения температуры для нескольких вариантов использования.
| Use case | Temperature | Top-p | Description |
|---|---|---|---|
| Code generation | 0.2 | 0.1 | Generates code that adheres to established patterns and conventions. Output is more deterministic and focused. Useful for generating syntactically correct code. |
| Creative writing | 0.7 | 0.8 | Generates creative and diverse text for storytelling. Output is more exploratory and less constrained by patterns. |
| Chatbot responses | 0.5 | 0.5 | Generates conversational responses that balance coherence and diversity. Output is more natural and engaging. |
Таблица 5: Эмпирические диапазоны температуры и top-p для различных задач
Штраф за повторение
Аналогично, применение штрафа за повторение может явно снизить вероятность генерации повторяющихся последовательностей токенов. Это может быть достигнуто путём завершения генерации при обнаружении повторяющихся n-грамм (как контролируется параметром «no_repeat_ngram_size» в моделях Hugging Face) или путём прямого изменения вероятностей токенов, которые уже были выбраны ранее в последовательности (как параметр «frequency_penalty» в API ChatGPT).
Если вы хотите узнать больше о методах сэмплирования в LLM, обратитесь к [30].
Оценка
Офлайн-метрики оценки
Оценка LLM, таких как ChatGPT, требует большего, чем традиционные метрики, такие как перплексия. Эти модели ведут себя сложным образом и показывают различную производительность на разных задачах. Поэтому нам необходимо оценить их навыки в различных задачах, чтобы убедиться, что модель одновременно эффективна и безопасна.
В этом разделе мы оцениваем нашу LLM со следующих перспектив:
- Традиционная оценка
- Оценка по конкретным задачам
- Оценка безопасности
- Оценка людьми
Традиционная оценка
Традиционная оценка обеспечивает начальное понимание производительности LLM с использованием типичных офлайн-метрик. Распространённая метрика — перплексия, которая измеряет, насколько точно модель предсказывает точную последовательность токенов в обучающих данных. Низкое значение перплексии указывает, что модель в среднем присваивает более высокие вероятности токенам в обучающих данных.
Хотя эти метрики важны для начальной оценки, они не дают представления о возможностях или ограничениях LLM. Например, низкая перплексия указывает, что модель хорошо предсказывает следующие токены, но не измеряет её способность понимать код или решать математические задачи.
Оценка по конкретным задачам
Для эффективной оценки LLM необходимо оценить её производительность на разнообразных задачах, таких как математика, генерация кода и рассуждение на основе здравого смысла. Этот комплексный подход помогает выявить сильные и слабые стороны модели. Часто используемые задачи для оценки возможностей LLM:
- Рассуждение на основе здравого смысла
- Знания о мире
- Понимание прочитанного
- Математическое рассуждение
- Генерация кода
- Комплексные бенчмарки
Рассуждение на основе здравого смысла
Рассуждение на основе здравого смысла оценивает способность модели делать выводы на основе повседневных ситуаций и общих знаний. Оно проверяет понимание моделью базового человеческого опыта, логических связей и предположений, которые люди делают естественным образом. Примеры включают интерпретацию идиом, понимание причинно-следственных связей в типичных сценариях и предсказание вероятных исходов в социальных ситуациях.
Типичные бенчмарки для рассуждения на основе здравого смысла — PIQA (Physical Interaction QA) [34], SIQA [35], HellaSwag [36], WinoGrande [37], OpenBookQA [38] и CommonsenseQA [39], каждый из которых фокусируется на различных аспектах. Например, бенчмарк CommonsenseQA — это набор данных с вопросами с множественным выбором, для ответа на которые требуются знания здравого смысла. PIQA фокусируется на рассуждениях о физических взаимодействиях в повседневных ситуациях, а HellaSwag — на повседневных событиях.
Знания о мире
Знания о мире относятся к фактическим знаниям модели о мире, включая исторические факты, научную информацию, географию и текущие события. Примером может быть ответ на вопросы о значимых исторических событиях или научных принципах.
Распространённые бенчмарки для этой задачи включают:
- TriviaQA [40]: Вопросы собраны с сайтов викторин и квиз-лиг.
- Natural Questions (NQ) [41]: Набор данных от Google, включающий вопросы и ответы, найденные в естественных веб-запросах.
- SQuAD (Stanford Question Answering Dataset) [42]: Содержит вопросы на основе статей Wikipedia.
Понимание прочитанного
Задачи на понимание прочитанного оценивают способность модели понимать и интерпретировать текстовые отрывки и отвечать на вопросы по ним. Это критически важно для оценки способности модели извлекать информацию из текстов и рассуждать на их основе.
Типичные бенчмарки для понимания прочитанного — SQuAD [42], QuAC [43] и BoolQ [44].
Математическое рассуждение
Задачи на математическое рассуждение оценивают способность модели решать математические задачи.
Два распространённых бенчмарка для задач математического рассуждения:
- MATH [46]: Набор данных, содержащий задачи из школьных математических олимпиад.
- GSM8K (Grade School Math 8K) [45]: Набор данных с математическими задачами начальной школы для проверки навыков решения задач моделью.
Генерация кода
Генерация кода оценивает способность модели писать синтаксически корректный и функциональный код на основе промпта на естественном языке.
Распространённые бенчмарки для генерации кода:
- HumanEval [47]: Задачи по программированию на Python.
- MBPP (MultiPL-E Benchmarks for Programming Problems) [48]: Задачи на нескольких языках программирования для оценки возможностей мультиязычной генерации кода.
Комплексные бенчмарки
Помимо конкретных бенчмарков, описанных выше, комплексные бенчмарки объединяют несколько задач для более широкой оценки. Популярные комплексные бенчмарки:
- MMLU (Massive Multitask Language Understanding) [49]: Состоит из вопросов с множественным выбором из широкого спектра предметов, включая гуманитарные науки, STEM, социальные науки и другие, с различными уровнями сложности.
- MMMU (Massive Multilingual Multitask Understanding) [50]: MMMU включает широкий спектр вопросов с множественным выбором, охватывающих множество предметов с различными уровнями сложности. В отличие от MMLU, который фокусируется на английском, MMMU тестирует способность моделей понимать и генерировать точные ответы на разных языках, оценивая не только мультиязычные возможности, но и способность к рассуждению и межкультурные знания.
- AGIEval [51]: Комплексный бенчмарк, предназначенный для тестирования искусственного общего интеллекта в нескольких областях и задачах.
- Meta Llama 3 human evaluation [13]: Высококачественный набор человеческой оценки, содержащий 1800 промптов, охватывающих 12 ключевых вариантов использования: запрос совета, мозговой штурм, классификация, закрытые вопросы, программирование, творческое письмо, извлечение информации, перевоплощение в персонажа, открытые вопросы, рассуждение, переформулирование и суммаризация.
Подводя итог, мы используем различные задачи и бенчмарки для оценки производительности LLM на конкретных задачах. Эта оценка охватывает понимание и генерацию ответов, подобных человеческим, в различных областях. Однако оценка на этом не заканчивается. Оценка безопасности критически важна для ответственного развёртывания этих моделей. Давайте рассмотрим подробнее.
Оценка безопасности
Оценки безопасности LLM критически важны для обеспечения того, чтобы эти модели генерировали безопасные и этичные ответы. Эти оценки фокусируются на различных задачах, которые помогают выявить и смягчить риски, такие как генерация вредоносного контента. Основные аспекты оценки безопасности включают:
- Токсичность и вредоносный контент
- Предвзятость и справедливость
- Достоверность
- Конфиденциальность пользователей и утечка данных
- Устойчивость к состязательным атакам
Токсичность и вредоносный контент
Мы оцениваем способность модели избегать генерации токсичного контента. Токсичность включает:
- Язык ненависти
- Оскорбительный язык
- Контент, который может нанести вред отдельным лицам, группам или обществу
- Контент, полезный для планирования нападений или насилия
- Инструкции для поиска нелегального контента
Часто используемые бенчмарки для оценки токсичности модели:
- RealToxicityPrompts [52]: Состоит из примерно 100 000 промптов, которые модель должна дополнить; затем оценка токсичности автоматически вычисляется с помощью PerspectiveAPI [53].
- ToxiGen [54]: Этот бенчмарк тестирует способность модели избегать генерации дискриминационного языка.
- HateCheck [55]: Набор тестов, специально предназначенных для обнаружения языка ненависти, охватывающий различные типы языка ненависти.
Оценка моделей с использованием этих бенчмарков помогает выявить потенциальные риски и улучшить способность моделей генерировать безопасный и уважительный контент.
Предвзятость и справедливость
Мы оцениваем ответы модели на наличие потенциальной предвзятости. Это включает обнаружение гендерной, расовой и других предвзятостей в сгенерированном контенте.
Типичные бенчмарки:
- CrowS-Pairs [56]: Содержит парные предложения, отличающиеся только одним атрибутом (например, полом) для тестирования предвзятости. Этот набор данных позволяет измерять предвзятости в 9 категориях: пол, религия, раса/цвет кожи, сексуальная ориентация, возраст, национальность, инвалидность, физическая внешность и социоэкономический статус.
- BBQ [57]: Набор данных из вручную написанных вопросов, нацеленных на подтверждённые социальные предвзятости в отношении различных социально значимых категорий.
- BOLD [58]: Крупномасштабный набор данных, состоящий из 23 679 промптов для генерации текста на английском языке для бенчмаркинга предвзятости в пяти областях.
Эти бенчмарки помогают нам убедиться, что модель относится ко всем демографическим группам справедливо и одинаково.
Достоверность
Мы оцениваем способность LLM генерировать правдивые и фактически точные ответы. Это включает различение фактической информации и распространённых заблуждений или ложных утверждений.
Распространённый бенчмарк для оценки достоверности — TruthfulQA [59]. Он измеряет правдивость модели, то есть её способность определять, когда утверждение истинно. Этот бенчмарк позволяет оценить риски генерации моделью дезинформации или ложных утверждений.
Конфиденциальность пользователей и утечка данных
Мы оцениваем склонность LLM к утечке конфиденциальной информации, с которой она могла столкнуться во время обучения. Поскольку LLM обучаются на различных общедоступных источниках данных, они могут знать о людях, имеющих публичное присутствие в интернете. Эти оценки гарантируют, что LLM не раскрывают случайно персональную информацию. Распространённый бенчмарк для этой цели — PrivacyQA [60].
Устойчивость к состязательным атакам
Устойчивость к состязательным атакам тестирует способность LLM обрабатывать входные данные, намеренно созданные для того, чтобы запутать или обмануть модель. Это критически важно для обеспечения надёжности и безопасности модели на практике. Типичные бенчмарки для тестирования устойчивости LLM к состязательным атакам включают AdvGLUE [61], TextFooler [62] и AdvBench [63].
Подводя итог, мы используем различные бенчмарки для оценки безопасности LLM, что критически важно для обеспечения безопасности пользователей. Хотя и оценка по конкретным задачам, и оценка безопасности необходимы, оценка людьми остаётся наиболее надёжным методом комплексной оценки.
Оценка людьми
В этом подходе людям-оценщикам предлагается оценить различные аспекты LLM, такие как полезность и безопасность. Оценка людьми критически важна для оценки нюансированных аспектов полезности и безопасности, которые бенчмарки по конкретным задачам и безопасности могут упустить.
Онлайн-метрики оценки
Онлайн-метрики оценки измеряют производительность LLM при развёртывании в продакшене. Часто используемые метрики:
- Обратная связь и оценки пользователей
- Вовлечённость пользователей
- Коэффициент конверсии
- Онлайн-рейтинги
Обратная связь и оценки пользователей: Пользователи могут оценивать свою удовлетворённость ответами модели. Эта прямая обратная связь от пользователей выявляет области, требующие улучшения.
Вовлечённость пользователей: Метрики, такие как «количество выполненных запросов» и «продолжительность сессии», могут быть показательными сигналами для измерения вовлечённости пользователей. Высокий уровень вовлечённости часто указывает на то, что LLM эффективно предоставляет полезную информацию.
Коэффициент конверсии: Коэффициент конверсии — это процент пользователей, которые совершают покупку или подписываются на сервис после взаимодействия с LLM. Коэффициент конверсии — важная метрика для мониторинга, поскольку более высокие коэффициенты указывают на то, что пользователи находят LLM достаточно полезной, чтобы платить за сервис.
Онлайн-рейтинги: Онлайн-рейтинги отслеживают производительность различных LLM в режиме реального времени. Заметным примером является LMSYS Chatbot Arena [64] — краудсорсинговая открытая платформа, предназначенная для оценки LLM. Модели ранжируются на основе более чем 800 000 парных сравнений людьми.
Общий дизайн ML-системы
Проектирование системы чатбота, такой как ChatGPT, требует слаженной работы нескольких компонентов. В отличие от традиционных моделей, эта система объединяет множество сервисов и пайплайнов для обеспечения эффективности, безопасности и непрерывного улучшения. В этом разделе мы рассмотрим два ключевых пайплайна:
- Пайплайн обучения
- Пайплайн инференса
Пайплайн обучения
Пайплайн обучения включает три критических этапа: pre-training, SFT и RLHF. Эти этапы в совокупности обеспечивают способность модели генерировать полезные и безопасные ответы.
Пайплайн инференса
Пайплайн инференса включает несколько компонентов, обеспечивающих безопасность, релевантность и качество генерируемых ответов. Этот пайплайн отвечает за взаимодействие с пользователями в реальном времени. Ключевые компоненты пайплайна инференса:
- Фильтрация безопасности
- Улучшение промпта
- Генератор ответов
- Оценщик безопасности ответов
- Генератор ответов отклонения
- Управление сессиями
Давайте подробнее рассмотрим каждый компонент.
Фильтрация безопасности
Этот компонент анализирует промпт пользователя для обнаружения вредоносных, неуместных или небезопасных запросов до их обработки моделью. Например, промпт с просьбой об инструкциях по созданию опасного устройства будет отклонён и помечен.
Улучшение промпта
Компонент улучшения промпта уточняет и обогащает входной промпт, делая его более информативным и детальным. Он расшифровывает аббревиатуры, исправляет опечатки и добавляет контекст при необходимости.
Этот компонент обеспечивает ясность, однозначность и грамматическую корректность текстовых промптов перед их передачей модели, что помогает модели генерировать лучшие ответы.
Генератор ответов
Генератор ответов взаимодействует с обученной LLM и использует top-p сэмплирование для генерации полезного ответа. Этот компонент может дополнительно использовать другие техники для улучшения качества и безопасности сгенерированного ответа. Например, он может генерировать несколько возможных ответов и затем выбирать наиболее подходящий на основе набора предопределённых критериев.
Оценщик безопасности ответов
Этот компонент оценивает сгенерированный ответ для обнаружения вредоносного или неуместного контента до его показа пользователю. Он действует как финальная защита для обеспечения соответствия ответов этическим стандартам и стандартам безопасности.
Генератор ответов отклонения
Этот компонент генерирует надлежащий ответ, когда входной промпт небезопасен или сгенерированный ответ неподходящий. Он предоставляет ясное и вежливое объяснение, почему запрос не может быть выполнен.
Управление сессиями
Для эффективного поддержания контекста диалога и обработки уточняющих вопросов требуется специальная обработка. Например, когда пользователь общается с моделью о своих любимых фильмах, модель должна помнить не только текущий вопрос, но и предыдущие упоминания различных жанров или фильмов.
Этот компонент поддерживает непрерывность и связность диалога, отслеживая историю чата и управляя потоком диалога. Это достигается путём подачи истории чата вместе с улучшенным промптом в генератор ответов. Такой дизайн обеспечивает контекстуальную релевантность каждого ответа путём обращения к предыдущим взаимодействиям и надлежащего управления состоянием диалога.
Другие темы для обсуждения
Если в конце собеседования останется дополнительное время, вот несколько дополнительных тем для обсуждения:
- Техники управления состояниями диалога и отслеживания контекста в многооборотных диалогах [65].
- Применение продвинутых или более эффективных целевых функций ML, таких как предсказание нескольких токенов [66].
- Обработка очень длинных последовательностей [67][68].
- Как разрабатывать мультимодальные LLM [69][70].
- Техники, такие как RAG, для использования внешних баз знаний и баз данных для улучшения вывода LLM [71]. Мы рассмотрим это в Главе 6.
- Техники повышения эффективности (например, дистилляция) для более быстрой генерации текста.
- Техники адаптации LLM к конкретным областям (например, обслуживание клиентов, здравоохранение) без потери предыдущих знаний [72].
- Решение проблем безопасности и конфиденциальности в LLM.
- Различные алгоритмы оптимизации, такие как PPO, DPO и rejection sampling [73].
- Red-teaming LLM для снижения вреда [74].
- Super-alignment и его важность в разработке LLM [75].
- Как работает обучение в контексте (in-context learning) [76].
- Grouped query attention и его преимущества [77].
- Применение техник chain-of-thought промптинга [78]. Мы рассмотрим это в Главе 6.
- Реализация KV cache [79].
- Повышение доверия путём требования от моделей создавать ясные и проверяемые обоснования своих выходов [80].
Резюме
Справочные материалы
[1] OpenAI’s ChatGPT. https://openai.com/index/chatgpt/. [2] ChatGPT wiki. https://en.wikipedia.org/wiki/ChatGPT. [3] OpenAI’s models. https://platform.openai.com/docs/models. [4] Google’s Gemini. https://gemini.google.com/. [5] Meta’s Llama. https://llama.meta.com/. [6] Beautiful Soup. https://beautiful-soup-4.readthedocs.io/en/latest/. [7] Lxml. https://lxml.de/. [8] Document Object Model. https://en.wikipedia.org/wiki/Document_Object_Model. [9] Boilerplate removal tool. https://github.com/miso-belica/jusText. [10] fastText. https://fasttext.cc/. [11] langid. https://github.com/saffsd/langid.py. [12] RoFormer: Enhanced Transformer with Rotary Position Embedding. https://arxiv.org/abs/2104.09864. [13] Llama 3 human evaluation. https://github.com/meta-llama/llama3/blob/main/eval_details.md. [14] Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer. https://arxiv.org/abs/1910.10683. [15] DeBERTa: Decoding-enhanced BERT with Disentangled Attention. https://arxiv.org/abs/2006.03654. [16] Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention. https://arxiv.org/abs/2006.16236. [17] Common Crawl. https://commoncrawl.org/. [18] C4 dataset. https://www.tensorflow.org/datasets/catalog/c4. [19] Stack Exchange dataset. https://github.com/EleutherAI/stackexchange-dataset. [20] Training language models to follow instructions with human feedback. https://arxiv.org/abs/2203.02155. [21] Alpaca. https://crfm.stanford.edu/2023/03/13/alpaca.html. [22] Dolly-15K. https://www.databricks.com/blog/2023/04/12/dolly-first-open-commercially-viable-instruction-tuned-llm. [23] Introducing FLAN: More generalizable Language Models with Instruction Fine-Tuning. https://research.google/blog/introducing-flan-more-generalizable-language-models-with-instruction-fine-tuning/. [24] Training a Helpful and Harmless Assistant with Reinforcement Learning from Human Feedback. https://arxiv.org/abs/2204.05862. [25] Proximal Policy Optimization Algorithms. https://arxiv.org/abs/1707.06347. [26] Direct Preference Optimization: Your Language Model is Secretly a Reward Model. https://arxiv.org/abs/2305.18290. [27] Illustrating RLHF. https://huggingface.co/blog/rlhf. [28] RLHF progress and challenges. https://www.youtube.com/watch?v=hhiLw5Q_UFg. [29] State of GPT. https://www.youtube.com/watch?v=bZQun8Y4L2A. [30] Different sampling methods. https://huggingface.co/blog/how-to-generate. [31] The Curious Case of Neural Text Degeneration. https://arxiv.org/abs/1904.09751. [32] OpenAI’s API reference. https://platform.openai.com/docs/api-reference/chat/create. [33] Cheat Sheet: Mastering Temperature and Top_p in ChatGPT API. https://community.openai.com/t/cheat-sheet-mastering-temperature-and-top-p-in-chatgpt-api/172683. [34] PIQA: Reasoning about Physical Commonsense in Natural Language. https://arxiv.org/abs/1911.11641. [35] SocialIQA: Commonsense Reasoning about Social Interactions. https://arxiv.org/abs/1904.09728. [36] HellaSwag: Can a Machine Really Finish Your Sentence? https://arxiv.org/abs/1905.07830. [37] WinoGrande: An Adversarial Winograd Schema Challenge at Scale. https://arxiv.org/abs/1907.10641. [38] Can a Suit of Armor Conduct Electricity? A New Dataset for Open Book Question Answering. https://arxiv.org/abs/1809.02789. [39] CommonsenseQA: A Question Answering Challenge Targeting Commonsense Knowledge. https://arxiv.org/abs/1811.00937. [40] TriviaQA: A Large Scale Dataset for Reading Comprehension and Question Answering. https://nlp.cs.washington.edu/triviaqa/. [41] The Natural Questions Dataset. https://ai.google.com/research/NaturalQuestions. [42] SQuAD: 100,000+ Questions for Machine Comprehension of Text. https://arxiv.org/abs/1606.05250. [43] QuAC dataset. https://quac.ai/. [44] BoolQ: Exploring the Surprising Difficulty of Natural Yes/No Questions. https://arxiv.org/abs/1905.10044. [45] GSM8K dataset. https://github.com/openai/grade-school-math. [46] MATH dataset. https://github.com/hendrycks/math/. [47] HumanEval dataset. https://github.com/openai/human-eval. [48] MBPP dataset. https://github.com/google-research/google-research/tree/master/mbpp. [49] Measuring Massive Multitask Language Understanding. https://arxiv.org/abs/2009.03300. [50] Measuring Massive Multilingual Multitask Language Understanding. https://huggingface.co/datasets/openai/MMMLU. [51] AGIEval: A Human-Centric Benchmark for Evaluating Foundation Models. https://arxiv.org/abs/2304.06364. [52] RealToxicityPrompts: Evaluating Neural Toxic Degeneration in Language Models. https://arxiv.org/abs/2009.11462. [53] Perspective API. https://perspectiveapi.com/. [54] ToxiGen: A Large-Scale Machine-Generated Dataset for Adversarial and Implicit Hate Speech Detection. https://arxiv.org/abs/2203.09509. [55] HateCheck: Functional Tests for Hate Speech Detection Models. https://arxiv.org/abs/2012.15606. [56] CrowS-Pairs: A Challenge Dataset for Measuring Social Biases in Masked Language Models. https://arxiv.org/abs/2010.00133. [57] BBQ: A Hand-Built Bias Benchmark for Question Answering. https://arxiv.org/abs/2110.08193. [58] BOLD: Dataset and Metrics for Measuring Biases in Open-Ended Language Generation. https://arxiv.org/abs/2101.11718. [59] TruthfulQA: Measuring How Models Mimic Human Falsehoods. https://arxiv.org/abs/2109.07958. [60] Question Answering for Privacy Policies: Combining Computational and Legal Perspectives. https://arxiv.org/abs/1911.00841. [61] AdvGLUE Benchmark. https://adversarialglue.github.io/. [62] Is BERT Really Robust? A Strong Baseline for Natural Language Attack on Text Classification and Entailment. https://arxiv.org/abs/1907.11932. [63] AdvBench. https://github.com/llm-attacks/llm-attacks. [64] Chatbot Arena leaderboard. https://lmarena.ai/leaderboard. [65] A Survey on Recent Advances in LLM-Based Multi-turn Dialogue Systems. https://arxiv.org/abs/2402.18013. [66] Better & Faster Large Language Models via Multi-token Prediction. https://arxiv.org/abs/2404.19737. [67] Gemini 1.5: Unlocking multimodal understanding across millions of tokens of context. https://arxiv.org/abs/2403.05530. [68] HyperAttention: Long-context Attention in Near-Linear Time. https://arxiv.org/abs/2310.05869. [69] MM-LLMs: Recent Advances in MultiModal Large Language Models. https://arxiv.org/abs/2401.13601. [70] Multimodality and Large Multimodal Models. https://huyenchip.com/2023/10/10/multimodal.html. [71] What is Retrieval-Augmented Generation? https://cloud.google.com/use-cases/retrieval-augmented-generation. [72] How to Customize an LLM: A Deep Dive to Tailoring an LLM for Your Business. https://techcommunity.microsoft.com/t5/ai-machine-learning-blog/how-to-customize-an-llm-a-deep-dive-to-tailoring-an-llm-for-your/ba-p/4110204. [73] Llama 2: Open Foundation and Fine-Tuned Chat Models. https://arxiv.org/abs/2307.09288. [74] Red Teaming Language Models to Reduce Harms: Methods, Scaling Behaviors, and Lessons Learned. https://arxiv.org/abs/2209.07858. [75] Introducing superalignment. https://openai.com/index/introducing-superalignment/. [76] Language Models are Few-Shot Learners. https://arxiv.org/abs/2005.14165. [77] GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. https://arxiv.org/abs/2305.13245. [78] Chain-of-Thought Prompting Elicits Reasoning in Large Language Models. https://arxiv.org/abs/2201.11903. [79] Efficiently Scaling Transformer Inference. https://arxiv.org/abs/2211.05102. [80] Prover-Verifier Games improve legibility of language model outputs. https://openai.com/index/prover-verifier-games-improve-legibility/.
Генерация подписей к изображениям
Введение
Генерация подписей к изображениям — это процесс создания текста, описывающего изображение. Сгенерированный текст, также известный как подпись, должен точно отражать содержимое изображения.
Генерация подписей к изображениям имеет множество применений. Например, на платформах социальных сетей она автоматически предлагает подписи к изображениям, экономя время создателей контента. В интернет-торговле она генерирует подписи для изображений товаров, тем самым улучшая опыт совершения покупок.
Помимо пользовательских приложений, генерация подписей также используется в системах, работающих за кулисами. Например, в модерации NSFW-контента (нежелательного для просмотра на работе) системы генерации подписей могут создавать описательные подписи, помогающие выявлять и помечать неуместный или явный контент, предоставляя текстовые интерпретации изображений. Кроме того, генерация подписей может решить проблему холодного старта в рекомендательных системах, которая возникает, когда системе не хватает данных о новых пользователях или элементах для точных рекомендаций. Генерируя описательные подписи, система получает текстовую информацию, которая помогает классифицировать и рекомендовать новые элементы на основе их содержимого.
В этой главе мы проектируем систему машинного обучения (ML), которая генерирует описательные подписи к изображениям.
Уточнение требований
Вот типичный диалог между кандидатом и интервьюером:
Кандидат: Существуют различные типы изображений, включая общие повседневные изображения и предметно-специфические изображения, такие как медицинские снимки или технические диаграммы. Могу ли я сосредоточиться на общих повседневных изображениях? Интервьюер: Конечно.
Кандидат: Есть ли конкретные приложения или варианты использования, на которые мы ориентируемся с этой системой? Интервьюер: Мы ориентируемся на предложение имён дизайнерам при загрузке их ресурсов.
Кандидат: Поскольку генератор подписей будет использоваться для предложения имён ресурсам, подписи не должны быть слишком длинными и детальными. Это справедливое предположение? Интервьюер: Логично. Подписи должны быть краткими, но описательными и чёткими.
Кандидат: Должна ли система поддерживать несколько языков, или она будет ориентирована только на английский? Интервьюер: Давайте сосредоточимся только на английском.
Кандидат: Каков предполагаемый размер и разнообразие набора данных? Интервьюер: У нас есть доступ к большому набору данных с 400 миллионами пар изображение–подпись, ориентированных на повседневные изображения.
Кандидат: Состоит ли набор данных исключительно из подписей на английском? Интервьюер: Набор данных не предобработан. Там могут быть подписи на разных языках, и некоторые подписи могут быть зашумлёнными или неточными. Кроме того, подписи для некоторых изображений могут отсутствовать.
Кандидат: Требуется ли генерация подписей в реальном времени? Интервьюер: Система должна генерировать подпись быстро, хотя скорость реального времени не обязательна. Задержка в 1–2 секунды допустима.
Кандидат: Как система должна обрабатывать изображения с неоднозначным содержимым или нечётким фокусом? Интервьюер: В таких случаях система должна пропустить предложение подписи.
Кандидат: Я предполагаю, что система должна избегать генерации предвзятых подписей или подписей с оскорбительными словами. Это справедливое предположение? Интервьюер: Отличное замечание. Да, крайне важно обеспечить справедливость и безопасность нашей системы для пользователей.
Кандидат: Каковы типичные размеры изображений? Очень маленькие изображения могут быть нечёткими, что приводит к неправильным подписям. Интервьюер: Давайте предположим, что система предлагает имена только для изображений с минимальным разрешением 256 x 256 пикселей.
Формулировка задачи как ML-задачи
Определение входных и выходных данных системы
Входными данными для системы генерации подписей является изображение. Это изображение обрабатывается моделью для создания описательной подписи. Выходными данными, следовательно, является текст, точно описывающий содержимое изображения.
Выбор подходящего ML-подхода
Задача генерации подписей к изображениям представляет уникальный вызов: ML-модель требует визуального понимания для обработки входного изображения, языкового понимания для генерации подписи и способности устранить разрыв между визуальными и текстовыми модальностями. Это требует разработки мультимодальной системы.
Распространённым подходом к построению мультимодальных систем является использование фреймворка encoder-decoder. Подобно языковому переводу — где мы использовали архитектуру encoder-decoder — мы рассматриваем изображение как новый «язык» в данном контексте. Конкретно, мы используем два основных компонента, каждый из которых обрабатывает одну модальность:
- Encoder изображений
- Текстовый decoder
Encoder изображений
Encoder изображений отвечает за понимание визуального содержимого изображения и кодирование изображения в представление меньшей размерности.
Текстовый decoder
Текстовый decoder использует закодированную визуальную информацию от encoder'а изображений для генерации описательной подписи.
Мы подробно рассмотрим архитектуру этих компонентов в разделе разработки модели. Важно отметить, что существуют различные подходы к решению задачи генерации подписей к изображениям. Хотя мы сосредоточимся на фреймворке encoder-decoder, альтернативные модели, такие как BLIP-2 [1], BLIP-3 [2] и InternVL [3], предлагают различные техники и архитектуры для генерации подписей. Если вас интересуют эти другие методы, вы можете обратиться к [1] [2] [3] для более широкого понимания области генерации подписей к изображениям.
Подготовка данных
В этом разделе мы подготавливаем набор данных для обучения нашей системы генерации подписей к изображениям.
Набор данных состоит из 400 миллионов пар изображений и подписей. Однако не все изображения или подписи подходят для обучения. Рассмотрим подготовку данных для подписей и изображений отдельно.
Подготовка подписей
Необработанные подписи часто зашумлены и не имеют формата, пригодного для использования ML-моделью. В процессе подготовки подписей мы удаляем неподходящие подписи и обеспечиваем согласованность и токенизацию оставшихся. В частности, мы выполняем следующие шаги:
- Удалить пары с подписями не на английском языке: Мы удаляем пары изображение–подпись, где подпись не на английском языке, поскольку эта модель будет ориентирована на английский.
- Удалить дубликаты изображений или подписей: Для обеспечения разнообразия и качества обучающих данных мы исключаем дублирующиеся изображения и подписи. Дублирующиеся изображения выявляются с помощью методов перцептивного хеширования или моделей схожести изображений (например, CLIP image encoder), а дублирующиеся подписи обнаруживаются по точному совпадению или проверкам семантического сходства (например, CLIP text encoder). Удаление дубликатов предотвращает переобучение модели на избыточных данных и помогает ей изучить более широкий спектр связей между изображениями и текстом.
- Удалить нерелевантные подписи: Мы используем предобученную vision–language модель (например, CLIP) для оценки релевантности между изображениями и соответствующими подписями. Более высокая оценка обычно указывает на большую семантическую релевантность между изображением и текстом. Мы удаляем пары с оценками ниже определённого порога, например 0,25. Это гарантирует, что наша модель учится на высококачественных, релевантных парах. Для получения дополнительной информации о том, как CLIP оценивает релевантность между текстом и изображениями, обратитесь к Главе 9.
- Суммаризировать длинные подписи: Подписи часто бывают длинными и подробными. Обучение модели с такими подписями приводит к генерации аналогично длинных подписей, что не соответствует нашему варианту использования. Для решения этой проблемы мы суммаризируем подписи с помощью большой языковой модели, такой как Llama [4], для создания кратких, лаконичных описаний, отвечающих нашим требованиям.
- Нормализовать подписи: Мы применяем стандартные методы нормализации текста, включая приведение к нижнему регистру и обрезку пробелов, для поддержания согласованности между подписями.
- Токенизировать подписи: Мы используем алгоритм токенизации на уровне подслов, такой как Byte-Pair Encoding (BPE) [5], для токенизации подписей в последовательность идентификаторов. Для подробного обзора методов токенизации текста и алгоритма BPE обратитесь к Главе 2 и Главе 3.
Подготовка изображений
Как и в случае с подписями, не все изображения полезны. Мы удаляем изображения, которые могут навредить обучению, и обеспечиваем согласованность и пригодность оставшихся изображений для обучения модели. В частности, мы выполняем следующие шаги:
- Удалить изображения с низким разрешением: Мы удаляем пары изображение–подпись, в которых разрешение изображения меньше 256×\times×256, поскольку такие изображения с низким разрешением могут не предоставлять достаточно деталей для точной генерации подписей.
- Нормализовать изображения: Мы масштабируем значения пикселей до нормализованного диапазона, например от 0 до 1. Эта нормализация делает процесс обучения более стабильным.
- Удалить изображения низкого качества: Для поддержания высокого качества обучающих данных мы фильтруем изображения с такими условиями, как размытость, переэкспозиция, недоэкспозиция или другие дефекты, снижающие визуальную чёткость. Методы оценки качества изображений, такие как LAION Aesthetics Predictor [6], помогают выявлять и удалять некачественные изображения, оценивая их по таким факторам, как резкость, контраст и освещение.
- Скорректировать размеры изображений: Изображения обычно имеют разнообразные размеры и соотношения сторон. Мы приводим все изображения к единому размеру. Это критически важно, поскольку ML-модели требуют входных данных фиксированного размера во время обучения. При приведении размеров изображений к единому размеру важно сохранять их исходные соотношения сторон. Для этого мы часто следуем двум шагам:
- Изменение размера: Сначала мы изменяем размер изображения так, чтобы меньшее измерение соответствовало целевому размеру. Например, если наш целевой размер 256×\times×256, а исходное изображение 512×\times×768, мы изменяем его размер до 256×\times×384.
- Центральная обрезка: Затем мы выполняем центральную обрезку изображения до целевых размеров. Из нашего предыдущего примера мы обрезаем изображение 256×\times×384 до 256×\times×256.
Этот двухэтапный метод обеспечивает сохранение соотношений сторон изображений и соответствие требуемому размеру для нашей ML-модели.
Разработка модели
Архитектура
Мы сформулировали генерацию подписей к изображениям как мультимодальную задачу генерации языка, где encoder изображений обрабатывает входное изображение, а текстовый decoder генерирует описательную подпись. В этом разделе мы исследуем архитектуру encoder'а изображений и текстового decoder'а.
Encoder изображений
Encoder изображений отвечает за обработку изображения и кодирование содержащейся в нём информации.
Выход encoder'а играет ключевую роль в определении качества и специфичности сгенерированных подписей. Выход encoder'а может быть либо одним token'ом, представляющим всё изображение как один вектор признаков, либо последовательностью token'ов, где каждый token соответствует конкретной области или аспекту изображения. Выбор между этими двумя подходами имеет существенные последствия для того, насколько эффективно система захватывает и представляет визуальное содержимое, и исследования изучили оба варианта для понимания их сильных и слабых сторон.
Когда encoder выдаёт один token в качестве выхода, он эффективно сжимает всё изображение в один вектор. Этот вектор служит кратким описанием изображения, инкапсулируя его глобальные признаки и общий контекст. Основное преимущество этого подхода заключается в его простоте: архитектура остаётся прямолинейной, со сниженной вычислительной сложностью и меньшими требованиями к ресурсам. Один вектор акцентирует внимание на общем содержимом изображения, что может быть особенно полезно для генерации кратких и высокоуровневых подписей, отражающих общую суть сцены. Однако этот подход также имеет существенные недостатки. Сжатие всей визуальной информации в один вектор часто означает потерю локальных деталей и специфических нюансов, которые критически важны для генерации описательных и контекстно-богатых подписей. В результате подписи, сгенерированные из одного token'а, могут быть более обобщёнными и могут не справляться со сложными изображениями, требующими детального представления.
С другой стороны, генерация последовательности token'ов от encoder'а позволяет системе захватить более детальный вид изображения. Каждый token в последовательности соответствует отдельной части или патчу изображения, что приводит к более богатому и всестороннему представлению, включающему как глобальные, так и локальные признаки. Этот подход особенно хорошо согласуется с механизмом attention, который является краеугольным камнем современных генеративных моделей, таких как Transformer. Механизм attention лучше всего работает с последовательными входными данными, поскольку он позволяет decoder'у динамически фокусироваться на разных областях изображения во время генерации подписи. Эта способность избирательно обращать внимание на различные части изображения приводит к более точным, релевантным и детальным подписям. Используя последовательность token'ов, модель может генерировать подписи, которые не только более описательны, но и лучше соответствуют конкретным объектам, действиям и контексту, присутствующим на изображении.
Архитектуры encoder'а изображений можно разделить на следующие:
- На основе CNN
- На основе Transformer
На основе CNN
Сверточные нейронные сети (CNN) традиционно используются для задач кодирования изображений. CNN отлично справляются с захватом пространственных иерархий в изображениях с использованием сверточных фильтров. Эти фильтры обнаруживают паттерны, такие как края, текстуры и объекты в разных масштабах.
Encoder'ы на основе CNN обрабатывают входное изображение и выдают сетку векторов признаков. Например, как показано на Рисунке 7, входное изображение проходит через CNN, производя вектор признаков размером 3 x 3 x c. Здесь c представляет размер канала, который зависит от архитектуры CNN. В то время как CNN производит выход 3 x 3 x c, Transformer в текстовом decoder'е нуждается в последовательности признаков (т.е. 9 x c). Для этого мы используем операцию сплющивания или изменения формы, которая реорганизует признаки из каждой из девяти позиций в сетке 3 x 3 в последовательный формат.
На основе Transformer
Модели Transformer, изначально разработанные для обработки естественного языка, в последнее время были адаптированы для кодирования изображений с значительным успехом. В этой архитектуре Transformer анализирует изображения, извлекает признаки и кодирует их в последовательность embedding'ов. Конкретно, encoder изображений на основе Transformer состоит из:
- Patchify
- Позиционное кодирование
- Transformer
Patchify
Поскольку Transformer'ы работают с последовательностями, изображение сначала должно быть преобразовано в последовательность. Этот процесс включает три шага:
- Разделить изображение на патчи фиксированного размера
- Сплющить каждый патч
- Линейно спроецировать каждый патч
Например, входное изображение 256 x 256 делится на патчи 64 x 64. Эти патчи сплющиваются в векторы размером 4096 и линейно проецируются в embedding-векторы размера c, где c — желаемый размер embedding'а.
Позиционное кодирование
Позиционное кодирование назначает позиционную информацию каждому патчу, указывая, где каждый патч располагался в исходном изображении. Это помогает Transformer'ам понимать позиции в последовательности.
Позиционное кодирование может быть реализовано различными способами. Давайте кратко рассмотрим следующие варианты:
- 1D vs. 2D позиционное кодирование
- Обучаемое vs. фиксированное позиционное кодирование
D vs. 2D позиционное кодирование
1D позиционное кодирование использует функцию, которая отображает целое число (позицию в последовательности) в c-мерный вектор, где c обычно является скрытым измерением Transformer'а. Это обычно используется в текстовых последовательностях, где каждый token получает позиционный вектор на основе своего места. При применении к изображениям 1D позиционное кодирование кодирует позицию каждого патча в сплющенной последовательности, что может не захватывать двумерные пространственные отношения в изображениях.
2D позиционное кодирование, с другой стороны, отображает два целых числа — представляющих позиции строки и столбца в сетке изображения — в c-мерный вектор. Этот метод кодирования более подходит для изображений, так как он сохраняет пространственную структуру.
Обучаемое vs. фиксированное позиционное кодирование
В обучаемом позиционном кодировании модель изучает позиционные кодировки во время обучения. Нейронная сеть отображает позиции (1D или 2D) в c-мерный вектор. В фиксированном подходе позиционные кодировки определяются фиксированной функцией, такой как синус-косинус. Для получения более подробной информации обратитесь к Главе 2.
При выборе между 1D и 2D, а также обучаемым и фиксированным позиционным кодированием часто не существует наилучшего решения. Хотя Vision Transformer (ViT) [7] использует обучаемое 1D позиционное кодирование, на практике мы часто тестируем различные комбинации, чтобы увидеть, какая из них лучше всего работает для конкретной задачи.
Какая архитектура подходит для нашего encoder'а изображений?
CNN эффективны при захвате локальных паттернов в изображениях, но испытывают затруднения с долгосрочными зависимостями между удалёнными областями изображения. Напротив, Transformer'ы захватывают как локальные, так и глобальные отношения в изображении с помощью механизма self-attention. Это позволяет Transformer'ам моделировать сложные зависимости, что делает их идеальными для задач, требующих детального, контекстно-осведомлённого понимания изображений, например для генерации описательных подписей. По этим причинам мы следуем ViT [7] и выбираем архитектуру на основе Transformer в качестве нашего encoder'а изображений.
Текстовый decoder
Текстовый decoder отвечает за генерацию подписи. Как мы видели в предыдущих главах, decoder-only Transformer является стандартным выбором для генерации текста. Входом в decoder-only Transformer является последовательность векторов, соответствующих входному изображению. Его выходом является подпись, генерируемая по одному token'у за раз.
Обучение
Подход к обучению модели генерации подписей к изображениям схож со стратегиями, рассмотренными в предыдущих главах. Мы следуем двухэтапной стратегии обучения:
- Неконтролируемое pre-training
- Контролируемое fine-tuning
. Неконтролируемое pre-training
На этом этапе текстовый decoder — который является decoder-only Transformer — обучается на общих данных. Цель этого этапа — разработать базовую модель, обладающую широким пониманием языковой структуры и способной генерировать связный текст. Эти знания крайне важны для того, чтобы модель хорошо работала при последующем fine-tuning'е на более конкретной задаче, такой как генерация подписей.
Этап pre-training'а является вычислительно затратным. Общепринятой практикой является использование существующих предобученных моделей для пропуска этого этапа и, таким образом, значительного снижения вычислительных затрат. В этой главе мы используем предобученный decoder-only Transformer, такой как GPT-2 [8] или Llama [4].
Аналогично, encoder изображений также может быть получен из предобученных моделей. Вместо обучения encoder'а изображений с нуля мы можем воспользоваться мощными предобученными визуальными моделями, такими как CLIP [9] или ViT [7].
. Контролируемое fine-tuning
На этом этапе мы обучаем как encoder изображений, так и текстовый decoder на 400 миллионах пар изображение–подпись. Encoder изображений улучшает свою способность эффективно кодировать информацию об изображении, а текстовый decoder учится понимать последовательность embedding'ов изображений и генерировать описательную подпись.
ML-цель и функция потерь
Текстовый decoder генерирует подпись по одному token'у за раз. В соответствии с предыдущими главами, мы используем предсказание следующего token'а как нашу ML-цель и применяем кросс-энтропийные потери [10] для управления процессом обучения.
Сэмплирование
Во время сэмплирования token'ы подписи генерируются по одному за раз.
Хотя стохастические методы сэмплирования могут создавать творческие подписи, beam search обеспечивает предсказуемость. Мы используем beam search для нашей системы генерации подписей по следующим причинам:
- Качество: Beam search обычно генерирует подписи более высокого качества, что критически важно для точного описания содержимого изображения.
- Согласованность: Детерминированная природа beam search гарантирует, что модель всегда производит одну и ту же подпись для одного и того же изображения. Эта согласованность крайне важна для генерации подписей к изображениям.
- Связность: Beam search обычно производит связные подписи, что важно для генерации подписей к изображениям. Это позволяет избежать внезапных смен темы или противоречий, таких как «Человек идёт дом» или «Собака читает человека».
Оценка
Метрики офлайн-оценки
Во время офлайн-оценки мы оцениваем производительность обученной модели на валидационном наборе данных. Это достигается путём сравнения сгенерированных подписей с эталонными (т.е. правильными) подписями и измерения их сходства.
Прежде чем рассматривать общие метрики, давайте рассмотрим данные для валидации. Данные для валидации содержат примеры, которые модель не видела во время обучения. Каждый пример включает изображение и набор эталонных подписей. Эти подписи обычно собираются несколькими аннотаторами-людьми, описывающими каждое изображение.
В системах генерации подписей к изображениям обычно имеется несколько эталонных подписей для каждого изображения. Это полезно как для обучения, так и для оценки по следующим причинам:
- Устойчивое обучение: Разные люди описывают одно и то же изображение по-разному. Несколько эталонов позволяют модели изучать различные способы описания изображения. Это приводит к более устойчивой модели, способной описывать изображения более точно.
- Комплексная оценка: Несколько подписей обеспечивают более тщательную оценку производительности модели. Сравнение сгенерированной подписи с несколькими правильными эталонными подписями приводит к более справедливой оценке.
Следующие метрики обычно используются при офлайн-оценке моделей генерации подписей к изображениям:
- BLEU
- ROUGE
- METEOR
- CIDEr
Первые три метрики в списке подробно рассмотрены в Главе 3. В этой главе мы сосредоточимся на CIDEr, который был разработан специально для оценки моделей генерации подписей к изображениям.
CIDEr
CIDEr [11] — популярная метрика для оценки моделей генерации подписей к изображениям. Она использует консенсус для оценки сходства сгенерированной подписи с набором эталонных подписей. CIDEr присваивает более высокие оценки подписям, которые похожи на несколько эталонных подписей, а не только на одну. Для одного примера CIDEr рассчитывается в три шага:
- Представить подписи с помощью Term Frequency–Inverse Document Frequency (TF-IDF)
- Вычислить сходства
- Агрегировать оценки сходства
. Представление подписей с помощью TF-IDF
На первом шаге мы преобразуем сгенерированную подпись и каждую эталонную подпись в числовые представления с помощью TF-IDF. TF-IDF оценивает важность слова для документа, учитывая, как часто оно встречается в этом документе и насколько оно распространено или редко во всём корпусе. Эти оценки важности используются для числового представления предложения. Для получения дополнительной информации о TF-IDF обратитесь к [12][13].
. Вычисление сходства
Затем мы вычисляем сходство между сгенерированной подписью и каждой эталонной подписью. Мы делаем это путём вычисления косинусного сходства между их TF-IDF представлениями.
Более высокое косинусное сходство (т.е. оценка ближе к 1) указывает на большее сходство, тогда как более низкое значение (ближе к 0) указывает на меньшее сходство.
. Агрегация оценок сходства
После получения оценок косинусного сходства между сгенерированной подписью и каждой из эталонных подписей, мы вычисляем среднее значение этих оценок. Это среднее значение отражает общее сходство между сгенерированной подписью и эталонными подписями.
Итоговая оценка CIDEr рассчитывается путём усреднения оценок сходства для всех сгенерированных подписей в валидационном наборе данных. Это обеспечивает единую метрику для оценки общей производительности модели.
Рассмотрим некоторые плюсы и минусы метрики CIDEr.
Плюсы:
- Основан на консенсусе: CIDEr подчёркивает консенсус, вознаграждая подписи, похожие на несколько эталонных подписей. Это приводит к более надёжной оценке производительности модели.
- Чувствителен к важным словам: TF-IDF присваивает больший вес уникальным словам в их представлении. Это гарантирует, что оценка CIDEr отражает важность слов и вознаграждает подписи, использующие эти слова.
- Устойчив к различным вариациям подписей: CIDEr устойчив к различным вариациям генерации, поскольку рассчитывается на основе нескольких эталонных подписей.
Минусы:
- Вычислительно сложный: Вычисление TF-IDF представлений в больших наборах данных может быть вычислительно затратным.
- Чувствителен к качеству эталонных подписей: Качество и разнообразие эталонных подписей влияют на оценку CIDEr. Плохие эталоны могут приводить к вводящим в заблуждение оценкам.
- Штрафует за новые, но точные подписи: CIDEr может штрафовать творческие или новые фразы, которые всё ещё точны, но не присутствуют в эталонном наборе.
- Отсутствие семантического понимания: CIDEr опирается на TF-IDF для измерения сходства между двумя предложениями. Это может не всегда захватывать семантическое сходство, когда подписи текстуально схожи, но семантически различны. Например, «Кофе на столе» и «Стол на кофе» могут иметь схожие TF-IDF представления из-за схожих слов, но они не являются семантически схожими.
Метрики онлайн-оценки
Метрики онлайн-оценки важны для оценки производительности ML-систем. Однако они часто не являются основным фокусом в системах генерации подписей к изображениям по двум основным причинам. Во-первых, системы генерации подписей к изображениям обычно являются частью более крупной системы, что затрудняет сбор данных о взаимодействии пользователей. Во-вторых, получение отзывов от пользователей является сложной задачей. В отличие от задач, где мы можем легко измерить удовлетворённость пользователей, оценка качества подписей к изображениям требует субъективного суждения, которое по определению варьируется между пользователями. Например, подпись может быть приемлемой для одного пользователя, но не для другого, в зависимости от их личной интерпретации изображения.
В заключение, стандартные офлайн-метрики остаются основным методом оценки нашей системы генерации подписей к изображениям. Для немногих вариантов использования, где генерация подписей напрямую влияет на пользовательский опыт, метрики вовлечённости и обратная связь от пользователей могут предоставить ценные сведения о производительности системы.
Общая архитектура ML-системы
Построение системы генерации подписей к изображениям — это больше, чем просто обучение модели. Это требует совместной работы различных компонентов. В этом разделе мы обсуждаем следующие ключевые компоненты, необходимые для построения системы генерации подписей к изображениям:
- Предобработка изображений
- Генератор подписей
- Постобработка
Давайте кратко рассмотрим каждый компонент и поймём его роль.
Предобработка изображений
Предобработка изображений — это начальный шаг, который подготавливает входное изображение для обученной модели. Это включает изменение размера изображений до стандартного размера, преобразование их в согласованный формат и стандартизацию значений пикселей. Этот шаг обеспечивает согласованность изображений с тем, что модель ожидает в качестве входных данных.
Генератор подписей
Генератор подписей — это основной компонент, который создаёт подписи на основе подготовленного изображения. Этот компонент взаимодействует с обученной моделью и использует beam search для генерации связной подписи. Если совокупная вероятность сгенерированной подписи падает ниже заранее определённого порога достоверности, предложение имени отключается; в противном случае подпись передаётся компоненту постобработки. Это гарантирует, что система избегает создания нерелевантных подписей для неоднозначных изображений.
Постобработка
Компонент постобработки выявляет предвзятые термины или фразы в подписи и заменяет их нейтральными альтернативами. Это обеспечивает справедливость и инклюзивность в сгенерированных подписях. Кроме того, он проверяет наличие оскорбительных слов и отключает сервис предложения имён, если они обнаружены.
Дополнительные темы для обсуждения
Если собеседование завершится раньше времени, вы можете поднять следующие темы:
- Расширение генератора подписей для поддержки других задач, таких как визуальный ответ на вопросы (VQA) [14].
- Адаптация моделей для создания подписей к изображениям из различных областей [15].
- Генерация подписей на нескольких языках с использованием многоязычных наборов данных и кросс-лингвального transfer learning [16].
- Методы оптимизации для генерации подписей на периферийных устройствах [17].
- Генерация и ранжирование нескольких правдоподобных подписей на основе релевантности [18].
- Подробности методов BLIP-2 и BLIP-3 и дополнительные функции потерь, используемые для улучшения генерации подписей [1] [2].
Резюме
Справочные материалы
[1] BLIP-2: Bootstrapping Language-Image Pre-training with Frozen Image Encoders and Large Language Models. https://arxiv.org/abs/2301.12597. [2] xGen-MM (BLIP-3): A Family of Open Large Multimodal Models. https://www.arxiv.org/abs/2408.08872. [3] InternVL: Scaling up Vision Foundation Models and Aligning for Generic Visual-Linguistic Tasks. https://arxiv.org/abs/2312.14238. [4] Meta's Llama. https://llama.meta.com/. [5] Byte-pair encoding tokenization. https://huggingface.co/learn/nlp-course/en/chapter6/5. [6] LAION-5B: An open large-scale dataset for training next generation image-text models. https://arxiv.org/abs/2210.08402. [7] An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. https://arxiv.org/abs/2010.11929. [8] Language Models are Unsupervised Multitask Learners. https://cdn.openai.com/better-language-models/language_models_are_unsupervised_multitask_learners.pdf. [9] Learning Transferable Visual Models From Natural Language Supervision. https://arxiv.org/abs/2103.00020. [10] Cross-entropy. https://en.wikipedia.org/wiki/Cross-entropy. [11] CIDEr: Consensus-based Image Description Evaluation. https://arxiv.org/abs/1411.5726. [12] TF-IDF introduction. https://web.stanford.edu/class/cs276/19handouts/lecture6-tfidf-1per.pdf. [13] TF-IDF. https://en.wikipedia.org/wiki/Tf%E2%80%93idf. [14] Visual question answering introduction. https://huggingface.co/tasks/visual-question-answering. [15] Cross-Domain Image Captioning with Discriminative Finetuning. https://arxiv.org/abs/2304.01662. [16] Crossmodal-3600 — Multilingual Reference Captions for Geographically Diverse Images. https://research.google/blog/crossmodal-3600-multilingual-reference-captions-for-geographically-diverse-images/. [17] Efficient Image Captioning for Edge Devices. https://arxiv.org/abs/2212.08985. [18] Ensemble model using an image captioning and ranking example. https://cloud.google.com/dataflow/docs/notebooks/run_inference_multi_model.
Генерация с дополнением извлечением (RAG)
Введение
В Главе 4 мы разработали чатбот, способный отвечать на вопросы из открытой области. Однако многие приложения нуждаются в доступе к дополнительной информации, такой как корпоративные базы данных (например, внутренняя документация), данные в реальном времени (например, спортивные результаты) или файлы, предоставленные пользователем (например, загруженные PDF).
Предоставление чатботам доступа к этой информации улучшает точность и релевантность их ответов, особенно для задач, основанных на фактах или специализированных задач. Реальным примером такой системы является Perplexity.ai [1] — поисковая система на основе ИИ, которая использует информацию из веба для ответа на запросы пользователей.
В этой главе мы строим систему, аналогичную ChatPDF [2], которая отвечает на вопросы сотрудников, используя внутренние корпоративные документы. Вместо чтения FAQ сотрудники могут напрямую спросить чатбот и получить ответы на основе этих документов.
Уточнение требований
Вот типичное взаимодействие между кандидатом и интервьюером:
Кандидат: Из чего состоит внешняя база знаний? Меняется ли она со временем? Интервьюер: База знаний включает Wiki-страницы компании и корпоративный форум в стиле «Stack Overflow». Документация меняется, но медленнее по сравнению с обновлениями в реальном времени.
Кандидат: Содержат ли Wiki-страницы и форумы текст, изображения и другие модальности? Интервьюер: Предположим, что каждая страница в формате PDF и содержит текст, таблицы и диаграммы. Для простоты другие модальности можно не учитывать.
Кандидат: Следуют ли страницы фиксированному формату или шаблону? Интервьюер: Нет, форматы варьируются. Некоторые двухколоночные, некоторые одноколоночные, а другие смешанные.
Кандидат: Сколько всего страниц? Интервьюер: У нас около 5 миллионов страниц.
Кандидат: Необходимо ли системе включать ссылки на документы? Интервьюер: Да.
Кандидат: Должна ли система отвечать в реальном времени? Интервьюер: Пользователи могут допустить небольшую задержку в несколько секунд.
Кандидат: Должна ли система поддерживать несколько языков? Интервьюер: Для простоты давайте ограничимся английским.
Кандидат: Должна ли система поддерживать обратную связь от пользователей или уточняющие вопросы? Интервьюер: Изначально нет. Однако ваш дизайн должен быть достаточно гибким, чтобы в будущем добавить поддержку циклов обратной связи или уточняющих вопросов.
Кандидат: Каков ожидаемый рост количества документов? Интервьюер: Ожидается, что база документов будет расти на двадцать процентов ежегодно.
Кандидат: Нужно ли решать вопросы безопасности, такие как предотвращение вредоносных, предвзятых или вводящих в заблуждение ответов? Интервьюер: Безопасность важна, но давайте приоритизируем обработку данных, архитектуру и эффективность производительности.
Формулировка задачи как задачи ML
Определение входа и выхода системы
Вход системы ChatPDF — это текстовый промпт, предоставленный пользователем. Модель обрабатывает этот промпт вместе с постоянно обновляемой базой данных документов, содержащей текст и изображения. Выход — это текстовый ответ, точно отвечающий на запрос пользователя.
Выбор подходящего подхода ML
Учитывая характер задачи, большие языковые модели (LLM) хорошо подходят для генерации текста и часто являются выбором по умолчанию. Однако LLM общего назначения могут испытывать трудности с конкретными доменами и, следовательно, могут нуждаться в настройке для работы с внешними источниками данных. Чтобы LLM могла отвечать на запросы на основе данных конкретной компании, существуют три основных подхода:
- Fine-tuning
- Prompt engineering
- Retrieval-augmented generation (RAG)
Давайте подробно рассмотрим каждый из них и обсудим их компромиссы.
Fine-tuning
В этом подходе предобученная LLM общего назначения дообучается на данных, специфичных для компании, таких как внутренние документы. Обновляя свои веса, LLM адаптируется для лучшего понимания уникальной терминологии, процессов и FAQ компании. В Главе 10 будут рассмотрены продвинутые техники fine-tuning, такие как LoRA [3], для адаптации больших моделей к конкретным данным.
Преимущества:
- Настраиваемость: Fine-tuning позволяет модели генерировать ответы, адаптированные к конкретным доменам.
- Повышенная точность: Путём дообучения модели на специализированных данных она становится более точной и лучше справляется с нишевыми темами.
Недостатки:
- Вычислительно затратно: Обновление всех параметров модели требует значительных вычислительных ресурсов, что может быть дорого.
- Частое переобучение: Этот подход требует частого дообучения для непрерывного включения актуальных данных в модель.
- Требует технической экспертизы: Этот подход требует понимания принципов ML и архитектур языковых моделей, что может быть барьером для тех, кто не имеет специализированных знаний.
- Обширные требования к данным: Fine-tuning требует существенного, высококачественного набора данных, который может быть сложно и долго собирать.
- Отсутствие ссылок: Дообученные модели обычно не могут предоставить ссылки на свои ответы, что затрудняет проверку или отслеживание информации до её источника.
Prompt Engineering
Prompt engineering направляет LLM общего назначения на создание конкретных ответов через тщательно спроектированные промпты. В отличие от fine-tuning, этот метод оставляет базовую LLM неизменной и включает релевантную информацию, такую как данные компании или инструкции, непосредственно в промпты для управления поведением модели. Например, промпт может включать информацию, такую как краткое изложение политик компании, как показано на Рисунке 4. Позже в этой главе мы рассмотрим более продвинутые техники prompt engineering, такие как few-shot и chain-of-thought промптинг.
Преимущества:
- Простота использования: Prompt engineering прост в использовании и не требует технических навыков, что делает его подходящим для широкого круга пользователей.
- Экономичность: Используя предобученную LLM, промптинг несёт минимальные вычислительные затраты по сравнению с fine-tuning.
- Гибкость: Промпты можно легко модифицировать для экспериментов с различными выходами без необходимости переобучения модели.
Недостатки:
- Непоследовательность: Качество и релевантность ответов могут сильно варьироваться в зависимости от формулировки промпта.
- Ограниченная настраиваемость: Возможность адаптации ответов ограничена эффективностью и креативностью дизайна промпта. Prompt engineering не обладает глубиной настройки, которую предоставляет fine-tuning.
- Ограниченность существующими знаниями LLM: Выходы ограничены информацией, на которой LLM была изначально обучена, что делает её менее эффективной для высокоспециализированных доменов или предоставления ответов на основе самой актуальной информации.
RAG
RAG — это продвинутый метод, который объединяет возможности LLM общего назначения с системой извлечения в реальном времени. Вместо того чтобы полагаться исключительно на предобученные знания LLM, RAG извлекает релевантную информацию из внешних источников, таких как внутренние документы компании, и подаёт её в LLM во время инференса. Этот подход гарантирует, что LLM генерирует ответы, которые одновременно релевантны и точны на основе доступной информации.
Система RAG, как показано на Рисунке 5, имеет два компонента:
- Извлечение: Компонент извлечения берёт исходный промпт пользователя, находит наиболее релевантную информацию из внешних источников и возвращает её в качестве контекста.
- Генерация: Как правило, LLM общего назначения использует промпт пользователя и извлечённую информацию для генерации ответа.
Преимущества:
- Доступ к самой актуальной информации: RAG может предоставлять актуальные ответы, извлекая данные из внешних источников, тем самым улучшая релевантность и точность ответов.
- Контекстуальная релевантность: Извлекая информацию из внешних источников, RAG может добавить контекст к ответам модели, делая ответы более детальными и релевантными.
Недостатки:
- Сложность реализации: Реализация RAG может быть технически сложной, поскольку требует слаженной работы двух компонентов (извлечение и генерация).
- Зависимость от качества извлечения: Качество ответов сильно зависит от релевантности и точности извлечённой информации, что может влиять на общую производительность системы.
Какой подход более подходит для ChatPDF?
Fine-tuning позволяет LLM генерировать более специализированные ответы, но вычислительно затратен и не ссылается на оригинальные документы, что делает его непригодным для наших нужд. Хотя prompt engineering предоставляет простой и гибкий способ направлять LLM общего назначения без fine-tuning, он не масштабируется. Это связано с тем, что включение информации из всех внешних источников в промпт обычно превышает контекстное окно LLM.
RAG предлагает сбалансированное решение с точки зрения простоты настройки, стоимости и масштабируемости, что делает его идеальным для работы с большими, развивающимися наборами данных и предоставления актуальной информации. Этот подход особенно эффективен для внутренних чатботов в корпоративных средах. Поэтому мы выбираем RAG для построения нашей системы ChatPDF. В разделе разработки модели мы углубимся в prompt engineering и обсудим, как мы комбинируем его с RAG для дальнейшего улучшения системы.
Подготовка данных
Производительность системы RAG зависит от качества базы знаний и способа её индексации. Когда база знаний получена из веб-сайтов, следует применять стратегии очистки данных, такие как удаление неуместного контента или анонимизация конфиденциальной информации, как обсуждалось в Главе 4.
В этом разделе мы сосредоточимся на подготовке данных из коллекции PDF-страниц. Это включает трёхэтапный процесс:
- Парсинг документов
- Разбиение документов на чанки
- Индексация
Парсинг документов
PDF — один из наиболее широко используемых форматов документов. Важно правильно извлечь их содержимое для подготовки данных к обучению LLM и для обеспечения того, чтобы LLM могла корректно отвечать на вопросы на основе содержимого PDF.
Парсинг PDF означает преобразование его текста, изображений и других элементов в структурированный формат, который языковая модель может понять. Существует два основных подхода к парсингу PDF:
- Парсер на основе правил
- Парсер на основе ИИ
Парсер документов на основе правил
Подход на основе правил опирается на предопределённые правила и паттерны, основанные на разметке и структуре документа. Он пытается «вычислить» разметку и соответственно извлечь контент, что делает его простым в реализации, когда формат документа последователен и предсказуем.
Однако методы на основе правил с трудом справляются с широким спектром типов и форматов PDF, поскольку PDF могут значительно различаться по дизайну. Жёсткая природа этого метода означает, что если документ не соответствует ожидаемому формату, это может привести к ошибкам при извлечении контента. Это делает парсинг на основе правил менее полезным при работе с различными или сложными макетами документов.
Парсер документов на основе ИИ
Методы на основе ИИ используют другой подход. Они применяют продвинутые техники, такие как обнаружение объектов и OCR (Optical Character Recognition) [4], для идентификации и извлечения различных элементов из документа, например, текста, таблиц и диаграмм. Эти методы могут обрабатывать широкий спектр макетов документов, что делает их более подходящими для работы со сложными документами.
Существуют различные инструменты для парсинга документов на основе ИИ. Например, Dedoc [5] поддерживает парсинг широкого спектра форматов документов и стандартизацию контента в единую структуру. Аналогично, Layout-Parser [6] использует высокоточные модели для точного обнаружения различных частей документа, хотя размер этих моделей может замедлять процесс. Чтобы лучше понять парсеры документов на основе ИИ, давайте подробнее рассмотрим, как работает Layout-Parser.
Layout-Parser принимает изображение документа на вход и генерирует структурированный выход, выполняя следующие шаги:
- Обнаружение макета: Парсер использует продвинутые модели обнаружения объектов для обнаружения и генерации прямоугольных рамок вокруг различных областей контента. Эти области могут включать такие элементы, как абзацы, таблицы, изображения или заголовки.
- Извлечение текста: Контент внутри каждой прямоугольной рамки обрабатывается с помощью OCR для извлечения текста. Координаты ограничивающей рамки обеспечивают распознавание текста в правильном порядке и формате, сохраняя исходную структуру документа.
- Генерация структурированного выхода: Парсер создаёт структурированный выход, содержащий два типа данных: Текстовые блоки: Включают координаты блока, извлечённый текст, порядок чтения и метаинформацию. Нетекстовые блоки: Включают координаты фигур или изображений.
Несколько онлайн-сервисов предоставляют услуги парсинга документов, например, Google Cloud Document AI [7] и PDF.co [8]. Эти сервисы позволяют пользователям загружать свои документы и получать их парсинг без необходимости настраивать и поддерживать систему парсинга самостоятельно.
Разбиение документов на чанки
После того как мы идентифицировали блоки текста, изображений или таблиц в документе, следующий шаг — индексировать их в поисковую базу данных. Для длинных текстовых блоков, например из отчётов или книг, индексация всего контента как единого элемента неэффективна. Это связано с тем, что embedding-вектор, представляющий целую книгу или отчёт, может захватить общий контекст, но упустить важные детали, что может привести к менее точным или неполным результатам поиска. Кроме того, если мы извлечём всю книгу или отчёт, это превысит лимит токенов большинства моделей, например, лимит в 128K токенов для модели GPT-4o.1
Разбиение документов на чанки решает эти проблемы, разбивая текст на более мелкие, управляемые части или «чанки». Чанкинг помогает улучшить качество и точность извлечения и гарантирует, что каждый чанк умещается в пределах лимита ввода модели.
Некоторые распространённые стратегии чанкинга:
- Чанкинг на основе длины: Этот простой подход разбивает текст на чанки заданной длины. Хотя его легко реализовать, он иногда может разделять предложения или логические разделы посередине, приводя к фрагментированным или менее осмысленным чанкам. Такие инструменты, как LangChain [9], предоставляют разделители текста, такие как CharacterTextSplitter и RecursiveCharacterTextSplitter, которые позволяют настраивать размеры чанков и параметры перекрытия. Эти разделители могут обрабатывать различные разделители и помогают поддерживать связность между чанками.
- Чанкинг на основе регулярных выражений: Этот подход использует регулярные выражения для разделения текста на основе конкретных знаков препинания, таких как точки, вопросительные знаки или восклицательные знаки. Он позволяет лучше разделять текст на уровне предложений, сохраняя логические разрывы, хотя ему может не хватать более глубокого семантического понимания текста.
- Разделители HTML, markdown или кода: Для документов в структурированных форматах, таких как HTML или Markdown, используются специализированные разделители. Эти инструменты разделяют текст на границах элементов, таких как заголовки, элементы списков или блоки кода, сохраняя общую структуру документа. Например, LangChain имеет MarkdownHeaderTextSplitter, HTMLHeaderTextSplitter и PythonCodeTextSplitter соответственно. Эти разделители полезны для веб-страниц или технической документации, где важно сохранять иерархическую структуру. Рисунок 7: Чанкинг текста на основе длины с LangChain
Индексация
После подготовки данных путём парсинга документов и чанкинга, завершающий критический шаг в системе RAG — индексация. Индексация — это процесс организации разбитых на чанки данных в структуру, обеспечивающую эффективное и точное извлечение. Этот шаг играет ключевую роль в обеспечении быстрого нахождения системой релевантных чанков информации при поступлении запроса.
Для определения процесса индексации важно понимать различные техники извлечения и выбрать ту, которая лучше всего подходит для задачи. Популярные техники извлечения включают:
- На основе ключевых слов
- Полнотекстовый поиск
- На основе графа знаний
- На основе векторов
Давайте сначала рассмотрим каждую технику, а затем проиндексируем наши данные для обеспечения эффективного извлечения.
На основе ключевых слов
Традиционное извлечение на основе ключевых слов опирается на точное совпадение терминов запроса с содержимым документов. Оно быстрое и простое, но не может понять смысл запроса. Например, оно может испытывать трудности с синонимами, что приводит к неполным или нерелевантным результатам. Этот подход неэффективен при работе с крупномасштабными наборами данных или когда цель — извлечение информации на основе семантического сходства, а не точного совпадения слов.
Полнотекстовый поиск
Полнотекстовые поисковые системы, такие как Elasticsearch [10], предлагают более продвинутый подход, сканируя целые документы для поиска релевантных совпадений. Этот метод позволяет проводить комплексный анализ содержимого документа, включая частичные совпадения и поиск фраз. Однако полнотекстовый поиск сопряжён с более высокими вычислительными затратами, особенно при работе с большими наборами данных, содержащими, например, миллионы PDF-документов. Хотя этот подход эффективен для поиска конкретного текста, он менее эффективен для семантического извлечения.
На основе графа знаний
Извлечение на основе графа знаний — это сложная техника, использующая структурированные связи между сущностями (например, людьми, местами или концепциями) для извлечения информации на основе связей между этими сущностями. Этот метод отлично подходит для ответа на сложные запросы и понимания взаимосвязей в данных. Однако построение и поддержание графа знаний требует значительных усилий, и это не всегда практично для больших неструктурированных наборов данных, таких как коллекции PDF или Wiki-страницы. Для получения дополнительной информации об извлечении на основе графа знаний обратитесь к [11].
На основе векторов
Вместо того чтобы полагаться на текстовые совпадения, этот метод использует высокоразмерные embedding — числовые представления текста и изображений — для измерения сходства между запросом и хранящимися чанками данных. Эта техника позволяет извлекать релевантную информацию, даже когда точные слова в запросе не совпадают с содержимым документа, что делает её более гибкой и мощной для крупномасштабных наборов данных.
Какая техника извлечения подходит для ChatPDF?
Для выбора подходящего метода извлечения давайте сначала поймём масштаб нашей системы и оценим количество чанков данных. В данном случае компания управляет большим набором данных из примерно 5 миллионов страниц. Предположим, что каждая страница содержит примерно 1500 символов и включает три изображения. Используя чанкинг на основе длины с размером чанка 500 символов и перекрытием 200 символов, каждая страница создаст 5 текстовых чанков и 3 чанка изображений. Таким образом, общее количество чанков, с которыми работает система RAG, составляет 5M(1500 / (500-200) + 3)=40M. Ожидается, что эта цифра будет расти примерно на 20 процентов ежегодно, как указано в разделе требований.
При примерно 40 миллионах чанков данных и прогнозируемом ежегодном увеличении на 20 процентов важно выбрать технику извлечения, которая масштабируема и может эффективно справляться с растущим объёмом.
Традиционные методы извлечения [12] [13], такие как на основе ключевых слов и полнотекстовый поиск, широко использовались, но имеют ограничения в скорости, масштабируемости и способности понимать семантический смысл запросов. Извлечение на основе графа знаний требует значительных усилий для построения и поддержания таких графов, что делает их дорогостоящим выбором.
Извлечение на основе векторов, с другой стороны, является основной техникой, используемой в современных системах RAG, благодаря следующим преимуществам:
- Семантическое понимание: Оно может захватывать семантический смысл запроса, обеспечивая более точное извлечение, даже когда точные термины запроса отсутствуют в документе.
- Масштабируемость: Использование embedding-векторов делает этот метод высокомасштабируемым и способным эффективно обрабатывать большие наборы данных.
- Эффективность: После индексации данных как embedding-векторов система может эффективно извлекать релевантные чанки.
Благодаря этим преимуществам мы выбираем извлечение на основе векторов и соответственно индексируем наши данные.
Индексация данных для извлечения на основе векторов
В системе извлечения на основе векторов каждый чанк данных преобразуется в embedding-вектор, представляющий содержимое в числовом формате. При индексации ML-модели используются для вычисления embedding и их хранения в векторной базе данных. Это позволяет системе RAG быстро сравнивать их с embedding запроса и извлекать наиболее релевантную информацию без лишней обработки во время инференса. Мы подробнее рассмотрим архитектуру этих ML-моделей и процесс извлечения в разделе разработки модели.
Подводя итог, мы используем трёхэтапный подход для подготовки PDF для системы RAG. Сначала мы применяем техники парсинга документов для преобразования PDF в структурированный формат, разбивая его на текст, таблицы и изображения. Затем мы используем чанкинг документов для разделения длинного текста на более мелкие, управляемые чанки. Наконец, каждый чанк преобразуется в embedding-вектор и индексируется индивидуально для улучшения точности извлечения.
Разработка модели
Архитектура
В этом разделе рассматривается архитектура системы RAG с фокусом на ML-моделях, используемых в компонентах индексации, извлечения и генерации.
Индексация
Как обсуждалось в разделе подготовки данных, мы используем ML-модели для преобразования чанков данных (например, текста или изображений) в embedding. Этот процесс включает две ML-модели: text encoder и image encoder.
Text encoder
Text encoder — это нейронная сеть, которая преобразует входной текст в плотные векторные представления, или «embedding». Эти embedding захватывают семантический смысл текста, позволяя оценивать сходство текстов. Во время процесса индексации text encoder преобразует каждый текстовый чанк в embedding, который затем сохраняется в базе данных для эффективного извлечения.
Архитектура text encoder обычно основана на encoder-only Transformer, аналогичном тому, что мы рассматривали в Главе 3.
Image encoder
Image encoder преобразует данные изображений в embedding. Его архитектура может быть на основе CNN или Transformer, как мы рассматривали в Главе 5.
Для эффективного извлечения важно выровнять embedding изображений с embedding текста. Например, если запрос «How many cats are in the company?», система должна обеспечить, чтобы закодированный запрос был близок к embedding релевантных изображений, например, с котами. Существуют два основных подхода для достижения этого выравнивания:
- Общее пространство embedding: Используйте encoder изображений и текста, которые генерируют embedding в общем пространстве. CLIP [14] предоставляет предобученные encoder с общим пространством embedding, обеспечивая кросс-модальное извлечение.
- Генерация подписей к изображениям: Сначала сгенерируйте текстовое описание изображения с помощью модели генерации подписей. Затем сгенерированную подпись можно закодировать с помощью text encoder, гарантируя, что данные изображений и текста существуют в одном пространстве embedding. Этот подход полезен при использовании отдельных моделей для text и image encoder или когда обучение совместной модели ресурсоёмко. Для получения дополнительной информации о построении системы генерации подписей к изображениям с нуля обратитесь к Главе 5.
Подводя итог, процесс индексации использует text encoder и image encoder для преобразования чанков данных в embedding. Эти модели часто предобучены, что означает, что их можно применять напрямую без дополнительного обучения. Для целей этой главы мы используем предобученную модель CLIP как text, так и image encoder.
Извлечение
Процесс извлечения включает преобразование запроса пользователя в то же пространство embedding, что и индексированные данные. Это делается с помощью того же text encoder, который использовался в процессе индексации. После вычисления embedding запроса он сравнивается с хранящимися embedding для извлечения наиболее релевантных чанков данных.
Генерация
Компонент генерации отвечает за создание финального ответа на основе запроса пользователя и извлечённого контекста. Эта задача обычно выполняется LLM, которая генерирует контекстуально релевантный текст.
Системы RAG могут работать с различными типами LLM независимо от их архитектуры, включая decoder-only Transformer (подробности см. в Главе 4) или облачные модели, поддерживающие fine-tuning через API [15] [16].
Обучение
Большинство компонентов системы RAG начинают с предобученных моделей, поэтому fine-tuning LLM обычно не является первым шагом в оптимизации производительности. Во многих случаях хорошо спроектированный процесс извлечения в сочетании с эффективным prompt engineering может дать удовлетворительные результаты. Fine-tuning следует рассматривать, когда система постоянно не может предоставить точные или релевантные ответы, даже после настройки параметров извлечения и разработки промптов. Например, если извлечённые документы релевантны, но LLM не генерирует качественные ответы, fine-tuning может помочь LLM лучше понять контекст и нюансы извлечённых данных.
Один многообещающий подход к fine-tuning LLM в системах RAG — Retrieval-Augmented Fine-Tuning (RAFT). Давайте кратко рассмотрим RAFT.
RAFT
RAFT [17] вводит новый метод обучения для улучшения способности LLM обрабатывать как релевантную, так и нерелевантную информацию в извлечённых документах.
В традиционных системах RAG выход LLM сильно зависит от качества извлечённых документов. Однако нерелевантные документы могут быть включены в результаты извлечения. Эти нерелевантные документы могут ввести LLM в заблуждение, заставив её генерировать субоптимальные ответы. RAFT решает эту проблему, включая различие между релевантными и нерелевантными документами в процесс fine-tuning. Этот процесс включает два ключевых шага:
- Разметка документов: Извлечённые документы помечаются как релевантные (золотые) или нерелевантные (отвлекающие). Это предоставляет LLM чёткие сигналы о том, на каких документах следует сосредоточиться.
- Совместное обучение: Во время fine-tuning LLM обучается генерировать ответы на основе релевантных документов, минимизируя влияние нерелевантных документов. Это требует корректировки функции потерь модели для штрафования использования нерелевантных документов при генерации ответа.
Обучая модель приоритизировать релевантный контент и игнорировать отвлекающие факторы, RAFT улучшает способность LLM справляться с зашумлёнными результатами извлечения и генерировать точные и релевантные ответы. Эта способность критически важна в реальных приложениях, где системы извлечения могут быть не всегда идеальными. Для получения дополнительной информации о RAFT обратитесь к [17].
Сэмплирование
Сэмплирование обычно включает генерацию новых данных с помощью генеративной модели. В системе RAG, однако, несколько компонентов работают вместе для создания ответа на запрос пользователя. В этом разделе мы рассмотрим эти компоненты и выделим техники для улучшения производительности на этапах извлечения и генерации системы RAG.
Извлечение
Процесс извлечения происходит в два основных шага:
- Вычисление embedding запроса
- Выполнение поиска ближайших соседей
. Вычисление embedding запроса
Первый шаг включает преобразование запроса пользователя в embedding с помощью text encoder. Этот embedding захватывает семантический смысл запроса, позволяя системе сравнивать его с индексированными embedding чанков данных.
. Выполнение поиска ближайших соседей
После вычисления embedding запроса система выполняет поиск ближайших соседей для нахождения чанков данных, наиболее похожих на запрос. Поиск ближайших соседей решает задачу идентификации точек данных в наборе данных, которые наиболее близки к заданной точке запроса, на основе выбранной меры сходства. Распространённые меры включают евклидово расстояние [18], косинусное сходство [19] или другие метрики расстояния, которые захватывают взаимосвязи между точками данных в пространстве embedding.
Поиск ближайших соседей — фундаментальный компонент информационного поиска, поисковых систем и рекомендательных систем. Даже небольшие улучшения его производительности могут привести к значительным общим улучшениям системы. Учитывая его важность, интервьюеры могут захотеть, чтобы вы углубились в эту тему.
Алгоритмы ближайших соседей обычно делятся на две категории:
- Точный поиск ближайших соседей
- Приближённый поиск ближайших соседей
Точный поиск ближайших соседей
Точный поиск ближайших соседей, также называемый линейным поиском, — это простейшая и наиболее точная форма поиска ближайших соседей. Он вычисляет расстояние между embedding запроса, EqE_qEq, и каждым элементом в наборе данных, извлекая kkk ближайших соседей.
Хотя этот метод гарантирует нахождение истинных ближайших соседей, он имеет временную сложность O(N×D)O(N\times D)O(N×D), где NNN — количество элементов в наборе данных, а DDD — размерность embedding. Эта линейная сложность может сделать процесс очень медленным при работе с крупномасштабными системами, такими как система RAG, индексирующая десятки миллионов элементов. Например, выполнение точного поиска по 40 миллионам элементов для одного запроса потребует 40 миллионов сравнений, что приводит к высоким вычислительным затратам и задержке. Поэтому точный поиск ближайших соседей часто слишком медленный и вычислительно затратный для практического использования.
Приближённый поиск ближайших соседей (ANN)
Во многих приложениях достаточно извлечь элементы, которые достаточно похожи, без необходимости находить точного ближайшего соседа. Алгоритмы ANN используют специализированные структуры данных, позволяющие системе извлекать «достаточно близких» соседей без поиска по всему набору данных, тем самым сокращая время поиска до сублинейной сложности, например, O(log(N)×D)O(\log(N)\times D)O(log(N)×D). Хотя эти алгоритмы обычно требуют некоторой предобработки или дополнительного хранения, они предлагают значительные преимущества в производительности.
Различные алгоритмы ANN можно разделить на следующие категории:
- На основе деревьев
- Locality-sensitive hashing
- На основе кластеризации
- На основе графов
Хотя интервьюер обычно не ожидает, что вы будете знать каждую деталь этих категорий, обычно полезно иметь общее представление о них. Давайте углубимся.
На основе деревьев
Алгоритмы на основе деревьев разбивают пространство данных на множество разделов. Затем они используют характеристики дерева для выполнения более быстрого поиска. Например, k-d дерево [20] разделяет пространство на основе значений признаков, обеспечивая более быстрый поиск путём сужения релевантных областей данных. Другие алгоритмы включают R-деревья [21] и Annoy (Approximate Nearest Neighbor Oh Yeah) [22].
Locality-sensitive hashing (LSH)
LSH группирует похожие точки в корзины с помощью специализированных хеш-функций. Эти функции гарантируют, что точки, близкие в пространстве, хешируются в одну корзину. Это кардинально сокращает пространство поиска, поскольку необходимо проверять только точки в той же корзине, что и запрос, что делает LSH высокоэффективным для больших наборов данных. Подробнее о LSH можно узнать из [23].
На основе кластеризации
Алгоритмы на основе кластеризации организуют данные в кластеры, используя метрики расстояния, такие как косинусное сходство или евклидово расстояние. Это позволяет ограничить поиск ближайшего соседа кластером (кластерами), наиболее релевантным запросу, сокращая количество необходимых сравнений, так как рассматриваются только точки данных внутри выбранного кластера. Конкретно, после организации индексированных элементов в кластеры, ближайшие соседи извлекаются в два шага:
- Межкластерный поиск: Embedding запроса сравнивается с центроидами всех кластеров, и выбираются кластеры, которые ближе заданного порога.
- Внутрикластерный поиск: Embedding запроса сравнивается с элементами в выбранных кластерах.
Этот двухэтапный процесс — сначала сужение поиска до кластера, затем более точный поиск внутри этого кластера — значительно повышает эффективность. Этот процесс показан на Рисунке 16.
На основе графов
Алгоритмы на основе графов, такие как HNSW (hierarchical navigable small world) [24], структурируют данные как граф, где узлы представляют точки данных, а рёбра соединяют их на основе близости в пространстве embedding. HNSW работает путём навигации по этому графу иерархическим способом, начиная с грубого графа более высокого уровня и постепенно переходя к более точным уровням. Поиск уточняется на каждом уровне, исследуя только ближайшие узлы, тем самым кардинально сокращая пространство поиска.
Какая категория поиска ближайших соседей лучше всего подходит для системы извлечения RAG?
В системах RAG количество индексированных элементов обычно огромно и растёт, часто превышая сотни миллионов embedding. Временная сложность точного поиска ближайших соседей слишком высока, поэтому мы полагаемся на алгоритмы ANN для эффективного извлечения релевантных чанков данных.
Различные алгоритмы ANN имеют свои сильные стороны. Выбор правильного алгоритма ANN обычно зависит от таких факторов, как размер набора данных, требуемая скорость и компромиссы по точности. Для простоты мы используем подход ANN на основе кластеризации в компоненте извлечения системы RAG.
Несколько современных фреймворков предоставляют встроенную поддержку ANN, включая:
- Elasticsearch [10]: Широко используемая поисковая система, поддерживающая векторный поиск по сходству.
- FAISS [25]: Популярная библиотека, разработанная Meta, обеспечивающая эффективный поиск ближайших соседей для больших наборов данных.
- ScaNN [26]: Библиотека, разработанная Google, предназначенная для быстрого и эффективного поиска ближайших соседей на больших наборах данных.
Эти фреймворки обычно используются на практике для обеспечения эффективности и масштабируемости компонентов извлечения крупномасштабных систем.
Генерация
Компонент генерации принимает запрос пользователя и извлечённый контекст на вход и генерирует ответ с помощью top-p сэмплирования. Однако мы можем дополнительно улучшить качество сгенерированного ответа, включив техники prompt engineering, как показано на Рисунке 17.
В этом разделе мы углубимся в prompt engineering и рассмотрим, как он улучшает генерацию ответов в системе RAG.
Prompt engineering
Prompt engineering — это мощная техника, которая оптимизирует входные промпты для помощи LLM в генерации более точных и контекстуально релевантных ответов. Тщательно проектируя промпты, мы можем направлять выход модели для лучшего соответствия конкретным задачам, улучшая общую производительность. Хотя prompt engineering можно применять как в извлечении (например, создание лучших запросов для оптимизации поиска), так и в генерации, мы фокусируемся на его применении к генерации в образовательных целях. Тот же подход можно использовать и для улучшения производительности извлечения.
Давайте начнём этот раздел с принципов проектирования промптов, а затем перейдём к техникам prompt engineering.
Принципы проектирования промптов
Эффективное проектирование промптов критически важно для максимизации производительности языковых моделей. Следуя ключевым принципам, мы можем повысить качество сгенерированного вывода и снизить количество нерелевантных или запутанных ответов. Ниже представлены некоторые основные принципы prompt engineering:
- Начинайте просто: Начните с простых промптов и постепенно добавляйте сложность. Итеративные эксперименты — ключ к совершенствованию промптов. Такие инструменты, как Cohere's Playground [27], позволяют легко тестировать и корректировать промпты по мере необходимости.
- Разбивайте сложные задачи: Разбивайте задачи, включающие несколько подзадач, на более мелкие, управляемые шаги. Это позволяет не перегружать LLM и обеспечивает лучшую фокусировку на отдельных подзадачах.
- Используйте ясные инструкции: Будьте явными с инструкциями, используя ясные, ориентированные на действие команды, такие как «Напишите», «Суммаризируйте» или «Переведите». Экспериментируйте с различными инструкциями, чтобы найти то, что лучше всего подходит для вашей задачи. Размещение инструкций в начале промпта, разделённых разделителями, такими как «###», также может помочь организовать промпт.
- Будьте конкретны: Конкретность ведёт к более точным ответам. Ясно опишите, что вы ожидаете в терминах формата, стиля или результатов. Однако избегайте перегрузки промпта ненужными деталями — включайте только то, что релевантно задаче.
- Экспериментируйте с длиной промпта: Учитывайте длину промпта. Слишком много ненужной информации может запутать LLM, а слишком мало может привести к размытым ответам. Найдите баланс, будучи лаконичными, но достаточно детальными для эффективного направления LLM.
Техники prompt engineering
Было разработано несколько техник prompt engineering для улучшения качества выхода LLM. Некоторые из наиболее эффективных включают:
- Chain-of-thought промптинг
- Few-shot промптинг
- Ролевой промптинг
- Пользовательско-контекстный промптинг
Chain-of-thought промптинг
Chain-of-thought (CoT) промптинг [28] включает направление модели через промежуточные шаги рассуждения перед получением финального ответа. Это особенно полезно для сложных запросов, требующих многошагового рассуждения, когда модель должна объединить информацию из нескольких документов для генерации полного ответа. CoT-промпты направляют модель на разбиение своего рассуждения на шаги, что приводит к более точным и содержательным ответам.
CoT был дополнительно расширен такими техниками, как [29], которые позволяют моделям оценивать несколько путей рассуждения перед выбором лучшего ответа. o12 от OpenAI [30] и [31] показали, что способность LLM решать более сложные задачи может быть улучшена путём выделения большего вычислительного бюджета во время инференса, также известного как масштабирование вычислений во время тестирования (test-time compute scaling).
Few-shot промптинг
Few-shot промптинг [32] включает предоставление модели нескольких примеров пар ввод-вывод перед фактическим запросом. Этот метод помогает модели понять желаемый формат и тон вывода, улучшая её способность генерировать ответы, согласованные с предоставленными примерами.
Ролевой промптинг
В некоторых случаях языковой модели может потребоваться принять конкретную «роль» для генерации подходящего ответа. Например, в юридических или медицинских доменах промптинг модели для действий в качестве эксперта в предметной области гарантирует, что ответ содержит необходимый тон, точность и авторитетность.
Пользовательско-контекстный промптинг
Пользовательско-контекстный промптинг адаптирует вывод модели на основе конкретной информации о пользователе, включённой в промпт. Включая профили пользователей, предпочтения или местоположение в запросы, модель может генерировать персонализированные ответы, более релевантные для пользователей.
Этот метод особенно эффективен, когда информация, специфичная для пользователя, критически важна для формирования ответа, например, в персонализированных рекомендациях или запросах, основанных на местоположении.
Собирая всё вместе: prompt engineering для генерации ответов
Комбинирование этих техник позволяет нам создавать высокоэффективные промпты для генерации ответов в системе RAG. Принципы, такие как ясность и конкретность, могут направлять модель на создание более точных выходов. Техники prompt engineering могут значительно усилить возможности генерации RAG, что приводит к более надёжным и контекстуально уместным результатам.
Оценка
В отличие от традиционных ML-моделей, которые оцениваются с помощью чётко определённых количественных метрик, оценка систем RAG более сложна. Эта сложность возникает потому, что качество финального текстового ответа зависит от эффективности множества компонентов в пайплайне. Для отражения этой многоаспектной оценки мы используем диаграмму триады для объяснения взаимосвязей между различными аспектами оценки.
Оценка системы RAG фокусируется на четырёх ключевых аспектах:
- Релевантность контекста
- Достоверность
- Релевантность ответа
- Корректность ответа
Эти аспекты помогают оценить, насколько хорошо система извлекает, генерирует и сопоставляет информацию, релевантную запросу пользователя. Давайте рассмотрим каждый подробнее.
Релевантность контекста
Релевантность контекста измеряет, насколько точно и полно компонент извлечения выбирает релевантные документы на основе запроса. Цель — обеспечить, чтобы весь релевантный контент появлялся вверху результатов извлечения. Этот аспект напрямую оценивает эффективность механизма извлечения. Распространённые метрики для релевантности контекста включают:
- Hit rate
- Mean reciprocal rank (MRR)
- Normalized discounted cumulative gain (NDCG)
- Precision@k
Для получения дополнительной информации о метриках оценки в системах извлечения и ранжирования обратитесь к [33][34].
Достоверность
Достоверность (faithfulness) оценивает, является ли сгенерированный ответ фактически согласованным с извлечённым контекстом. Она проверяет, галлюцинирует ли компонент генерации (т.е. вводит информацию, не обоснованную контекстом). Это критически важно, поскольку система должна создавать ответы, строго отражающие исходные материалы. Оценивая достоверность, мы снижаем риск генерации правдоподобно звучащих, но фактически несогласованных ответов, тем самым повышая надёжность и достоверность вывода.
Достоверность может быть оценена с помощью следующих методов:
- Оценка людьми: Эксперты вручную просматривают сгенерированные ответы, чтобы определить, являются ли они фактически согласованными и корректно ссылаются на извлечённые документы. Этот процесс включает перекрёстную проверку каждого утверждения с исходными материалами для обеспечения обоснованности всей сгенерированной информации.
- Автоматизированные инструменты проверки фактов: Такие инструменты, как [35] и [36], могут автоматизировать процесс валидации, сравнивая сгенерированный ответ с базой данных проверенных фактов. Они предлагают масштабируемое решение для выявления неточностей, тем самым снижая зависимость от оценщиков-людей.
- Проверки согласованности: Этот метод включает оценку того, предоставляет ли LLM согласованную фактическую информацию по нескольким запросам. Регулярные проверки согласованности обеспечивают, что LLM не создаёт противоречивую информацию, что необходимо для поддержания надёжности и связности ответов с течением времени.
Релевантность ответа
Релевантность ответа измеряет, насколько близко сгенерированный ответ соответствует исходному запросу с точки зрения полноты и отсутствия избыточности. Если ответ содержит нерелевантную или избыточную информацию или не содержит важных деталей, он получает низкую оценку релевантности. Этот аспект можно оценить, сравнив вопрос и ответ с помощью другой языковой модели (например, ChatGPT).
Корректность ответа
Корректность ответа фокусируется на том, насколько близко сгенерированный ответ соответствует правильному эталонному ответу. Она измеряет сходство между ними с помощью популярных метрик, включая BLEU, ROUGE и METEOR. Для обзора этих метрик обратитесь к Главе 3.
Общий дизайн ML-системы
Система RAG состоит из нескольких компонентов, которые работают вместе для эффективного извлечения и генерации ответов. В этом разделе мы рассмотрим следующие ключевые компоненты:
- Процесс индексации
- Фильтрация безопасности
- Расширение запроса
- Извлечение
- Генерация
Процесс индексации
Процесс индексации отвечает за преобразование базы знаний в embedding, которые затем сохраняются в индексной таблице для эффективного извлечения. Это начинается с парсинга и чанкинга документов, где текст и изображения в PDF разбиваются на осмысленные чанки данных. Затем эти чанки преобразуются в embedding с помощью text и image encoder CLIP, гарантируя, что embedding текста и изображений находятся в общем пространстве embedding. После получения embedding чанки данных сохраняются в индексной таблице, обеспечивая быстрое извлечение.
Фильтрация безопасности
Компонент фильтрации безопасности обеспечивает безопасность пользовательских запросов и их соответствие рекомендациям системы. Это включает проверку запросов на наличие неуместного или вредоносного контента перед их дальнейшей обработкой. Для получения дополнительной информации о фильтрации и оценке безопасности обратитесь к Главе 4.
Расширение запроса
Расширение запроса повышает качество процесса извлечения путём расширения запроса пользователя для улучшения его связности и устранения опечаток и грамматических ошибок. Расширяя область поиска, расширение запроса помогает системе выявлять дополнительные релевантные данные, которые могли не быть явно упомянуты в исходном запросе, тем самым увеличивая шансы на извлечение более релевантных результатов.
Для получения дополнительной информации о расширении запросов и его технических деталях обратитесь к [37].
Извлечение
Компонент извлечения отвечает за нахождение чанков данных, наиболее релевантных запросу пользователя. Запрос пользователя сначала преобразуется в embedding с помощью text encoder CLIP, а затем алгоритм ANN используется для эффективного извлечения наиболее похожих чанков данных в индексной таблице.
Генерация
После извлечения релевантных чанков данных компонент генерации создаёт финальный ответ. Это включает два основных шага:
- Prompt Engineering: Запрос пользователя и извлечённый контекст объединяются в промпт, который затем оптимизируется с помощью таких техник, как CoT, для структурирования процесса рассуждения модели.
- LLM: LLM генерирует финальный ответ с помощью top-p сэмплирования.
Другие темы для обсуждения
Если в конце собеседования останется время, рассмотрите обсуждение следующих дополнительных тем:
- Обнаружение таблиц при парсинге документов [38] [39] [40].
- Детали алгоритмов приближённого поиска ближайших соседей [20] [21] [23] [24].
- Поддержка документов, загружаемых пользователями [2].
- Динамическая стратегия извлечения [41] [42].
- Переписывание и расширение запросов [43] [37].
- CoT во время инференса и масштабирование вычислений во время тестирования [30] [31].
Резюме
Справочные материалы
[1] Perplexity. https://www.perplexity.ai/. [2] ChatPDF. https://www.chatpdf.com/. [3] LoRA: Low-Rank Adaptation of Large Language Models. https://arxiv.org/abs/2106.09685. [4] Optical character recognition. https://en.wikipedia.org/wiki/Optical_character_recognition. [5] Dedoc GitHub Repository. https://github.com/ispras/dedoc. [6] LayoutParser: A Unified Toolkit for Deep Learning Based Document Image Analysis. https://arxiv.org/abs/2103.15348. [7] Google Cloud document parser API. https://cloud.google.com/document-ai/docs/layout-parse-chunk. [8] PDF.CO document parser API. https://developer.pdf.co/api/document-parser/index.html. [9] Character text splitter in LangChain. https://python.langchain.com/v0.1/docs/modules/data_connection/document_transformers/character_text_splitter/. [10] Elasticsearch. https://www.elastic.co/elasticsearch. [11] A Survey on Knowledge Graphs: Representation, Acquisition, and Applications. https://ieeexplore.ieee.org/document/9416312. [12] Manning, Christopher D. "Introduction to information retrieval." (2008). [https://nlp.stanford.edu/IR-book/information-retrieval-book.html] (https://nlp.stanford.edu/IR-book/information-retrieval-book.html) [13] Modern information retrieval: A brief overview. http://singhal.info/ieee2001.pdf. [14] Learning Transferable Visual Models From Natural Language Supervision. https://arxiv.org/abs/2103.00020. [15] OpenAI finetuning documentation. https://platform.openai.com/docs/guides/fine-tuning. [16] Anthropic finetuning. https://www.anthropic.com/news/fine-tune-claude-3-haiku. [17] RAFT: Adapting Language Model to Domain Specific RAG. https://arxiv.org/abs/2403.10131. [18] Euclidean distance. https://en.wikipedia.org/wiki/Euclidean_distance. [19] Cosine similarity. https://en.wikipedia.org/wiki/Cosine_similarity. [20] Multidimensional binary search trees used for associative searching. https://dl.acm.org/doi/10.1145/361002.361007. [21] R-trees: A dynamic index structure for spatial searching. https://dl.acm.org/doi/10.1145/971697.602266. [22] Annoy library. https://github.com/spotify/annoy. [23] Similarity search in high dimensions via hashing. https://www.cs.princeton.edu/courses/archive/spring13/cos598C/Gionis.pdf. [24] Efficient and robust approximate nearest neighbor search using Hierarchical Navigable Small World graphs. https://arxiv.org/abs/1603.09320. [25] Faiss Documentation. https://faiss.ai/. [26] ScaNN. https://research.google/blog/announcing-scann-efficient-vector-similarity-search/. [27] Developer Playground. https://docs.cohere.com/v2/docs/playground-overview. [28] Chain-of-Thought Prompting Elicits Reasoning in Large Language Models. https://arxiv.org/abs/2201.11903. [29] Tree of Thoughts: Deliberate Problem Solving with Large Language Models. https://arxiv.org/abs/2305.10601. [30] OpenAI o1. https://openai.com/index/learning-to-reason-with-llms/. [31] Scaling LLM Test-Time Compute Optimally can be More Effective than Scaling Model Parameters. https://arxiv.org/abs/2408.03314. [32] Language Models are Few-Shot Learners. https://arxiv.org/abs/2005.14165. [33] Machine Learning System Design Interview. https://www.aliaminian.com/books. [34] Evaluation measure for information retrieval. https://en.wikipedia.org/wiki/Evaluationmeasures(information_retrieval). [35] Ragas. https://docs.ragas.io/en/stable/. [36] ARES: An Automated Evaluation Framework for Retrieval-Augmented Generation Systems. https://arxiv.org/abs/2311.09476. [37] Query2doc: Query Expansion with Large Language Models. https://arxiv.org/abs/2303.07678. [38] TableNet: Deep Learning model for end-to-end Table detection and Tabular data extraction from Scanned Document Images. https://arxiv.org/abs/2001.01469. [39] CascadeTabNet: An approach for end to end table detection and structure recognition from image-based documents. https://arxiv.org/abs/2004.12629. [40] Deepdesrt: Deep learning for detection and structure recognition of tables in document images. https://ieeexplore.ieee.org/document/8270123. [41] Active Retrieval Augmented Generation. https://arxiv.org/abs/2305.06983. [42] Self-RAG: Learning to Retrieve, Generate, and Critique through Self-Reflection. https://arxiv.org/abs/2310.11511. [43] Precise Zero-Shot Dense Retrieval without Relevance Labels. https://arxiv.org/abs/2212.10496.
Сноски
- Актуально на момент написания. ↩
- Подробности неизвестны публике на момент написания. ↩
Генерация реалистичных лиц
Введение
Одним из основных применений генеративного ИИ является создание реалистичных изображений лиц. Это может быть полезно в сфере развлечений, маркетинга и виртуальной реальности. В этой главе мы исследуем технологии, лежащие в основе генерации лиц.
Уточнение требований
Вот типичный диалог между кандидатом и интервьюером:
Кандидат: Каково основное применение системы генерации лиц? Интервьюер: Основной упор делается на развлечения и создание контента, но в будущем мы также рассмотрим её использование для сбора данных.
Кандидат: Фокус только на лицах? Или нужно генерировать тело целиком? Интервьюер: Давайте сосредоточимся только на лицах.
Кандидат: Должны ли сгенерированные лица представлять разнообразие этнических групп, возрастов и полов? Интервьюер: Да. Это крайне важно для обеспечения инклюзивности и исключения предвзятостей.
Кандидат: Должна ли система поддерживать управление атрибутами лица? Например, редактирование мимики на сгенерированном изображении с сохранением идентичности? Интервьюер: Хороший вопрос. Давайте начнём без управления атрибутами. Если останется время, можно дополнительно обсудить управление атрибутами.
Кандидат: Как мы будем собирать обучающие данные? Каков их объём? Интервьюер: Мы используем общедоступные наборы данных с соответствующими лицензиями, чтобы все данные соответствовали требованиям конфиденциальности. В наборе данных содержится 70 000 изображений разнообразных лиц.
Кандидат: Какое разрешение изображений желательно? Интервьюер: Давайте ориентироваться на 1024x1024.
Кандидат: Какова ожидаемая скорость генерации лица? Интервьюер: Система должна генерировать лица в режиме, близком к реальному времени — менее секунды.
Постановка задачи как задачи машинного обучения
Определение входных и выходных данных системы
В системе генерации лиц пользователи обычно не предоставляют конкретных входных данных — они просто запрашивают генерацию нового лица. Поскольку модели машинного обучения (ML) требуют числовых входных данных для начала работы, большинство моделей генерации изображений начинают с вектора случайного шума. Этот шум служит начальными входными данными, которые модель затем преобразует в реалистичное изображение. Если система поддерживает управление атрибутами, пользователи также могут указывать желаемые атрибуты в качестве входных данных для управления генерацией.
Выходные данные, генерируемые в ответ на случайный шум, представляют собой реалистичное изображение человеческого лица. Эти выходные данные также должны отражать желаемые атрибуты (если указаны), такие как возраст, пол и причёска.
Выбор подходящего подхода ML
В этом разделе мы рассмотрим распространённые подходы ML к генерации изображений. Мы обсудим сильные стороны и ограничения каждого подхода и выберем наиболее подходящий для нашего случая использования.
Хотя существуют различные подходы к генерации изображений, мы сосредоточимся на тех, которые наиболее широко используются в отрасли. Существует четыре основных подхода к генерации изображений:
- Вариационный автоэнкодер
- Генеративно-состязательная сеть
- Авторегрессионная модель
- Diffusion-модель
Вариационный автоэнкодер
Вариационный автоэнкодер (VAE) — это архитектура генеративной модели, разработанная для изучения распределения данных. Это позволяет VAE генерировать новые точки данных, выполняя выборку из этого изученного распределения.
VAE состоит из двух основных компонентов:
- Encoder
- Decoder
Encoder: Encoder — это нейронная сеть, которая отображает входное изображение в пространство меньшей размерности, известное как latent space. Выходными данными encoder являются вектор в latent space — закодированное представление входного изображения.
Decoder: Decoder — это ещё одна нейронная сеть, которая отображает закодированное представление обратно в изображение. Выходными данными decoder является изображение того же размера, что и исходное входное изображение.
В процессе обучения VAE кодирует входные данные в latent space, а затем восстанавливает исходные входные данные из этого закодированного представления. После обучения VAE может генерировать новые изображения, выбирая точки из изученного многомерного гауссова распределения и используя decoder для отображения этих точек в форму изображения.
С помощью трюка перепараметризации (подробнее см. в [2]) VAE моделируют вектор в latent space как выбранный из многомерного гауссова распределения. Моделирование latent space помогает VAE изучать осмысленные представления, которые можно плавно интерполировать, что выгодно для таких задач, как морфинг изображений и создание вариаций входных данных.
VAE имеют ряд сильных и слабых сторон.
Преимущества:
- Простая архитектура: encoder и decoder являются архитектурами нейронных сетей, простыми в реализации.
- Быстрая генерация: по сравнению с другими подходами VAE обеспечивают быструю генерацию изображений. Процесс включает выборку случайного шума из latent space и его декодирование в изображение с помощью decoder.
- Стабильное обучение: обучение VAE, как правило, простое и стабильное.
- Возможность сжатия: помимо генерации изображений, VAE являются мощным инструментом для сжатия изображений в представления меньшей размерности.
Недостатки:
- Менее реалистичные изображения: VAE плохо справляются с захватом высокочастотных деталей. Это приводит к изображениям, которые менее реалистичны по сравнению с теми, которые генерируются некоторыми другими подходами.
- Размытость: существенным ограничением VAE является их склонность к созданию размытых изображений, лишённых чётких деталей.
- Ограниченная новизна: VAE, как правило, с трудом генерируют изображения, существенно отличающиеся от их обучающих данных. Это ограничивает их способность создавать новые выходные данные.
- Ограниченное управление генерацией: VAE не предназначены для поддержки дополнительных управляющих входных данных, таких как текстовые описания или управление атрибутами желаемого изображения.
Подводя итог, можно сказать, что VAE не являются лучшим выбором для генерации высококачественных детальных изображений. Однако их сильная сторона заключается в эффективном кодировании изображений в компактные представления. В главе 11 мы рассмотрим VAE и используем их возможности сжатия для построения эффективной системы генерации видео.
Генеративно-состязательная сеть
Генеративно-состязательная сеть (GAN) [3] состоит из двух сетей:
- Генератор: нейронная сеть, которая преобразует случайный шум в изображение.
- Дискриминатор: ещё одна нейронная сеть, которая определяет, является ли данное изображение реальным или искусственно сгенерированным.
В процессе обучения эти две сети участвуют в непрерывной игре: генератор учится создавать более реалистичные изображения, а дискриминатор становится лучше в различении реальных и сгенерированных изображений. Если сгенерированное изображение правильно классифицируется как «сгенерированное», генератор получает штраф за создание нереалистичного изображения. Этот состязательный процесс продолжается до тех пор, пока генератор не создаёт изображения, которые дискриминатор больше не может отличить от реальных.
Преимущества:
- Высококачественная генерация: GAN известны своей способностью генерировать высококачественные изображения.
- Быстрая генерация: хотя GAN в целом медленнее VAE, генератор всё равно может генерировать изображение за один прямой проход.
- Управление атрибутами: архитектура GAN может быть изменена для управления конкретными атрибутами, такими как возраст или выражение лица. Например, пользователь может запросить изображение лица, которое выглядит счастливым и старым.
Недостатки:
- Нестабильность обучения: GAN сложно обучать. Распространёнными проблемами обучения являются вырождение моды [4], когда генератор создаёт ограниченное разнообразие выходных данных, и отсутствие сходимости [5], когда модель GAN не стабилизируется в процессе обучения.
- Ограниченное управление: хотя GAN допускают управление атрибутами, выйти за его рамки сложно, например, использовать текстовое описание для генерации изображения [6].
- Ограниченная новизна: хотя GAN хорошо справляются с генерацией вариаций изображений в определённой области, они, как правило, с трудом генерируют новые изображения, существенно отличающиеся от их обучающих данных.
Подводя итог, можно сказать, что GAN сложно обучать и они предлагают ограниченный контроль над сгенерированными изображениями. Однако они могут генерировать детальные изображения и поддерживают управление атрибутами лица, что делает их подходящими для таких приложений, как генерация лиц и редактирование изображений.
Авторегрессионная модель
В авторегрессионном моделировании генерация изображений формулируется как задача генерации последовательности, где каждая часть изображения генерируется последовательно. Эта последовательная генерация позволяет использовать архитектуру Transformer, позволяя нам воспользоваться её мощной способностью захватывать дальние зависимости.
Преимущества:
- Высокая детализация и реалистичность: авторегрессионные модели генерируют изображения с высоким уровнем детализации и чёткости.
- Стабильное обучение: по сравнению с GAN обучение авторегрессионных моделей обычно более стабильно.
- Управление генерацией: возможно управление генерацией изображений с использованием дополнительных входных данных, например текстового запроса, описывающего желаемое содержимое изображения. Эта гибкость обусловлена архитектурой Transformer, которая может поддерживать любое количество входных данных в качестве части входной последовательности.
- Поддержка мультимодального кондиционирования: авторегрессионные модели легко поддерживают кондиционирование по различным модальностям. Например, если в качестве входных данных предоставить праздничное аудио, сгенерированное изображение будет соответствовать звуку. Эта гибкость обусловлена архитектурой Transformer, которая может поддерживать различные модальности в качестве входных данных, если они предоставлены в виде последовательности числовых векторов.
- Новизна: авторегрессионные модели способны генерировать новые и сложные изображения. Например, они могут сгенерировать изображение «авокадо на стуле на Марсе», даже если не видели подобных примеров в своих обучающих данных.
Недостатки:
- Медленная генерация: авторегрессионные модели генерируют изображение последовательно, по одному token за раз. Эта последовательная генерация делает их медленнее по сравнению с VAE или GAN.
- Ресурсоёмкость: эти модели, как правило, очень большие, с миллиардами параметров. Обучение таких больших моделей требует значительных вычислительных ресурсов, что увеличивает стоимость.
- Ограниченные манипуляции с изображениями: в отличие от VAE и GAN, авторегрессионные модели не имеют структурированного latent space, который можно легко исследовать или манипулировать. Это ограничивает определённые типы манипуляций с изображениями, такие как управление атрибутами лица.
Подводя итог, хотя авторегрессионные модели медленны в генерации из-за своей последовательной природы, они могут генерировать высоко детализированные и новые изображения. Многие популярные модели генерации изображений, такие как DALL-E [7] от OpenAI и Muse [8] от Google, основаны на авторегрессионном моделировании. Глава 8 рассмотрит этот подход подробнее.
Diffusion-модель
Diffusion-модели — ещё один популярный подход к генерации изображений, показавший замечательные возможности. Diffusion-модели формулируют генерацию изображений как итерационный процесс. В процессе обучения шум постепенно добавляется к изображениям, а нейронная сеть обучается предсказывать этот шум. При генерации изображений в процессе инференса начинается со случайного шума. Затем обученная нейронная сеть используется для итерационного удаления шума из изображения, преобразуя шум в осмысленное изображение. Это преобразование происходит за фиксированное количество шагов, при этом на каждом шаге модель добавляет детали к изображению.
Преимущества:
- Высокая детализация и реалистичность: diffusion-модели могут генерировать изображения исключительного качества и реалистичности.
- Стабильное обучение: по сравнению с GAN обучение diffusion-моделей, как правило, стабильно.
- Управление генерацией: подобно авторегрессионным моделям, diffusion-модели можно контролировать с помощью различных входных данных, например текста, описывающего желаемое изображение.
- Новизна и креативность: diffusion-модели могут генерировать новые и образные изображения.
- Устойчивость к зашумлённым изображениям: diffusion-модели эффективно удаляют шум из изображений благодаря своему процессу удаления шума. Это может быть полезно в определённых приложениях, таких как удаление шума из изображений.
Недостатки:
- Медленная генерация: diffusion-модели генерируют изображения за несколько шагов удаления шума. Этот итерационный процесс делает их медленнее по сравнению с другими методами.
- Ресурсоёмкость: diffusion-модели обычно большие, с миллиардами параметров. Это делает их вычислительно интенсивными и, следовательно, дорогостоящими для обучения.
- Ограниченные манипуляции с изображениями: в отличие от VAE и GAN, у diffusion-моделей нет структурированного latent space для манипуляций с изображениями.
Подводя итог, хотя diffusion-модели медленные, они показали впечатляющую производительность в генерации высоко детализированных, разнообразных и образных изображений. Большинство современных моделей генерации изображений, таких как DALL·E 3 [9], основаны на diffusion-моделях. Глава 9 рассмотрит diffusion-модели подробно.
| Characteristics | VAE | GAN | Autoregressive | Diffusion |
|---|---|---|---|---|
| Quality | Low | Moderate | High | Exceptional |
| Speed | Fast | Fast | Slow | Slow |
| Training stability | Stable | Unstable | Stable | Stable |
| Control over generation | Limited | Limited | Flexible | Moderate |
| Facial manipulation | No | Yes | No | No |
| Novelty | Limited | Limited | High | High |
| Resource intensity | Moderate | Moderate | High | High |
Таблица 1: Сравнение различных подходов к генерации изображений
Для реалистичной генерации лиц мы выбираем GAN в качестве основного подхода. GAN особенно эффективны, поскольку позволяют манипулировать атрибутами лица через структурированный latent space, что является опциональным требованием в этой главе.
Подготовка данных
Разработка реалистичной системы генерации лиц требует большой коллекции изображений. Здесь у нас есть 70 000 разнообразных изображений человеческих лиц. Для подготовки этих изображений к обучению мы применяем следующие шаги:
- Удаление низкокачественных или низкоразрешённых изображений: мы удаляем изображения с низким разрешением и используем ML-модели для фильтрации низкокачественных, размытых. Это обеспечивает обучение модели только на высококачественных изображениях.
- Аугментация изображений: мы применяем методы аугментации данных, такие как отражение, поворот или настройка цвета, чтобы искусственно увеличить размер обучающих данных. Это помогает модели видеть больше вариаций изображения в процессе обучения и, таким образом, лучше обобщать впоследствии.
- Нормализация и изменение размера изображений: мы приводим все изображения к стандартному размеру, например 1024x1024. Мы также нормализуем изображения до стандартного диапазона, как правило, от -1 до 1.
- Повышение разнообразия: мы используем ML-классификаторы для маркировки изображений по полу, возрасту и другим атрибутам. Затем мы корректируем набор данных для обеспечения сбалансированного представления различных групп. Этот шаг имеет решающее значение для предотвращения предвзятостей в сгенерированных лицах.
Разработка модели
Архитектура
GAN состоят из двух компонентов: генератора и дискриминатора. Давайте кратко рассмотрим архитектуру каждого компонента.
Генератор
Компонент генератора принимает случайный шум в качестве входных данных и преобразует его в изображение. Его архитектура состоит из серии блоков upsampling, каждый из которых увеличивает пространственные размеры (высоту и ширину) своих входных данных. Эти блоки постепенно преобразуют низкоразмерный вектор шума в двумерное изображение желаемого размера.
Давайте поговорим о трёх основных компонентах блока upsampling:
- Транспонированная convolution
- Слой нормализации
- Нелинейная активация
Транспонированная convolution
Транспонированная convolution, также известная как деконволюция или convolution upsampling, — это операция, используемая в нейронных сетях для увеличения пространственного разрешения карт признаков — по сути, выполняющая обратное обычной convolution. Она широко используется в таких приложениях, как генерация изображений, семантическая сегментация и суперразрешение, где цель состоит в восстановлении выходных данных более высокого разрешения из входных данных более низкого разрешения.
В отличие от стандартной convolution, которая скользит фильтром по входным данным, транспонированная convolution начинается с вставки нулей между пикселями входной карты признаков, фактически расширяя её. Затем расширенные входные данные свёртываются с фильтром, где шаг фильтра1 и отступ2 настраиваются для достижения желаемого размера выходных данных. Например, начиная с входных данных 1x1x100, с 1024 фильтрами размера ядра 1x1 и шагом 1, мы получаем карту признаков 4x4x1024. На следующем шаге с 512 фильтрами размера ядра 3x3 и шагом 1 мы получаем карту признаков 8x8x512. Это первые два этапа upsampling, показанные на Рисунке 8.
В PyTorch этот слой обычно реализуется с помощью «ConvTranspose2d». Чтобы узнать больше о convolution и транспонированных convolution, обратитесь к [11].
Слой нормализации
Слой нормализации улучшает стабильность обучения, масштабируя входные данные до единообразного распределения.
Обучение GAN нестабильно, поскольку в нём участвуют две сети (генератор и дискриминатор), конкурирующие друг с другом. Это может привести к таким проблемам, как вырождение моды, когда генератор производит ограниченное разнообразие, или осцилляции, когда генератор и дискриминатор не могут сойтись в процессе обучения. Нормализация помогает стабилизировать обучение, масштабируя активации на каждом слое, тем самым снижая риск затухания или взрыва градиентов. Это помогает поддерживать единообразное распределение активаций, что критично для сбалансированной конкуренции между генератором и дискриминатором. С более надёжным процессом оптимизации мы можем использовать более высокую скорость обучения, ускорить обучение и уменьшить время, необходимое для сходимости. Мы обсудим проблемы обучения и меры по их устранению позже в этой главе.
Существует несколько слоёв нормализации, каждый со своим способом нормализации данных:
- Batch normalization
- Нормализация по слою
- Instance normalization
- Групповая нормализация
Batch Normalization (BN)
BN [12] нормализует входные данные слоя по измерению batch, вычисляя среднее и дисперсию для каждого признака. Затем нормализованные данные масштабируются и сдвигаются с использованием обучаемых параметров.
- Преимущества: BN помогает стабилизировать процесс обучения и позволяет использовать более высокие скорости обучения, ускоряя обучение. Также выступает регуляризатором, снижая вероятность переобучения.
- Применение: широко используется в глубоких сетях, включая свёрточные нейронные сети (CNN) и генераторы GAN.
Layer Normalization (LN)
LN [13] нормализует входные данные по признакам каждого отдельного образца, а не по измерению batch. Поэтому он вычисляет среднее и дисперсию для каждого признака по всему вектору признаков каждого образца.
- Преимущества: LN эффективен в условиях, когда размеры batch малы или переменны, например в рекуррентных нейронных сетях (RNN) и Transformer.
- Применение: часто используется в последовательностных моделях и сценариях, где важно единообразное поведение между образцами.
Instance Normalization (IN)
IN [14] работает путём нормализации по каждой карте признаков индивидуально для каждого образца.
- Преимущества: IN полезен для задач, в которых внешний вид отдельных образцов сильно варьируется, поскольку позволяет сети сосредоточиться на содержимом, а не на стиле.
- Применение: широко используется в задачах переноса стиля и генерации изображений.
Group Normalization (GN)
GN [15] нормализует входные данные, разделяя признаки на группы и нормализуя внутри каждой группы. Он предлагает баланс между BN и LN.
- Преимущества: GN полезен в случаях с очень маленькими размерами batch, где BN может быть неэффективен.
- Применение: часто применяется в задачах, где BN не работает из-за малых размеров batch или когда необходима единообразность поведения слоёв по группам признаков.
Нелинейная активация
Нелинейные функции активации, такие как ReLU [16], вносят нелинейность в модель, позволяя ей изучать сложные закономерности и представления. Без нелинейности сеть по сути была бы линейным преобразованием, независимо от глубины, что делало бы её неспособной моделировать сложные распределения данных, такие как изображения, речь или сложные функции.
Как показано на Рисунке 11, наш генератор состоит из блоков upsampling (ConvTranspose2D), каждый из которых сопровождается слоем нормализации (BatchNorm2D) и нелинейной активацией (ReLU). Финальный блок использует «Tanh» [17] вместо «ReLU». Этот выбор обеспечивает диапазон финальных выходных данных от -1 до 1, соответствующий диапазону пикселей нашего изображения после подготовки данных.
Дискриминатор
Задача дискриминатора — различать реальные и сгенерированные изображения. Он функционирует как бинарный классификатор, принимая изображение в качестве входных данных и выдавая вероятность того, что изображение является реальным.
Дискриминатор включает серию блоков downsampling, за которыми следует классификационная голова. Блоки downsampling постепенно уменьшают пространственные размеры входного изображения, извлекая его признаки. Затем классификационная голова обрабатывает извлечённые признаки для предсказания вероятности того, что входное изображение является реальным.
Блок downsampling состоит из нескольких операций convolution для постепенного уменьшения пространственного измерения входных данных. В PyTorch мы обычно используем слой «Conv2D» с шагом 2 для уменьшения пространственных размеров вдвое. Как и в генераторе, batch normalization (BatchNorm2D) и нелинейная функция активации (ReLU) используются между слоями convolution для повышения стабильности обучения и производительности.
Классификационная голова включает один или два полностью связанных слоя, за которыми следует функция активации sigmoid. Функция sigmoid обеспечивает диапазон финальных выходных данных от 0 до 1, что критично для интерпретации выходных данных как вероятности.
На протяжении многих лет были разработаны различные версии GAN для различных целей. Например, StyleGAN [18] изменяет архитектуру генератора для управления атрибутами сгенерированных лиц, такими как возраст, цвет волос и мимика. Для получения подробной информации об архитектуре StyleGAN и ключевых архитектурных решениях обратитесь к [18].
Обучение
Для создания реалистичных изображений мы обучаем GAN с использованием уникального процесса, называемого состязательным обучением. При состязательном обучении генератор и дискриминатор обучаются одновременно в игровом сценарии. Генератор стремится создавать изображения, выглядящие как настоящие, в то время как дискриминатор улучшает свою способность различать реальные и сгенерированные изображения. В ходе этого состязательного процесса генератор учится создавать всё более убедительные изображения. В то же время дискриминатор становится лучше в выявлении поддельных изображений. Этот конкурентный процесс продолжается до тех пор, пока генератор не создаёт изображения, которые дискриминатор больше не может обнаружить как поддельные.
При обучении GAN важно обеспечить совместное улучшение как генератора, так и дискриминатора, избегая сценариев, в которых один доминирует над другим. Эмпирически показано, что такой баланс имеет решающее значение для успешного обучения. Для поддержания этого баланса принято чередовать следующие два шага:
- Обучать дискриминатор несколько итераций, держа генератор замороженным.
- Обучать генератор несколько итераций, держа дискриминатор замороженным.
Далее рассмотрим цель ML и функцию потерь для обучения модели GAN.
Цель ML и функция потерь
Генератор и дискриминатор имеют свои собственные конкретные, конфликтующие цели. Дискриминатор стремится точно различать реальные и сгенерированные изображения. Генератор стремится создавать изображения, которые дискриминатор не может отличить от реальных. Мы сначала изучим функцию потерь и цель ML каждого компонента, а затем объединим их в единую функцию потерь для GAN.
Дискриминатор
Мы используем бинарную перекрёстную энтропию в качестве функции потерь для дискриминатора, поскольку она обычно применяется в моделях бинарной классификации.
Где:
- D(x(i))D\left(x^{(i)}\right)D(x(i)) — предсказанные дискриминатором вероятности для реального изображения,
- G(z(j))G\left(z^{(j)}\right)G(z(j)) — выходные данные генератора (поддельное изображение) для случайного шума,
- mmm — количество реальных изображений,
- nnn — количество поддельных изображений.
Цель ML дискриминатора — минимизировать функцию потерь бинарной перекрёстной энтропии.
Генератор
Генератор стремится создавать реалистичные изображения, которые дискриминатор не может отличить от реальных. В идеале дискриминатор должен предсказывать вероятности, близкие к 1, для изображений, созданных генератором. Для достижения этого цель ML формулируется как максимизация log(D(G(z(j))))\log \left(D\left(G\left(z^{(j)}\right)\right)\right)log(D(G(z(j)))) для всех поддельных изображений или, эквивалентно, минимизация следующей функции потерь:
Minimax loss GAN
Minimax loss [19], изначально использованный в статье о GAN, объединяет потери генератора и дискриминатора в единую функцию:
Дискриминатор стремится максимизировать потери, а генератор — минимизировать их. Следовательно, общая цель ML такова:
Помимо minimax loss, исследователи предложили другие функции потерь для улучшения стабильности обучения в GAN. Чтобы узнать больше об этих функциях потерь, обратитесь к [20].
Типичные проблемы обучения GAN
GAN сложнее обучать по сравнению с другими генеративными моделями, такими как авторегрессионные или diffusion-модели. Обсуждение этих проблем полезно на собеседовании по проектированию ML-систем. В этом разделе мы обсудим три основные проблемы обучения GAN:
- Затухание градиентов
- Вырождение моды
- Отсутствие сходимости
Затухание градиентов
Проблема затухания градиентов [21] возникает, когда градиенты становятся очень маленькими в процессе обучения. Эта проблема в первую очередь затрагивает генератор. Когда дискриминатор становится слишком хорошим в различении реальных и поддельных изображений, он предоставляет генератору очень маленькие значения градиентов для обновления его параметров. Это замедляет или останавливает процесс обучения генератора.
Два распространённых метода смягчения проблемы затухания градиентов:
- Модифицированный minimax loss
- Wasserstein loss
Модифицированный minimax loss: оригинальная статья о GAN указывает, что минимизация исходной цели ML может привести к остановке обучения. Чтобы преодолеть это, в статье рекомендуется изменить цель генератора на максимизацию
Это небольшое изменение вдохновлено формулировкой цели ML с другой точки зрения. С этим изменением генератор стремится максимизировать вероятность того, что поддельные изображения будут идентифицированы как реальные, а не минимизировать вероятность того, что поддельные изображения будут идентифицированы как поддельные.
Wasserstein loss: эта функция потерь используется в модифицированном GAN, называемом «Wasserstein GAN» или «WGAN». Давайте рассмотрим цель ML для дискриминатора и генератора WGAN.
- Дискриминатор WGAN: дискриминатор для WGAN, также известный как «критик», отличается от дискриминатора в традиционном GAN. Вместо классификации изображений как реальных или поддельных критик выдаёт оценку, представляющую «реалистичность» изображения. Потери критика определяются как разница между выходными данными критика для реальных и поддельных изображений. Следовательно, цель ML критика — максимизировать потери критика.
- Генератор WGAN: цель ML для генератора WGAN — максимизировать вероятность того, что поддельные изображения будут идентифицированы как реальные:
Вырождение моды
В идеале модель GAN должна создавать различные вариации изображений с разными случайными входными данными. Вырождение моды (mode collapse) относится к ситуациям, когда генератор создаёт ограниченное разнообразие изображений. Давайте разберёмся, почему это может происходить.
В процессе обучения генератор может научиться обманывать систему, находя одно изображение, которое кажется наиболее правдоподобным для дискриминатора. Как только генератор находит это наиболее правдоподобное изображение, он может продолжать создавать одно и то же изображение, чтобы обмануть дискриминатор. Следовательно, генератор никогда не учится генерировать другие изображения. Этот тип сбоя в GAN известен как «вырождение моды» (mode collapse). Два распространённых метода смягчения вырождения моды:
- Wasserstein loss
- Unrolled GAN [22]
Чтобы узнать больше об unrolled GAN и о том, как он смягчает вырождение моды, обратитесь к [4].
Отсутствие сходимости
Обучение GAN, как правило, сложное, а процесс зачастую нестабильный. Это обусловлено проблемой «отсутствия сходимости», распространённой в обучении GAN. Давайте разберёмся, почему это происходит.
По мере улучшения генератора в процессе обучения производительность дискриминатора снижается, поскольку становится всё труднее различать реальные и поддельные изображения. Если генератор достигает точки, в которой он может идеально имитировать реальные данные, точность дискриминатора падает до 50 процентов; он фактически начинает делать случайные предположения, подобно подбрасыванию монеты. Это снижение производительности дискриминатора препятствует сходимости GAN, поскольку его обратная связь постепенно становится всё менее и менее полезной для генератора. Когда обучение продолжается за определённую точку, генератор начинает обучаться на бесполезной обратной связи, и его качество может ухудшиться в результате.
Существуют различные подходы для улучшения стабильности обучения и сходимости GAN:
- Нормализация: применение таких методов, как batch normalization, помогает стабилизировать обучение, обеспечивая единообразное распределение по слоям.
- Различные скорости обучения: использование различных скоростей обучения для генератора и дискриминатора может помочь сбалансировать их прогресс и избежать нестабильности обучения.
- Регуляризация: применение методов регуляризации, таких как затухание весов, предотвращает переобучение и помогает поддерживать стабильность обучения.
- Добавление шума к входным данным дискриминатора: внедрение шума во входные данные дискриминатора может предотвратить его слишком быстрое усиление, что помогает сбалансировать конкуренцию между генератором и дискриминатором.
Чтобы узнать больше о конкретных деталях этих подходов, обратитесь к [21] и [23].
Сэмплирование
Сэмплирование — это процесс генерации новых изображений из обученной модели GAN. Прежде чем обсуждать методы сэмплирования, давайте сначала рассмотрим концепцию latent space в GAN.
В процессе обучения генератор учится преобразовывать различные векторы шума в изображения. Этот процесс формирует latent space — многомерное пространство, где каждая точка представляет потенциальный вектор шума. Этот latent space важен для GAN, поскольку его генератор может отображать каждую точку на соответствующее изображение.
Для генерации реалистичного изображения лица мы выбираем точку из этого latent space, известную как вектор в latent space. Затем генератор берёт этот вектор и преобразует его в изображение.
Существует два метода выборки вектора из изученного latent space:
- Случайное сэмплирование
- Усечённое сэмплирование
Случайное сэмплирование
Случайное сэмплирование использует стандартное гауссово распределение для извлечения векторов из latent space. Это обеспечивает разнообразный выбор векторов, что приводит к генерации разнообразных изображений.
Усечённое сэмплирование
Усечённое сэмплирование ограничивает векторы меньшей областью высокой вероятности в latent space. Усекая распределение, метод снижает вероятность генерации выбросов, что приводит к изображениям более высокого качества. Этот подход выгоден, когда основная цель — поддерживать высокий реализм в сгенерированных лицах. Если вас интересуют детали и реализация усечённого сэмплирования, обратитесь к [24].
Подводя итог, случайное сэмплирование обеспечивает разнообразие, исследуя весь latent space, тогда как усечённое сэмплирование фокусируется на области высокой вероятности для повышения реализма. Для реалистичной генерации лиц мы используем случайное сэмплирование, так как оно обеспечивает разнообразие и обычно хорошо работает на практике.
Оценка
Метрики офлайн-оценки
Оценка систем генерации изображений включает оценку как качества, так и разнообразия сгенерированных изображений. Для этой цели разработано несколько метрик, таких как Inception score [25], расстояние Фреше для Inception (FID) [26] и расстояние ядра Inception (KID) [27]. Среди них наиболее широко используются Inception score и FID. Оценка людьми по-прежнему остаётся важным методом оценки генеративных моделей.
На собеседовании по проектированию ML-систем цель интервьюера — оценить вашу интуицию и практическое понимание, а не проверить ваши детальные знания формул и теорий. Однако широкое понимание некоторых из этих метрик всё же может быть полезным. Давайте кратко рассмотрим Inception score и FID.
Inception score
Inception score — широко используемая метрика для оценки качества сгенерированных изображений в генеративных моделях, таких как GAN. Метрика опирается на предобученную модель классификации изображений, такую как «Inception v3» [28], для оценки того, насколько хорошо сгенерированные изображения напоминают реальные объекты.
Вот пошаговое объяснение того, как вычисляется метрика:
- Генерация изображений: мы начинаем с генерации большого набора изображений с использованием модели, которую хотим оценить.
- Вычисление вероятностей классов: для каждого сгенерированного изображения модель Inception предоставляет распределение вероятностей по всем 1000 классам объектов. Высококачественное изображение должно приводить к распределению с пиком (т.е. высокой вероятностью для одного класса), указывающим, что модель распознаёт его как чёткий экземпляр класса.
- Вычисление маргинального распределения: маргинальное распределение — это среднее предсказанных вероятностей классов по всем изображениям. Это помогает нам понять общее распределение классов, представленных в сгенерированном наборе. Если изображения разнообразны, маргинальное распределение будет плоским и распределённым по многим классам.
- Вычисление дивергенции KL: дивергенция KL измеряет, насколько предсказанное распределение классов для каждого изображения отличается от маргинального распределения. Высококачественные изображения будут иметь распределение, сильно отличающееся от маргинального распределения. Это связано с тем, что ожидается, что высококачественное изображение будет иметь пик в своём распределении, тогда как маргинальное распределение, как ожидается, будет близко к равномерному, если изображения разнообразны.
- Вычисление Inception score: Inception score — это экспоненцированное среднее дивергенции KL по всем изображениям. Высокий Inception score указывает на то, что отдельные изображения были уверенно классифицированы в разнообразные классы, что означает, что сгенерированные изображения одновременно разнообразны и высокого качества.
Как Inception score измеряет как разнообразие, так и качество?
- Разнообразие: Inception score оценивает разнообразие, проверяя, приводят ли сгенерированные изображения к почти равномерному маргинальному распределению по классам, что указывает на равномерное распределение изображений по различным классам.
- Качество: высококачественные изображения приводят к резкому, пикообразному распределению вероятностей, указывающему, что изображение чётко распознаётся как принадлежащее определённому классу. Inception score сравнивает это распределение с маргинальным распределением для оценки качества изображения.
Расстояние Фреше для Inception (FID)
FID — ещё одна популярная метрика для оценки качества изображений, создаваемых генеративными моделями. Она оценивает, насколько распределение сгенерированных изображений похоже на распределение реальных изображений. В отличие от Inception score, который использует вероятности классов, FID учитывает статистику признаков, извлечённых предобученной моделью, такой как Inception v3. Модель Inception выбрана потому, что она обучена на большом и разнообразном наборе данных (ImageNet) и может извлекать значимые признаки, представляющие содержимое и стиль изображений.
Вот пошаговое объяснение того, как вычисляется FID:
- Генерация изображений: мы начинаем с генерации большого набора изображений с использованием модели, которую хотим оценить. Эти изображения будут сравниваться с набором реальных изображений для оценки их качества и разнообразия.
- Извлечение признаков: мы пропускаем каждое изображение (как сгенерированное, так и реальное) через модель Inception v3 и извлекаем признаки («активации») из определённого слоя, обычно одного из ближних к концу сети. Признаки из этого глубокого слоя захватывают высокоуровневую информацию — такую как формы, текстуры и объекты — что критично для оценки реалистичности изображений.
- Вычисление среднего и ковариации: мы вычисляем среднее и ковариацию извлечённых признаков отдельно для сгенерированных и реальных изображений. Эти статистические меры суммируют распределения признаков для обоих наборов изображений.
- Вычисление расстояния Фреше: мы вычисляем FID как расстояние Фреше между средним и ковариацией сгенерированных и реальных изображений. Расстояние Фреше измеряет, насколько близки два распределения. Более низкий FID указывает на большее сходство между распределениями, что означает, что сгенерированные изображения более реалистичны и разнообразны. Чтобы узнать больше о расстоянии Фреше и его формуле, обратитесь к [29].
Как FID измеряет как разнообразие, так и качество?
- Разнообразие: FID учитывает ковариацию признаков, отражающую разброс и вариацию в признаках изображений. Разнообразный набор сгенерированных изображений будет иметь распределение признаков, аналогичное распределению реальных изображений, демонстрируя способность модели создавать широкий спектр различных изображений.
- Качество: FID гарантирует высокое качество сгенерированных изображений, сравнивая их распределения признаков с распределениями реальных изображений. Если среднее и ковариация признаков сгенерированных изображений схожи с признаками реальных изображений, это указывает на то, что сгенерированные изображения, вероятно, будут высокого качества.
FID и Inception score полезны для оценки качества и разнообразия моделей генерации изображений, но они не всегда согласуются с человеческими суждениями. Это в основном связано с тем, что эти метрики опираются на классы ImageNet, которые могут вносить артефакты. Авторы [30] предполагают, что использование модели, не обученной на ImageNet, такой как CLIP, может обеспечить лучшее соответствие с человеческой оценкой.
Хотя метрики на основе CLIP демонстрируют перспективы в улучшении соответствия с человеческими суждениями, они всё ещё не могут полностью заменить инсайты, полученные от прямой обратной связи с людьми. Человеческая оценка по-прежнему остаётся наиболее надёжным методом оценки качества сгенерированных изображений, поскольку она улавливает нюансы, которые автоматические метрики могут упустить.
Оценка людьми
Оценка людьми имеет решающее значение для оценки систем генерации изображений, поскольку автоматические метрики могут упустить субъективные качества, такие как эстетическая привлекательность. Существуют различные протоколы для проведения оценки людьми. Один из протоколов, описанный в [31], включает представление пользователям пар изображений, сгенерированных разными моделями. Оценщики-люди просят выбрать, какое изображение выглядит более фотореалистичным. Этот подход позволяет нам сравнивать модели по критериям, более тесно согласующимся с человеческими суждениями.
Метрики онлайн-оценки
На практике для обеспечения хорошей работы системы генерации изображений и соответствия ожиданиям пользователей обычно отслеживают различные метрики. Две распространённые метрики:
- Обратная связь пользователей: эта метрика жизненно важна, поскольку напрямую отражает мнения пользователей о сгенерированных изображениях. Обратную связь пользователей можно собирать через опросы, оценки или прямые комментарии.
- Задержка (latency): задержка — это время от момента запроса до полной генерации изображения и его доставки пользователю. Быстрое время отклика имеет решающее значение для поддержания хорошего пользовательского опыта, особенно в интерактивных приложениях. Мониторинг задержки помогает выявлять узкие места в производительности и гарантирует, что система соответствует ожиданиям пользователей.
Общий дизайн ML-системы
В этом разделе мы рассмотрим целостный дизайн реалистичной системы генерации лиц. Ключевые компоненты, которые мы изучим:
- Генератор лиц
- Сервис обучения
- Сервис оценки
- Сервис развёртывания
Генератор лиц
Генератор лиц — это основной компонент, отвечающий за создание реалистичных лиц. Он обрабатывает запросы пользователей и взаимодействует с обученной моделью GAN для выборки высококачественных изображений. Пользователи могут опционально указывать желаемые атрибуты, такие как возраст, пол, причёска и выражение лица. Сервис использует свойства StyleGAN для настройки вектора шума в latent space на основе этих атрибутов.
Архитектурные детали управления атрибутами и манипуляций с latent space обычно не обсуждаются на собеседованиях по проектированию ML-систем. Если вас интересует более глубокое изучение этой области, читайте [32][33].
Сервис обучения
Сервис обучения непрерывно совершенствует модель GAN, периодически переобучая её с помощью одобренных пользователями сгенерированных изображений и новых обучающих данных.
Сервис оценки
Сервис оценки автоматически оценивает вновь обученные модели. Он использует заранее определённые метрики для оценки их производительности. Результаты оценки определяют, соответствует ли новая модель стандартам качества и должна ли она заменить существующую.
Сервис развёртывания
Сервис развёртывания развёртывает улучшенные модели в производственной среде. Он обеспечивает плавный переход с минимальным простоем во время обновлений. Сервис развёртывания также отслеживает производительность развёрнутых моделей, чтобы убедиться, что они функционируют должным образом.
Дополнительные темы для обсуждения
Если после интервью останется дополнительное время, можно обсудить следующие темы:
- Различные архитектуры GAN, такие как DCGAN, WGAN и StyleGAN, и компромиссы каждой архитектуры [34][35][18].
- Методы стабилизации обучения GAN для избежания вырождения моды и проблем сходимости, включая использование Wasserstein loss, штрафа за градиент и других надёжных методов обучения [36].
- Использование условных GAN (cGAN) для генерации лиц на основе конкретных условий или входных данных [37].
- Метрики оценки согласованности условий [38][39].
- Смешение стилей в генерации лиц [18].
Резюме
Справочные материалы
[1] StyleGAN2. https://arxiv.org/abs/1912.04958. [2] Auto-Encoding Variational Bayes. https://arxiv.org/abs/1312.6114. [3] Generative Adversarial Networks. https://arxiv.org/abs/1406.2661. [4] Combating Mode Collapse in GAN training: An Empirical Analysis using Hessian Eigenvalues. https://arxiv.org/abs/2012.09673. [5] Google's GAN course. https://developers.google.com/machine-learning/gan/training. [6] StackGAN: Text to Photo-realistic Image Synthesis with Stacked Generative Adversarial Networks. https://arxiv.org/abs/1612.03242. [7] Zero-Shot Text-to-Image Generation. https://arxiv.org/abs/2102.12092. [8] Muse: Text-To-Image Generation via Masked Generative Transformers. https://arxiv.org/abs/2301.00704. [9] DALL·E 3. https://openai.com/index/dall-e-3/. [10] Attribute-specific Control Units in StyleGAN for Fine-grained Image Manipulation. https://arxiv.org/abs/2111.13010. [11] A guide to convolution arithmetic for deep learning. https://arxiv.org/abs/1603.07285. [12] Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. https://arxiv.org/abs/1502.03167. [13] Layer Normalization. https://arxiv.org/abs/1607.06450. [14] Instance Normalization: The Missing Ingredient for Fast Stylization. https://arxiv.org/abs/1607.08022. [15] Group Normalization. https://arxiv.org/abs/1803.08494. [16] Deep Learning using Rectified Linear Units (ReLU). https://arxiv.org/abs/1803.08375. [17] PyTorch's Tanh layer. https://pytorch.org/docs/stable/generated/torch.nn.Tanh.html. [18] A Style-Based Generator Architecture for Generative Adversarial Networks. https://arxiv.org/abs/1812.04948. [19] Minimax. https://en.wikipedia.org/wiki/Minimax. [20] Loss functions in GANs. https://developers.google.com/machine-learning/gan/loss. [21] Towards Principled Methods for Training Generative Adversarial Networks. https://arxiv.org/abs/1701.04862. [22] Unrolled Generative Adversarial Networks. https://arxiv.org/abs/1611.02163. [23] Stabilizing Training of Generative Adversarial Networks through Regularization. https://arxiv.org/abs/1705.09367. [24] Megapixel Size Image Creation using Generative Adversarial Networks. https://arxiv.org/abs/1706.00082v1. [25] Inception score. https://en.wikipedia.org/wiki/Inception_score. [26] GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium. https://arxiv.org/abs/1706.08500. [27] Demystifying MMD GANs. https://arxiv.org/abs/1801.01401. [28] Rethinking the Inception Architecture for Computer Vision. https://arxiv.org/abs/1512.00567. [29] FID calculation. https://en.wikipedia.org/wiki/Fr%C3%A9chet_inception_distance. [30] The Role of ImageNet Classes in Fréchet Inception Distance. https://arxiv.org/abs/2203.06026. [31] Hierarchical Text-Conditional Image Generation with CLIP Latents. https://arxiv.org/abs/2204.06125. [32] Alias-Free Generative Adversarial Networks. https://arxiv.org/abs/2106.12423. [33] StyleGAN3. https://nvlabs.github.io/stylegan3/. [34] Unsupervised Representation Learning with Deep Convolutional Generative Adversarial Networks. https://arxiv.org/abs/1511.06434. [35] Wasserstein GAN. https://arxiv.org/abs/1701.07875. [36] Stabilizing Generative Adversarial Networks: A Survey. https://arxiv.org/abs/1910.00927. [37] Conditional Generative Adversarial Nets. https://arxiv.org/abs/1411.1784. [38] CLIPScore: A Reference-free Evaluation Metric for Image Captioning. https://arxiv.org/abs/2104.08718. [39] DreamBooth: Fine Tuning Text-to-Image Diffusion Models for Subject-Driven Generation. https://arxiv.org/abs/2208.12242.
Сноски
- Шаг (stride) управляет тем, насколько фильтр перемещается по входным данным во время convolution — большие шаги пропускают больше пикселей ↩
- Отступ (padding) добавляет дополнительные границы вокруг входных данных для управления размером выходных данных во время convolution ↩
Синтез изображений высокого разрешения
Введение
Генеративный ИИ обладает удивительной способностью создавать высокореалистичные и разнообразные изображения. В этой главе мы рассматриваем технику, позволяющую генерировать детализированные и разнообразные изображения всего за несколько секунд.
Уточнение требований
Вот типичный диалог между кандидатом и интервьюером:
Кандидат: Должна ли система с самого начала фокусироваться на определённых категориях изображений? Интервьюер: Для простоты начнём с природных пейзажей и городских ландшафтов. Другие категории можно будет добавить позже.
Кандидат: Есть ли у нас обучающие данные с природными пейзажами? Каков размер набора данных? Интервьюер: У нас есть большой набор данных, содержащий около 5 миллионов изображений высокого разрешения с природными пейзажами и ландшафтами.
Кандидат: Должна ли система поддерживать дополнительные условия, например текстовый prompt, описывающий желаемое изображение? Интервьюер: Хороший вопрос. Мы сосредоточимся на генерации изображений без входных условий. Однако система должна быть гибкой и поддерживать входные prompt-ы.
Кандидат: На какой диапазон разрешений следует ориентироваться при генерации изображений? Интервьюер: Система должна генерировать изображения с разрешением 1024×\times×1024 или 2048×\times×2048 пикселей по запросу пользователя.
Кандидат: Изображения должны генерироваться в реальном времени или допустима некоторая задержка? Интервьюер: Генерация в реальном времени не обязательна. Однако важно обеспечить разумное время обработки. Давайте ориентироваться на пять секунд на изображение.
Формулировка задачи как задачи ML
Определение входных и выходных данных системы
Для синтеза изображений высокого разрешения пользователь просто запрашивает новое изображение. Выходными данными является изображение высокого разрешения.
Выбор подходящего подхода ML
Как обсуждалось в Главе 7, существует несколько подходов к генерации изображений, включая VAE, GAN, авторегрессионные модели и diffusion-модели. В этом разделе мы выбираем наиболее подходящий для данной задачи.
Большинство вариантов VAE и GAN с трудом справляются с генерацией изображений высокого разрешения, например с разрешением 512×512 пикселей и выше. Они сталкиваются с проблемой, известной как posterior collapse (коллапс апостериорного распределения). Это происходит потому, что по мере увеличения разрешения этим моделям требуется decoder с большей ёмкостью для захвата дополнительных деталей. В процессе обучения decoder может стать настолько мощным, что начнёт игнорировать входные данные из latent space, поскольку способен моделировать выходные данные самостоятельно. В результате латентные переменные почти не вносят вклада в процесс генерации, снижая разнообразие изображений.
Хотя как авторегрессионные, так и diffusion-модели способны генерировать изображения высокого разрешения, они существенно различаются по сложности и требованиям к ресурсам. Авторегрессионные модели часто считаются медленными из-за своей последовательной природы, где каждый пиксель зависит от ранее сгенерированных. Эта зависимость приводит к временной сложности, линейно возрастающей с количеством пикселей, что даёт сложность O(N2)O(N^2)O(N2) для изображения размером N×NN\times NN×N, и процесс сложно распараллелить. Чтобы решить эту проблему, авторегрессионные модели генерируют изображения по фрагментам, а не попиксельно. Например, генерация изображения 1024×\times×1024 с использованием фрагментов 64×\times×64 пикселей требует всего 256 шагов или token-ов, что значительно снижает вычислительные затраты по сравнению с традиционными попиксельными методами.
С другой стороны, сложность diffusion-моделей возрастает сверхлинейно с размером изображения, что даёт вычислительную сложность O(TN2)O(TN^2)O(TN2), где NNN — количество пикселей, а T — количество шагов удаления шума. Более крупные изображения часто требуют большего числа шагов уточнения для поддержания качества и согласованности, что ещё больше увеличивает вычислительную нагрузку.
На практике генерация изображения высокого разрешения с использованием стандартных diffusion-моделей может занимать несколько минут1. Напротив, авторегрессионные модели на основе Transformer справляются с аналогичными задачами за секунды благодаря своему подходу к генерации по фрагментам. В этой главе мы рассматриваем авторегрессионные модели в образовательных целях. В Главе 9 мы подробно рассмотрим diffusion-модели. Теперь давайте углубимся в авторегрессионные модели и их ключевые компоненты.
Авторегрессионные модели генерируют изображения, рассматривая их как задачу генерации последовательностей. Этот подход опирается на два основных компонента:
- Токенизатор изображений
- Генератор изображений
Токенизатор изображений
Токенизация изображений означает представление изображения в виде последовательности дискретных token-ов. Это ключевой элемент авторегрессионных моделей, где изображение генерируется последовательно, фрагмент за фрагментом.
Токенизатор изображений — это отдельная модель, обучаемая независимо. Его основные функции — кодировать изображение в последовательность дискретных token-ов и декодировать последовательность дискретных token-ов обратно в изображение.
Генератор изображений
Генератор изображений — это основная модель для генерации изображений фрагмент за фрагментом. Среди различных архитектур для генерации последовательностей decoder-only Transformer является наиболее эффективным выбором по двум причинам. Во-первых, decoder-only Transformer обладает гибкой архитектурой, способной работать с различными модальностями. В чат-боте он принимает текстовые token-ы на вход и генерирует текстовые token-ы на выходе. В задаче подписей к изображениям он принимает изображение на вход и выдаёт текстовые token-ы. Для генерации изображений он производит последовательность token-ов изображений на выходе, которые затем декодируются в изображение.
Во-вторых, архитектура Transformer эффективно улавливает дальние зависимости благодаря механизму attention, что полезно для генерации согласованных изображений.
В итоге мы подходим к генерации изображений с помощью авторегрессионной модели на основе Transformer. Сначала генератор изображений (decoder-only Transformer) генерирует последовательность дискретных token-ов. Затем токенизатор изображений декодирует эти token-ы в итоговое изображение. Мы подробно изучим архитектуру, обучение и процессы сэмплирования этих компонентов в разделе о разработке модели.
Подготовка данных
Процесс подготовки данных включает два ключевых шага:
- Очистка и нормализация изображений
- Токенизация изображений
Очистка и нормализация изображений
На этом шаге мы удаляем низкокачественные изображения из обучающих данных и обеспечиваем согласованность оставшихся. Это достигается с помощью следующих операций:
- Удаление низкокачественных изображений: мы удаляем изображения с низким разрешением, избыточным шумом или нерелевантным содержимым. Также мы обеспечиваем, чтобы набор данных охватывал широкий спектр стилей, объектов и композиций. Этот шаг критически важен для того, чтобы генеративная модель могла создавать разнообразные изображения высокого качества.
- Нормализация изображений: нормализация предполагает масштабирование значений пикселей до определённого диапазона, как правило от 0 до 1, для стабилизации процесса обучения.
- Изменение размера изображений: изображения часто имеют разные размеры и соотношения сторон. Изменение размера до единого значения обеспечивает получение моделью согласованных входных данных. В соответствии с требованиями интервьюера мы изменяем размер всех изображений до 1024×1024.
Токенизация изображений
Генератор изображений требует представления изображений в виде последовательности дискретных token-ов. Для этого, после обучения токенизатора изображений, мы токенизируем все изображения в нашем обучающем наборе данных в дискретные token-ы. Важно отметить, что этот шаг подготовки данных предназначен прежде всего для генератора изображений, а не для токенизатора.
Эти два шага обеспечивают высокое качество, согласованность обучающих данных и их представление в виде последовательности числовых входных данных.
Разработка модели
Архитектура
В этом разделе мы рассматриваем архитектуру как токенизатора изображений, так и генератора изображений.
Токенизатор изображений
Модель токенизатора изображений выполняет две функции:
- Кодирование изображения в последовательность дискретных token-ов
- Декодирование последовательности дискретных token-ов обратно в изображение
Распространённая архитектура, специально разработанная для токенизации изображений, — это Vector-Quantized VAE (VQ-VAE) [2], вариант стандартного VAE, рассмотренного в Главе 7. VQ-VAE состоит из трёх компонентов:
- Encoder
- Quantizer
- Decoder
Encoder
Encoder отображает входное изображение в низкоразмерный latent space. Этот компонент кодирует важные признаки изображения в закодированное представление.
Архитектура encoder — это глубокая сверточная нейронная сеть (CNN) с несколькими сверточными слоями, каждый из которых сопровождается функцией активации ReLU [3]. Эти слои обрабатывают входное изображение и извлекают визуальные признаки.
Quantizer
Quantizer преобразует непрерывные латентные векторы в дискретные token-ы. Есть две основные причины, по которым VQ-VAE вводит компонент quantizer в стандартный VAE:
- Предотвращение posterior collapse
- Сокращение пространства обучения
Предотвращение posterior collapse
Posterior collapse — распространённая проблема в стандартных VAE, когда латентные переменные почти не вносят вклада или игнорируются, поскольку decoder генерирует точные выходные данные без использования latent space. Шаг квантизации устраняет эту проблему путём дискретизации латентных переменных, тем самым вынуждая модель использовать их при реконструкции. Это гарантирует, что decoder не подавляет latent space и латентные переменные активно участвуют в формировании выходных данных.
Сокращение пространства обучения
Непрерывные векторы сложно предсказывать последовательно, поскольку они имеют бесконечное множество возможностей и малые различия. Преобразуя эти векторы в дискретные token-ы, quantizer упрощает процесс, позволяя Transformer сосредоточиться на меньшем числе вариантов.
Quantizer использует внутреннюю кодовую книгу для преобразования непрерывных латентных векторов в дискретные token-ы. Эта кодовая книга содержит обучаемые embedding-и, представляющие различные паттерны во входных изображениях. Каждый embedding выступает token-ом, представленным целым числом от 1 до k. Quantizer заменяет каждый непрерывный вектор ближайшим token-ом в кодовой книге на основе евклидова расстояния [4].
Обратите внимание, что quantizer — это таблица embedding-ов. Его единственный параметр — кодовая книга, которая изучается в процессе обучения. Единственная ответственность quantizer — отображать каждый непрерывный вектор на ближайший token в кодовой книге; поэтому на выходе получается набор идентификаторов token-ов.
Decoder
Decoder преобразует дискретные token-ы обратно в исходное изображение. Обычно он использует глубокую CNN с транспонированными свёртками (ConvTranspose2d) для постепенного преобразования представления к исходному размеру изображения. Подробнее о свёртках и транспонированных свёртках см. в [5].
Генератор изображений
Генератор изображений создаёт последовательность дискретных token-ов, представляющих изображение. Как упоминалось ранее, для задач генерации последовательностей часто используется decoder-only Transformer, который включает следующие компоненты:
- Embedding lookup: заменяет каждый дискретный token его embedding-ом из кодовой книги.
- Проекция: проецирует каждый embedding token-а в размерность, соответствующую внутреннему представлению Transformer.
- Позиционное кодирование: добавляет позиционные кодировки к последовательности для предоставления пространственной информации.
- Transformer: обрабатывает входную последовательность и выдаёт обновлённую последовательность векторов.
- Предсказывающая голова: использует обновлённые embedding-и для предсказания следующего token-а.
Обучение
В авторегрессионной генерации изображений выделяют два этапа обучения:
- Этап I: Обучение токенизатора изображений
- Этап II: Обучение генератора изображений
Этап I: Обучение токенизатора изображений
Процесс обучения предполагает оптимизацию encoder, decoder и кодовой книги, чтобы модель могла точно реконструировать исходные изображения. Этот процесс можно описать в три шага:
- Encoder обрабатывает входное изображение и преобразует его в непрерывное представление.
- Quantizer заменяет непрерывное представление дискретными token-ами с помощью своей внутренней кодовой книги.
- Decoder использует дискретные token-ы для реконструкции исходного изображения.
Поскольку операция поиска в quantizer не имеет чётко определённого градиента для обратного распространения ошибки, в статье VQ-VAE предлагается аппроксимировать градиент путём его копирования непосредственно с входа decoder на выход encoder. Этот подход означает, что только выбранные token-ы получают градиенты от decoder, тогда как невыбранные token-ы не получают никаких градиентов.
Обучающие данные
Мы обучаем токенизатор изображений на 5 миллионах изображений. Поскольку обучение самонаблюдаемое и не требует меток изображений, мы включаем другие общедоступные наборы данных изображений для повышения устойчивости токенизатора. В частности, мы используем набор данных LAION-400M [6], содержащий 400 миллионов изображений. Это даёт более богатую кодовую книгу, охватывающую разнообразные визуальные паттерны.
Цель ML и функция потерь
Цель ML токенизатора изображений — точно реконструировать исходные изображения из их квантизованных token-ов. Для достижения этой цели в процессе обучения обычно используются следующие функции потерь:
- Потери реконструкции
- Потери квантизации
Потери реконструкции: потери реконструкции измеряют разницу между исходным изображением и его реконструкцией из квантизованных token-ов. Обычно они рассчитываются по формуле среднеквадратичной ошибки (MSE):
Где:
- xix_ixi — значение пикселя исходного изображения,
- x^i\hat{x}_ix^i — значение пикселя реконструированного изображения,
- nnn — общее количество пикселей в изображении.
Потери квантизации: потери квантизации измеряют расстояние между выходными данными encoder и ближайшим embedding-ом в кодовой книге. Эти потери побуждают encoder производить выходные данные, более близкие к embedding-ам кодовой книги.
Где:
- E(x)E(x)E(x) — непрерывный латентный вектор, произведённый encoder E из входных данных xxx,
- zqz_qzq — квантизованный латентный вектор, выбранный из кодовой книги ZZZ,
- sg(.)\operatorname{sg}(.)sg(.) представляет операцию остановки градиента, которая блокирует прохождение градиентов через данный член. Здесь она используется для предотвращения обновления кодовой книги при оптимизации encoder.
Для получения дополнительных сведений о формуле потерь квантизации обратитесь к статье VQGAN [1].
На практике использование как потерь реконструкции, так и потерь квантизации в процессе обучения хорошо работает для реконструкции изображений низкого разрешения. Однако для изображений высокого разрешения модель всё ещё может давать артефакты. Для улучшения качества реконструкции при высоких разрешениях обычно применяются две дополнительные функции потерь:
- Перцептивные потери
- Состязательные потери
Перцептивные потери: перцептивные потери измеряют разницу между признаками исходного и реконструированного изображений, извлечёнными из определённого слоя предобученной модели, например VGG [7]. Формула:
Где:
- ϕl\phi_lϕl обозначает карту признаков слоя l из предобученной модели VGG,
- xxx — исходное изображение,
- x^\hat{x}x^ — реконструированное изображение.
Перцептивные потери побуждают модель реконструировать изображения, перцептивно схожие с исходными. Признаки VGG кодируют высокоуровневые детали, такие как содержание и стиль. Перцептивные потери направляют процесс обучения таким образом, чтобы модель лучше сохраняла эти детали в реконструированных изображениях.
Состязательные потери: состязательные потери заимствованы из GAN [8], где дискриминатор пытается отличить реальные изображения от реконструированных. Эти потери используются для измерения того, насколько хорошо реконструированное токенизатором изображение может обмануть дискриминатор. Формула, как мы видели в Главе 7:
Где:
- DDD — сеть дискриминатора,
- x^\hat{x}x^ — реконструированное изображение.
Эта функция потерь побуждает модель создавать реконструированные изображения, которые обученный дискриминатор не может отличить от реальных. В статье VQGAN была введена патч-основанная версия этих потерь для уменьшения неестественных артефактов и улучшения реалистичности реконструкций.
Общие потери: общая функция потерь часто представляет собой взвешенную сумму отдельных потерь, описанных выше. Веса (λi)(\lambda_i)(λi) являются гиперпараметрами, которые необходимо подбирать в зависимости от конкретных целей производительности и экспериментов.
После обучения токенизатора изображений мы преобразуем все 5 миллионов обучающих изображений в дискретные token-ы и кэшируем их, как подробно описано в разделе подготовки данных. Этот шаг гарантирует, что все изображения представлены в виде последовательности дискретных token-ов, что необходимо для обучения генератора изображений.
Этап II: Обучение генератора изображений
Обучение генератора изображений, представляющего собой decoder-only Transformer, аналогично процессу, описанному в предыдущих главах. Обучающие данные состоят из последовательностей дискретных token-ов, и модель учится предсказывать эти token-ы последовательно в процессе обучения.
В качестве цели ML мы используем предсказание следующего token-а, а в качестве функции потерь — кросс-энтропию для измерения точности предсказанных вероятностей по сравнению с правильными визуальными token-ами.
Сэмплирование
В авторегрессионных моделях генерация нового изображения включает два шага:
- Генерация последовательности дискретных token-ов
- Декодирование дискретных token-ов в изображение
. Генерация последовательности дискретных token-ов
На первом шаге генератор изображений производит последовательность token-ов. Авторегрессионная природа генерации гарантирует, что каждый token обусловлен предшествующими token-ами, что обеспечивает создание согласованных изображений.
Вот пошаговый процесс авторегрессивной генерации последовательности token-ов:
- Случайным образом выбирается token из кодовой книги в качестве первого token-а. Этот начальный token служит отправной точкой для всего процесса генерации.
- Авторегрессивно генерируются token-ы один за другим. Это включает: Передачу текущей последовательности token-ов генератору изображений для предсказания распределения вероятностей по кодовой книге Выбор следующего token-а с помощью метода сэмплирования, например top-p sampling Добавление выбранного token-а к текущей последовательности
Этот процесс продолжается до тех пор, пока не будет сгенерировано всё изображение. Количество итераций зависит от разрешения и размера желаемого выходного изображения. Например, для генерации изображения размером 1024×\times×1024 пикселей, где каждый визуальный token представляет блок 64×\times×64 пикселя, требуется 256 token-ов. Процесс продолжается до тех пор, пока не будут сгенерированы все 256 token-ов. После завершения последовательности token-ов она преобразуется в реальное изображение, о чём рассказывается на следующем шаге.
. Декодирование дискретных token-ов в изображение
На этом шаге последовательность дискретных token-ов преобразуется в изображение с использованием функции декодирования токенизатора изображений.
Оценка
Метрики оценки для синтеза изображений высокого разрешения схожи с теми, что рассматривались в Главе 7. В этом разделе мы кратко рассмотрим их, не вдаваясь в детали.
Метрики офлайн-оценки
Для измерения качества и разнообразия сгенерированных изображений обычно используются следующие метрики:
- Inception score: измеряет сходство сгенерированных изображений с изображениями реальных объектов с помощью предобученной модели Inception v3. Подробнее об Inception score см. в [9].
- Расстояние Фреше Inception (FID): сравнивает распределение сгенерированных изображений с реальными путём сравнения признаков, извлечённых из предобученной модели Inception v3. Эта метрика измеряет схожесть статистических характеристик сгенерированных и реальных изображений. Подробнее о FID см. в [10].
- Оценка людьми: оценщики-люди получают пары изображений и оценивают их фотореализм и эстетические качества. Голоса дают статистическую меру того, какие модели со временем создают более реалистичные изображения.
Помимо этих метрик, принято оценивать и другие аспекты модели, такие как задержка и стоимость.
- Время генерации изображения: измеряет время, которое модель тратит на генерацию изображения. Эта метрика важна для мониторинга, поскольку пользователи, как правило, ожидают быстрых результатов.
- Стоимость генерации: рассчитывает стоимость генерации одного изображения. Эта метрика зависит от таких факторов, как сложность модели, разрешение и инфраструктурные расходы. Мониторинг стоимости генерации критически важен, поскольку влияет на доходы бизнеса.
Метрики онлайн-оценки
На практике компании отслеживают различные метрики для оценки качества системы в реальном времени. Распространённые метрики включают:
- Обратная связь от пользователей: сбор прямых отзывов пользователей о сгенерированных изображениях.
- Периодические опросы: сбор мнений пользователей о качестве и релевантности сгенерированных изображений.
- Коэффициент подписок: измеряет, как часто пользователи подписываются на сервисы или функции, связанные с генерацией изображений.
- Коэффициент оттока: измеряет долю пользователей, прекративших использование сервиса.
Общий дизайн системы ML
Когда мы удовлетворены производительностью моделей генератора и токенизатора изображений, мы можем интегрировать их для создания системы синтеза изображений. Основные компоненты системы синтеза изображений высокого разрешения:
- Сервис генерации
- Сервис декодирования
- Сервис суперразрешения
Понимание назначения каждого компонента и их взаимодействий даст целостное представление о системе. Рассмотрим каждый из них подробнее.
Сервис генерации
Сервис генерации обрабатывает запросы пользователей и взаимодействует с обученной моделью генератора изображений для создания последовательности визуальных token-ов.
Сервис декодирования
Сервис декодирования взаимодействует с токенизатором изображений для преобразования сгенерированной последовательности визуальных token-ов в изображение. Обратите внимание, что при развёртывании модели encoder в токенизаторе изображений не нужен — он используется только в процессе обучения.
Разделение сервисов генерации и декодирования критически важно, поскольку генератор изображений и токенизатор — это разные модели с различными вычислительными потребностями и задержками. Такой подход позволяет каждому сервису масштабироваться независимо и эффективно управлять ресурсами.
Сервис суперразрешения
Сервис суперразрешения использует предобученную модель для увеличения разрешения сгенерированных изображений. Например, если желаемое разрешение — 2048×\times×2048, а генератор выдаёт только 1024×\times×1024, мы используем модель суперразрешения с коэффициентом масштабирования 2x.
Этот сервис критически важен для приложений, требующих детализированного и реалистичного визуального контента, например в медицинской визуализации. Существует множество устоявшихся решений для суперразрешения: от CNN-основанных [11] до улучшенных GAN [12]. Подробнее о современных подходах см. в [13].
Дополнительные темы для обсуждения
Если в конце интервью осталось время, можно рассмотреть следующие дополнительные темы:
- Расширение авторегрессионных моделей для поддержки текстовой генерации [14] [15].
- Поддержка приложений, таких как дополнение изображений и суперразрешение изображений [16].
- Балансировка разнообразия и точности при сэмплировании с использованием таких техник, как температурное масштабирование [17].
- Повышение стабильности с помощью состязательного обучения, обрезки градиентов и планирования скорости обучения [18][19].
- Использование прогрессивного роста и многомасштабных архитектур для улучшения качества и детализации изображений [20].
- Создание интерактивных систем, позволяющих пользователям уточнять и настраивать сгенерированные изображения [21].
Резюме
Список литературы
[1] Taming Transformers for High-Resolution Image Synthesis. https://arxiv.org/abs/2012.09841. [2] Neural Discrete Representation Learning. https://arxiv.org/abs/1711.00937. [3] Deep Learning using Rectified Linear Units (ReLU). https://arxiv.org/abs/1803.08375. [4] Euclidean distance. https://en.wikipedia.org/wiki/Euclidean_distance. [5] A guide to convolution arithmetic for deep learning. https://arxiv.org/abs/1603.07285. [6] LAION data set 400 million https://laion.ai/blog/laion-400-open-dataset/. [7] Very Deep Convolutional Networks for Large-Scale Image Recognition. https://arxiv.org/abs/1409.1556. [8] Generative Adversarial Networks. https://arxiv.org/abs/1406.2661. [9] Inception score. https://en.wikipedia.org/wiki/Inception_score. [10] FID calculation. https://en.wikipedia.org/wiki/Fr%C3%A9chet_inception_distance. [11] Image Super-Resolution Using Very Deep Residual Channel Attention Networks. https://arxiv.org/abs/1807.02758. [12] ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks. https://arxiv.org/abs/1809.00219. [13] NTIRE 2024 Challenge on Image Super-Resolution (×4): Methods and Results. https://arxiv.org/abs/2404.09790. [14] Muse: Text-To-Image Generation via Masked Generative Transformers. https://arxiv.org/abs/2301.00704. [15] VQGAN-CLIP: Open Domain Image Generation and Editing with Natural Language Guidance. https://arxiv.org/abs/2204.08583. [16] LAR-SR: A Local Autoregressive Model for Image Super-Resolution. https://openaccess.thecvf.com/content/CVPR2022/papers/ Guo_LAR-SR_A_Local_Autoregressive_Model_for_Image_Super-Resolution_CVPR_2022_paper.pdf. [17] Long Horizon Temperature Scaling. https://arxiv.org/abs/2302.03686. [18] Learning Rate Scheduling. https://d2l.ai/chapter_optimization/lr-scheduler.html. [19] Adversarial Training. https://adversarial-ml-tutorial.org/adversarial_training/. [20] Progressive Growing of GANs for Improved Quality, Stability, and Variation. https://arxiv.org/abs/1710.10196. [21] CogView2: Faster and Better Text-to-Image Generation via Hierarchical Transformers. https://arxiv.org/abs/2204.14217.
Сноски
- Определённые оптимизации и техники (например, latent diffusion model) могут значительно ускорить процесс генерации в diffusion-моделях. Эти методы подробно рассматриваются в Главах 10 и 11. ↩
Генерация изображений по тексту
Введение
Во многих случаях вместо того, чтобы позволить модели генерировать контент из случайного шума (как обсуждалось в главах 7 и 8), мы хотим управлять содержимым сгенерированного изображения. Генерация изображений по тексту (text-to-image) — это увлекательное приложение генеративного ИИ, позволяющее пользователям вводить текстовый prompt, который модель преобразует в детализированное изображение. Несколько коммерческих text-to-image сервисов уже доступны на рынке: DALL-E 3 от OpenAI [1], Imagen от Google [2] и Firefly от Adobe [3].
Уточнение требований
Ниже приведён типичный диалог между кандидатом и интервьюером:
Кандидат: На какое разрешение мы ориентируемся для сгенерированных изображений? Интервьюер: Мы стремимся к высокому разрешению — конкретно 1024x1024 пикселей.
Кандидат: Должна ли система поддерживать несколько языков ввода или только английский? Интервьюер: Сначала сосредоточимся на английском, но архитектура системы должна быть адаптируема для других языков в будущем.
Кандидат: Каков размер набора данных для обучения text-to-image модели? Интервьюер: У нас около 500 миллионов изображений из пользовательских активов, большинство с подписями.
Кандидат: Насколько детальными и сложными могут быть текстовые prompt'ы? Есть ли ограничения по сложности или длине? Интервьюер: Система должна обрабатывать детальные текстовые prompt'ы с максимальной длиной 128 слов.
Кандидат: Какую скорость генерации изображений должна достигать система? Интервьюер: Цель — генерация близкая к реальному времени. Ориентируемся на 10 секунд на изображение.
Кандидат: Какие типы изображений должна генерировать система? Мы сосредоточены на конкретной области, например пейзажи? Интервьюер: Система должна быть способна генерировать на основе текстовых prompt'ов широкий спектр изображений, включая реалистичные пейзажи, портреты и абстрактное или концептуальное искусство.
Кандидат: Важно обеспечить, чтобы изображения не были предвзятыми по возрасту, расе или полу. Могу ли я начать с фокуса на этих трёх атрибутах? Интервьюер: Отличное замечание. Справедливая система крайне важна. Начнём с решения вопросов по этим трём атрибутам.
Кандидат: Этические соображения критически важны. Нам нужны фильтры и проверки, чтобы не генерировать оскорбительные, неуместные или вредоносные изображения. Это верно? Интервьюер: Да, всё верно.
Формулировка задачи как ML-задачи
Определение вхо да и выхода системы
Вход системы — текстовый prompt, предоставленный пользователем, описывающий желаемое изображение. Такой prompt обычно включает детали: сцены, объекты, цвета, стили и эмоции.
Выход — визуально детализированное изображение, соответствующее текстовому prompt'у. Например, как показано на рисунке 2, prompt «Лодка в океане» создаёт изображение, изображающее эту сцену.
Выбор подходящего ML-подхода
Генерация изображений по тексту — это мультимодальная задача, включающая понимание текста и генерацию соответствующего изображения. Существует два основных подхода для построения text-to-image систем:
- Авторегрессионные модели
- Diffusion модели
Кратко рассмотрим каждый из них и выберем наиболее подходящий для наших нужд.
Авторегрессионные модели
Эти модели рассматривают генерацию text-to-image как задачу генерации последовательности. Decoder-only Transformer принимает последовательность текстовых token'ов на вход и выдаёт последовательность визуальных token'ов, представляющих изображение. Затем токенизатор изображений декодирует эти визуальные token'ы в фактическое изображение.
На основе этого подхода разработано несколько text-to-image моделей, таких как DALL-E от OpenAI [5] и Muse от Google [6].
Diffusion модели
Впервые представленные в 2019 году [7], diffusion модели привлекли широкое внимание примерно через три года. Они используют другой подход к генерации text-to-image: начинают со случайного шума и постепенно преобразуют его в чёткое изображение на основе текстового prompt'а. Этот процесс обычно включает text encoder, такой как CLIP от OpenAI [8] или T5 от Google [9], который преобразует текстовый prompt в embedding. Это embedding захватывает смысл prompt'а и направляет diffusion модель для генерации соответствующих изображений.1
Примеры text-to-image моделей на основе diffusion: Imagen 3 от Google [2], DALL-E 2 от OpenAI [10] и Stable Diffusion от Stability AI [11].
Diffusion против авторегрессионных моделей
Авторегрессионные модели рассматривают генерацию text-to-image как задачу генерации последовательности, тогда как diffusion модели подходят к ней как к итеративному процессу уточнения. Это ключевое различие в моделировании влияет на их возможности.
Оба типа моделей — diffusion и авторегрессионные — могут генерировать реалистичные изображения, медленно работают при генерации, как правило имеют миллиарды параметров и требуют значительных вычислительных ресурсов для обучения. Несмотря на эти сходства, они различаются в трёх ключевых аспектах:
- Сложность реализации: Авторегрессионные модели проще в реализации как при обучении, так и при inference. При обучении они статистически эффективнее, поскольку могут получать полезные сигналы градиента от всех шагов за один прямой-обратный проход. Напротив, diffusion модели менее статистически эффективны — они требуют сэмплирования различных уровней шума для каждого обучающего примера. При inference, как только Transformer в авторегрессионной модели генерирует последовательность визуальных token'ов, эти token'ы образуют итоговое изображение. Diffusion модели, однако, уточняют изображение через множество шагов, что усложняет реализацию.
- Качество изображений: Diffusion модели демонстрируют лучшие результаты при генерации высокодетализированных и реалистичных изображений. Их итеративный процесс позволяет модели непрерывно уточнять и улучшать мелкие детали, обеспечивая превосходный реализм сгенерированных изображений.
- Гибкость в сэмплировании: Diffusion модели более гибко балансируют между скоростью сэмплирования и качеством изображений. Они легко регулируют количество шагов сэмплирования — больше шагов обычно даёт более высокое качество, но занимает больше времени. Обученная авторегрессионная модель не может так просто вносить подобные коррективы.
В этой главе мы выбираем diffusion модели, отдавая приоритет исключительному качеству изображений. В разделе разработки модели мы рассмотрим архитектуру, методы обучения и техники сэмплирования diffusion моделей.
Подготовка данных
Наш набор данных состоит примерно из 500 миллионов пар изображение–подпись. Однако крупномасштабные наборы данных часто требуют значительной предобработки перед использованием в обучении модели. В этом разделе мы рассмотрим распространённые методы подготовки изображений и подписей для diffusion обучения.
Подготовка изображений
Мы сосредоточимся на двух основных этапах подготовки изображений: фильтрации неподходящих изображений и стандартизации оставшихся. Рассмотрим каждый этап подробнее.
Фильтрация неподходящих изображений
В крупных наборах данных многие изображения могут быть бесполезны для обучения. Важно удалить их, чтобы модель обучалась только на качественных и безопасных данных. Для этого выполняем следующие шаги:
- Удаление мелких изображений: Отбрасываем пары изображение–подпись, где размер изображения меньше определённого порога, например 64×64 пикселей. Мелкие изображения часто низкого качества и могут не предоставлять ценной информации для обучения.
- Дедупликация изображений: Используем методы дедупликации, такие как [12], для удаления идентичных или перцептуально похожих изображений. Это предотвращает смещение модели в сторону определённых изображений, встречающихся чаще.
- Удаление неподходящих изображений: Применяем модели обнаружения вредоносного контента и NSFW (Not Safe For Work) для фильтрации вредоносного контента, такого как насилие или обнажённость. Это гарантирует, что модель не научится генерировать неподходящие изображения.
- Удаление эстетически неприемлемых изображений: Используем специализированные ML-модели для исключения изображений с низкими эстетическими показателями. Это помогает модели сосредоточиться на качественных изображениях в процессе обучения.
Стандартизация изображений
- Корректировка размеров изображений: Diffusion модель требует входных данных определённых размеров. Поэтому важно иметь обучающие данные схожих размеров. Например, если ожидаемый вход модели — 128x128, сначала изменяем размер изображений так, чтобы меньшее измерение стало 128, сохраняя пропорции. Затем обрезаем по центру до финального размера 128x128.
- Нормализация изображений: Нормализуем значения пикселей до стандартного диапазона, например [0, 1] или [-1, 1], для более стабильного обучения.
Подготовка подписей
Помимо изображений, подписи часто нерелевантны или отсутствуют. Вот распространённые шаги для обеспечения согласованности и высокого качества подписей:
- Обработка отсутствующих или неанглийских подписей: Для изображений без подписей или с подписями на другом языке используем модель создания подписей к изображениям, такую как BLIP-3 [13], для автоматической генерации описательных подписей. Если хотите построить систему создания подписей с нуля, обратитесь к Главе 5.
- Улучшение подписей: Используем предобученную модель, такую как CLIP [8], для оценки релевантности каждой пары изображение–подпись. Для пар с оценкой ниже порога заменяем исходную подпись автоматически сгенерированной с помощью модели BLIP-3.
- Удаление плохо совпадающих пар: После улучшения подписей удаляем пары изображение–подпись с оценками CLIP similarity ниже порогового значения. Этот шаг гарантирует, что модель получает только пары, чьи подписи точно описывают изображения.
Разработка модели
Архитектура
Diffusion модель, как объяснялось ранее, постепенно денойзит зашумлённое изображение через множество шагов, пока оно не станет чётким. На каждом шаге, как показано на рисунке 7, модель принимает зашумлённое изображение на вход и предсказывает шум, который нужно удалить.2 Для этого обычно используются две распространённые архитектуры:
- U-Net
- Diffusion Transformer (DiT)
U-Net
U-Net [14] — это архитектура свёрточной нейронной сети (CNN), изначально разработанная для сегментации биомедицинских изображений. Она состоит из серии блоков понижения дискретизации (downsampling), за которыми следуют блоки повышения дискретизации (upsampling), как показано на рисунке 8.
Блоки downsampling
Блоки downsampling постепенно уменьшают пространственные размеры (высоту и ширину), увеличивая глубину (количество каналов), что приводит к сжатому представлению входных данных. Каждый блок downsampling обычно состоит из следующих компонентов:
- Операция свёртки: Извлекает визуальные признаки из входных данных.
- Batch normalization: Нормализует карты признаков для стабилизации обучения.
- Нелинейная активация: Вносит нелинейность для обучения сложным паттернам.
- Max-pooling: Уменьшает размеры карты признаков.
- Cross-attention: Обращается к дополнительным условиям, таким как token'ы текстового prompt'а. Это необходимо для обеспечения влияния текстового prompt'а на предсказанный шум.
Рассмотрим слой cross-attention подробнее, поскольку это первый случай, когда мы применяем cross-attention между разными модальностями — входными данными изображения и текста. Текстовый prompt обрабатывается text encoder'ом, таким как модель на основе Transformer, который преобразует слова или token'ы в последовательность непрерывных embedding'ов. Эти embedding'и захватывают семантический смысл текста. На каждом шаге денойзинга в diffusion процессе модель получает зашумлённое изображение и обрабатывает его через слои Conv2D и BatchNorm2D для извлечения визуальных признаков. В слое cross-attention запросы (queries) формируются из признаков изображения, а ключи (keys) и значения (values) поступают из текстовых embedding'ов. Это позволяет модели эффективно выравнивать и интегрировать информацию из текста в признаки изображения.
Блоки upsampling
Блоки upsampling симметрично увеличивают пространственные размеры и уменьшают глубину карты признаков. Итоговый выход соответствует исходному размеру входных данных — в данном случае это предсказанный шум. Каждый блок upsampling состоит из следующих компонентов:
- Транспонированная свёртка: Использует операции вроде ConvTranspose2D из PyTorch для обработки и увеличения размеров карты признаков.
- Batch normalization: Нормализует карты признаков для стабилизации обучения.
- Нелинейная активация: Вносит нелинейность для обучения сложным паттернам.
- Cross-attention: Сохраняет влияние дополнительных условий при upsampling.
Архитектура U-Net имеет множество деталей и вариаций. Разные реализации могут использовать различные слои и конфигурации. Однако понимания ключевых компонентов и структуры, как правило, достаточно для большинства интервью по проектированию ML-систем. Для более углублённой информации обратитесь к [14].
DiT
DiT [15] — ещё одна популярная архитектура в diffusion моделях. В отличие от U-Net, которая использует серию слоёв downsampling и upsampling, DiT в основном опирается на архитектуру Transformer для обработки зашумлённого входного изображения и предсказания шума.
DiT в основном вдохновлён архитектурой Vision Transformer (ViT) [16], обсуждавшейся в Главе 5. Компоненты DiT:
- Patchify: Преобразует входное изображение в последовательность patch embedding'ов.
- Positional encoding: Прикрепляет информацию о позиции к каждому patch embedding'у, указывая его расположение в исходном изображении.
- Transformer: Обрабатывает последовательность embedding'ов и другие conditioning сигналы, такие как текстовый prompt, для предсказания шума для каждого patch'а.
- Unpatchify: Преобразует последовательность предсказанных векторов шума обратно в изображение с теми же размерами, что и исходное входное изображение.
Подводя итог: архитектуры U-Net и DiT хорошо работают на практике. U-Net изначально использовалась в нескольких text-to-image моделях, таких как Imagen от Google [17] и Stable Diffusion от Stability AI [11]. В последнее время архитектура DiT показала большой потенциал для генерации text-to-image. В образовательных целях в этой главе мы используем архитектуру U-Net. В Главе 11 архитектура DiT рассматривается подробнее.
Обучение
Diffusion модель обучается посредством diffusion процесса. Diffusion процесс состоит из двух фаз:
- Прямой процесс
- Обратный процесс
Прямой процесс
В прямом процессе, также известном как процесс зашумления, шум постепенно добавляется к изображению через множество шагов (обозначаемых как t или timestep), пока изображение не станет полностью зашумлённым. Значение t, представляющее количество шагов, обычно выбирается случайным образом из диапазона, как правило от 1 до 1000. Прямой процесс не включает никаких ML-моделей или обновления параметров.
Обратный процесс
В обратном процессе, также известном как процесс денойзинга, ML-модель учится обращать прямой процесс. На каждом шаге модель предсказывает шум в зашумлённом изображении. Этот предсказанный шум затем используется для уменьшения шума во входном изображении. Как показано на рисунке 13, этот процесс повторяется до тех пор, пока изображение не станет чётким.
Понимая оба процесса — прямой и обратный — мы можем теперь рассмотреть, как они применяются в процессе diffusion обучения.
Процесс diffusion обучения
При обучении мы вносим шум в исходное изображение, симулируя прямой процесс, затем просим модель предсказать этот шум. Этот процесс включает четыре ключевых шага:
- Добавление шума
- Подготовка conditioning сигналов
- Предсказание шума
- ML objective и расчёт loss
Этот раздел содержит математику для тех, кому это интересно, но детали не влияют на проектирование ML-системы.
. Добавление шума
Первый шаг — симуляция прямого diffusion процесса путём добавления шума к исходному изображению через несколько timestep'ов. На каждом timestep'е мы немного повреждаем изображение, добавляя гауссовский шум. Постепенное добавление шума со временем превращает изображение в чистый шум.
Количество шума, добавляемого на каждом timestep'е, контролируется расписанием шума. Расписание шума определяется набором параметров дисперсии: 1, 2, ..., T, где T — общее количество timestep'ов. Каждый βt(0,1) управляет количеством шума, добавляемого на timestep'е t.
Расписание шума обычно инкрементально увеличивает значения β\betaβ:
Таким образом, на ранних шагах добавляется меньше шума, что сохраняет больше исходного изображения, а на поздних шагах добавляется больше шума, ускоряя diffusion процесс.
С определённым расписанием шума можно выразить зашумлённые данные на timestep'е ttt с помощью формулы добавления шума:
где:
- xtx_txt — зашумлённое изображение на timestep'е ttt,
- xt−1x_{t-1}xt−1 — зашумлённое изображение на timestep'е t−1t-1t−1,
- ϵ\epsilonϵ — гауссовский шум, сэмплированный из стандартного нормального распределения N(0,I)N(0, I)N(0,I),
- βt\beta_tβt — параметр расписания дисперсии на timestep'е t, управляющий количеством добавляемого шума.
Итеративное добавление шума через множество шагов может занимать много времени. Вместо этого можно показать, что зашумлённые данные на timestep'е ttt можно напрямую вывести из исходных данных, x0x_0x0:
где:
- xtx_txt — зашумлённое изображение на timestep'е ttt,
- αt=1−βt\alpha_t=1-\beta_tαt=1−βt и αt′=∏i=1tαi\alpha_t^{\prime}=\prod_{i=1}^t \alpha_iαt′=∏i=1tαi — репараметризации βt\beta_tβt,
- ϵ\epsilonϵ — гауссовский шум, сэмплированный из стандартного нормального распределения N(0,I)N(0, I)N(0,I).
Подводя итог: при добавлении шума мы случайным образом сэмплируем ttt и вычисляем xtx_txt напрямую из x0x_0x0 без необходимости итеративно добавлять шум на каждом timestep'е по следующей формуле:
. Подготовка conditioning сигналов
Для предсказания добавленного шума модель обычно ожидает две дополнительные части информации: подпись к изображению и сэмплированный timestep ttt, указывающий уровень шума. Для подготовки каждого из этих conditioning сигналов к обработке моделью мы используем отдельные encoder'ы (см. рисунок 14).
. Предсказание шума
Основная цель обучения diffusion модели — научиться обращать прямой diffusion процесс, то есть реконструировать исходные данные x0x_0x0 из их зашумлённой версии xtx_txt. Было показано, что напрямую предсказывать x0x_0x0 не так эффективно. Вместо этого обучение модели для предсказания шума ϵ\epsilonϵ, добавленного в ходе прямого процесса, упрощает задачу и улучшает производительность. Поэтому на этом шаге модель предсказывает шум ϵ\epsilonϵ, зная зашумлённый вход xtx_txt и timestep t.3
. ML objective и расчёт loss
ML objective — минимизировать разницу между истинным шумом ϵ\epsilonϵ и предсказанием модели. Используемая функция loss — среднеквадратичная ошибка (MSE) между истинным и предсказанным шумом:
где:
- ttt — timestep, равномерно сэмплированный из {1,2,…,T}\{1,2, \ldots, T\}{1,2,…,T},
- ϵ∼N(0,I)\epsilon \sim N(0, I)ϵ∼N(0,I) — гауссовский шум, используемый в прямом процессе,
- xtx_txt — зашумлённые данные на timestep'е ttt, вычисленные по формуле добавления шума,
- ϵθ(xt,t)\epsilon_\theta\left(x_t, t\right)ϵθ(xt,t) — предсказание модели нейронной сети (U-Net или DiT).
Для улучшения читаемости мы опустили математические детали, такие как вывод трактуемого среднего и упрощение функции loss. Для получения дополнительной информации о diffusion обучении обратитесь к [18].
Сэмплирование
Сэмплирование означает генерацию нового изображения из обученной diffusion модели. В этом разделе мы рассмотрим, как работает сэмплирование в diffusion моделях и как шумы преобразуются в связные изображения под руководством текстового prompt'а.
Процесс сэмплирования начинается с изображения случайных пикселей, обычно взятых из гауссовского распределения. Затем модель постепенно уточняет это изображение шаг за шагом. На каждом шаге diffusion модель предсказывает шум в текущем изображении и использует это предсказание для небольшой корректировки изображения в нужном направлении. Постепенное уточнение продолжается, каждый шаг даёт более чёткое изображение, пока не будет достигнуто чёткое и детализированное изображение.
Описанный выше базовый процесс сэмплирования имеет два недостатка. Во-первых, он часто не может генерировать изображения, точно соответствующие текстовому prompt'у. Во-вторых, он медленный, потому что для генерации каждого сэмпла требуется множество итеративных шагов. Следующие два метода широко применяются на практике для устранения указанных недостатков:
- Classifier-free guidance (CFG): CFG [19] улучшает соответствие между изображениями и текстовыми prompt'ами в diffusion моделях. Во время обучения модель учится генерировать изображения как с текстовым prompt'ом, так и без него. При сэмплировании CFG регулирует баланс между этими двумя режимами. CFG обеспечивает близкое соответствие сгенерированных изображений текстовым prompt'ам, усиливая влияние обусловленного (с текстовым prompt'ом) режима и уменьшая необусловленный (без текстового prompt'а) режим. Это корректировка направляет diffusion процесс для получения более точных результатов. Для подробного изучения CFG обратитесь к [19].
- Сокращение количества шагов diffusion: Алгоритмы сэмплирования, такие как DDIM [20], сокращают количество шагов diffusion со стандартных 1000 до 20. Это значительно ускоряет процесс генерации при сохранении качества изображений. Для подробного изучения DDIM обратитесь к [20].
В большинстве интервью по проектированию ML-систем фокус делается на высокоуровневых концепциях и взаимодействии компонентов, а не на запутанных деталях. Если вас интересует более глубокое изучение diffusion моделей, обратитесь к [21][19][20].
Проблемы text-to-image diffusion моделей
Diffusion модели обычно очень большие. Например, DALLE-2 имеет 3,5 миллиарда параметров [10]. Такая ёмкость необходима, поскольку эти модели должны изучить разнообразные концепции, формы и стили.
Обучение таких больших моделей создаёт несколько трудностей как при обучении, так и при сэмплировании. Наиболее распространённые из них:
- Ресурсоёмкое обучение модели
- Медленная генерация изображений
Ресурсоёмкое обучение модели
Обучение diffusion моделей вычислительно интенсивно и требует значительной вычислительной мощности. Оно также требует существенной памяти GPU из-за размера модели и высокоразмерной природы сгенерированных изображений. Большинство современных GPU могут не иметь достаточно памяти для хранения параметров модели, активаций и градиентов при обучении. Для преодоления этих трудностей широко применяются следующие стратегии:
- Обучение со смешанной точностью: Этот метод использует оба типа чисел с плавающей точкой — 16-битные и 32-битные — для уменьшения использования памяти и повышения вычислительной эффективности. Подробнее — в [22].
- Параллелизм моделей и данных: Эти методы распределяют обучение по нескольким устройствам. Большинство фреймворков распределённого обучения, таких как FSDP [23] и Deepspeed [24], поддерживают различные техники параллелизма.
- Latent diffusion модели: Эти модели работают в пространстве меньшей размерности вместо пиксельного пространства, значительно ускоряя обучение и inference. Глава 11 рассматривает этот подход подробнее.
Медленная генерация изображений
Генерация изображений из текста в diffusion моделях медленна по двум основным причинам. Во-первых, из-за последовательной природы процесса сэмплирования diffusion требуется множество шагов для уточнения изображения. Во-вторых, поскольку diffusion модели имеют миллиарды параметров, на каждом шаге происходят значительные вычисления.
Распространённые стратегии для снижения этой проблемы включают:
- Параллельное сэмплирование: Реализация параллельной обработки при сэмплировании сокращает время, необходимое для генерации изображений [25].
- Model distillation: Дистиллированная модель улучшает скорость генерации благодаря уменьшенному размеру, сохраняя поведение и производительность исходной модели. Для получения дополнительной информации о model distillation в diffusion моделях обратитесь к [26].
- Квантизация модели: Этот метод снижает точность весов модели, что уменьшает использование памяти и ускоряет генерацию.
Оценка
Офлайн-метрики оценки
Последовательный и надёжный бенчмарк является ключом к оценке text-to-image моделей. DrawBench [17] служит этой цели, предоставляя тщательно отобранный набор prompt'ов, проверяющих различные аспекты генерации изображений: композицию объектов, взаимодействие и понимание контекста. Эти prompt'ы — от простых до сложных — помогают оценить, насколько точно модель генерирует изображения по тексту. Благодаря своей полноте мы используем DrawBench для оценки нашей text-to-image модели. Рассмотрим как автоматические метрики, так и оценку людьми для анализа трёх ключевых областей способности модели генерировать изображения:
- Качество изображений
- Разнообразие изображений
- Соответствие изображений тексту
Как обсуждалось в предыдущих главах, Inception score (IS) [27] и Fréchet Inception distance (FID) [28] — две распространённые метрики для оценки качества и разнообразия в системах генерации изображений. В этом разделе мы сосредоточимся в основном на соответствии изображений тексту.
Соответствие изображений тексту
Соответствие изображений тексту означает, насколько точно сгенерированные изображения соответствуют текстовым prompt'ам. Измерение этого соответствия важно, поскольку оно гарантирует, что сгенерированные изображения соответствуют вводу пользователя. Распространённой метрикой для оценки этого является CLIPScore [29], который измеряет степень соответствия. Прежде чем перейти к CLIPScore, кратко рассмотрим CLIP.
CLIP
CLIP [8] — модель, разработанная OpenAI и обученная сопоставлять изображения с их соответствующими описаниями. Она состоит из двух encoder'ов: один для текста, другой для изображений. Text encoder преобразует входной текст в текстовый embedding; image encoder преобразует изображение в image embedding.
При обучении CLIP учится выравнивать embedding'и, сближая связанные текстовые и image embedding'и и отдаляя несвязанные. Это помогает CLIP разработать общее пространство embedding'ов, где изображение и связанный с ним текст будут отображаться в одно и то же пространство.
После обучения похожие текстовые описания располагаются близко друг к другу в пространстве embedding'ов, а изображения отображаются рядом с соответствующими описаниями.4
Понимая модель CLIP, мы теперь можем легко рассмотреть CLIPScore как метрику для оценки соответствия изображений тексту.
CLIPScore
CLIPScore измеряет косинусное сходство между CLIP embedding'ами текстового описания и изображения. Он показывает, насколько тесно изображение соответствует тексту в многомерном пространстве embedding'ов. Более высокая оценка указывает на лучшее соответствие между изображением и описанием.
Оценка людьми
Оценка людьми дополняет автоматические метрики, оценивая как качество изображений, так и соответствие тексту следующим образом:
- Качество изображений: Рейтеры-люди сравнивают сгенерированное изображение с эталонным, оценивая, какое из них более фотореалистично. Качество изображений определяется процентом случаев, когда сгенерированное изображение предпочтено эталонному.
- Соответствие тексту: Рейтерам-людям показывают изображение и его подпись, а затем задают вопрос: «Точно ли подпись описывает данное изображение?» Ответы «да», «отчасти» и «нет» оцениваются как 100, 50 и 0 соответственно. Эти оценки усредняются отдельно для сгенерированных и эталонных изображений для измерения соответствия тексту.
Онлайн-метрики оценки
Онлайн-метрики измеряют, как модель работает в продакшене. Распространённые метрики для оценки нашей text-to-image модели включают:
- Показатель кликабельности (CTR): Процент пользователей, которые кликают на сгенерированные изображения. Высокий CTR указывает на то, что пользователи считают сгенерированные изображения полезными.
- Время на странице: Среднее время, которое пользователи проводят с сервисом. Более длительное время просмотра указывает на большую вовлечённость пользователей.
- Обратная связь от пользователей: Прямая обратная связь от пользователей собирается через форму обратной связи. Положительная обратная связь указывает на удовлетворённость качеством изображений и соответствием тексту.
- Коэффициент конверсии: Процент пользователей, которые совершают желаемое действие (например, покупку, регистрацию) после взаимодействия со сгенерированными изображениями. Высокий коэффициент конверсии указывает на удовлетворённость производительностью модели.
- Задержка: Время, необходимое для генерации изображения по текстовому prompt'у. Меньшая задержка указывает на более высокую производительность, что важно для удовлетворённости пользователей.
- Пропускная способность: Количество генераций изображений, которые модель может обрабатывать в секунду. Высокая пропускная способность обеспечивает обслуживание большего количества пользователей.
- Использование ресурсов: Вычислительные ресурсы (например, CPU, GPU, память), используемые для запуска модели и обслуживания пользователей. Эффективное использование ресурсов является ключом к снижению затрат.
- Средняя стоимость на пользователя в месяц: Генерация изображений с моделями, имеющими миллиарды параметров, обходится дорого. Если пользователи недовольны изображениями, они могут многократно генерировать новые с тем же prompt'ом, но разными seeds, в надежде на лучшие результаты. Такое поведение увеличивает наши расходы. Отслеживая эту метрику, мы можем убедиться, что затраты остаются обоснованными.
Общее проектирование ML-системы
Хотя diffusion модель является ядром системы text-to-image генерации, несколько других конвейеров играют важную роль в обеспечении эффективности, безопасности и качества. В этом разделе мы погружаемся в целостный дизайн системы генерации text-to-image, рассматривая следующие конвейеры:
- Конвейер данных
- Конвейер обучения
- Конвейер оптимизации модели
- Конвейер inference
Конвейер данных
Конвейер данных подготавливает данные для обучения путём удаления неподходящих изображений, стандартизации остальных и их хранения. Он обеспечивает наличие и релевантность подписей, а также использует предобученную модель, такую как T5 от Google [9], для предварительного вычисления и кэширования embedding'ов подписей. Это кэширование снижает вычислительную нагрузку во время обучения.
Помимо подготовки пар текст–изображение из обучающих данных, конвейер также собирает и обрабатывает новые сгенерированные данные: пользовательские prompt'ы, сгенерированные изображения и обратную связь от пользователей. Эти новые данные добавляются в обучающий набор для будущего использования.
Конвейер обучения
Конвейер обучения обучает модель на последних обучающих данных, собранных конвейером данных.
Конвейер обучения обеспечивает адаптацию модели к последним пользовательским prompt'ам и обучение на изображениях более высокого качества.
Конвейер оценки
Конвейер оценки оценивает новые обученные модели с использованием предварительно определённых автоматических метрик, чтобы определить, соответствуют ли они стандартам производительности и качества для развёртывания.
Конвейер оптимизации модели
Конвейер оптимизации модели отвечает за повышение эффективности модели. Существует несколько методов оптимизации моделей:
- Сжатие модели: Использование техник квантизации и прунинга для уменьшения размера модели и времени генерации.
- Model distillation: Дистилляция модели в меньшую для уменьшения размера модели и времени генерации.
- Оптимизированные алгоритмы: Замена сэмплирования более эффективными алгоритмами для более быстрой генерации.
После завершения оптимизации модели оптимизированная модель может заменить существующую в продакшене.
Конвейер inference
Конвейер inference обрабатывает запросы пользователей и генерирует изображения на основе текстовых prompt'ов. Он включает несколько компонентов, каждый из которых играет важную роль в обеспечении качества и безопасности системы. В этом разделе мы рассмотрим ключевые компоненты:
- Сервис автодополнения prompt'а
- Сервис безопасности prompt'а
- Улучшение prompt'а
- Генерация изображений
- Обнаружение вреда
- Сервис суперразрешения
Сервис автодополнения prompt'а
Сервис автодополнения prompt'а использует специализированную модель для предложения возможных следующих слов или фраз в реальном времени по мере ввода пользователем своего prompt'а. Это улучшает пользовательский опыт, предоставляя возможные варианты завершения.
Сервис безопасности prompt'а
Этот сервис использует модель классификации текста для обработки пользовательских prompt'ов и отклонения тех, которые нарушают нашу политику использования, например, запросов на изображения насилия, ненависти или обнажённости.
Этот сервис обеспечивает соответствие системы стандартам безопасности и предотвращает генерацию неподходящих изображений.
Улучшение prompt'а
Компонент улучшения prompt'а уточняет пользовательские prompt'ы для повышения их ясности, связности и детализации.
Этот компонент широко используется в продвинутых системах генерации изображений и видео [30], поскольку эффективно помогает модели создавать лучшие результаты. Он улучшает качество сгенерированных изображений, предоставляя модели более связный и детализированный prompt.
Генерация изображений
Компонент генерации изображений является ядром конвейера inference. Он взаимодействует с T5 text encoder'ом для кодирования улучшенного текстового prompt'а в последовательность token'ов. Эти token'ы передаются в diffusion модель для генерации одного или нескольких изображений для каждого prompt'а.
Обнаружение вреда
Этот компонент обеспечивает безопасность сгенерированных изображений для пользователей. Если изображение всё ещё содержит насилие или обнажённость, несмотря на предыдущие защитные меры, компонент помечает его и блокирует отображение.
Сервис суперразрешения
Сервис суперразрешения увеличивает разрешение сгенерированных изображений. Этот шаг обеспечивает визуальную привлекательность итогового вывода и соответствие требованиям разрешения.
На практике системы text-to-image часто используют как минимум одну модель суперразрешения, поскольку diffusion модели обычно не могут напрямую генерировать изображения высокого разрешения. Вместо этого diffusion модель обучается при меньшем разрешении, а специализированные модели суперразрешения увеличивают разрешение. Например, базовая модель может генерировать изображение 64x64, которое первая модель суперразрешения увеличивает до 256x256, а вторая — до 1024x1024. Метод Google [31] следует этому подходу для достижения желаемого разрешения.
Подводя итог: различные конвейеры работают совместно, чтобы обеспечить надёжность, высокое качество и безопасность системы text-to-image. Конвейер данных создаёт фундамент для постоянного улучшения, конвейеры обучения и оптимизации модели повышают производительность модели. Конвейер inference обеспечивает безопасную, эффективную и высококачественную генерацию изображений. Эти конвейеры создают систему, готовую к реальным задачам.
Дополнительные темы для обсуждения
Если в конце интервью остаётся время, рассмотрите обсуждение следующих дополнительных тем:
- Использование consistency models для более быстрой генерации изображений [26].
- Применение RLHF для улучшения качества [32].
- Расширение text-to-image модели для поддержки приложений inpainting и outpainting [33].
- Персонализация text-to-image модели для конкретной концепции (Глава 10).
- Детали различных техник расписания шума [34].
- Детали DDPM и DDIM, включая их теоретические основы [20][18].
- Поддержка нескольких соотношений сторон и разрешений с использованием техник, таких как Patch n' Pack [35].
- Детали разработки модели повторной генерации подписей [36][37][13].
- Улучшение баланса разнообразие–точность с помощью guidance [19].
- Более продвинутый контроль над сгенерированными изображениями с использованием таких техник, как ControlNet [38].
- Управление стилем сгенерированных изображений [39].
Итоги
Справочные материалы
[1] OpenAI's DALL-E 3. https://openai.com/index/dall-e-3/. [2] Imagen 3. https://arxiv.org/abs/2408.07009. [3] Adobe's Firefly. https://www.adobe.com/products/firefly.html. [4] Introducing ChatGPT. https://openai.com/index/chatgpt/. [5] Zero-Shot Text-to-Image Generation. https://arxiv.org/abs/2102.12092. [6] Muse: Text-To-Image Generation via Masked Generative Transformers. https://arxiv.org/abs/2301.00704. [7] Generative Modeling by Estimating Gradients of the Data Distribution. https://arxiv.org/abs/1907.05600. [8] Learning Transferable Visual Models From Natural Language Supervision. https://arxiv.org/abs/2103.00020. [9] Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer. https://arxiv.org/abs/1910.10683. [10] Hierarchical Text-Conditional Image Generation with CLIP Latents. https://arxiv.org/abs/2204.06125. [11] High-Resolution Image Synthesis with Latent Diffusion Models. https://arxiv.org/abs/2112.10752. [12] On the De-duplication of LAION-2B. https://arxiv.org/abs/2303.12733. [13] xGen-MM (BLIP-3): A Family of Open Large Multimodal Models. https://www.arxiv.org/abs/2408.08872. [14] U-Net: Convolutional Networks for Biomedical Image Segmentation. https://arxiv.org/abs/1505.04597. [15] Scalable Diffusion Models with Transformers. https://arxiv.org/abs/2212.09748. [16] An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. https://arxiv.org/abs/2010.11929. [17] Photorealistic Text-to-Image Diffusion Models with Deep Language Understanding. https://arxiv.org/abs/2205.11487. [18] Denoising Diffusion Probabilistic Models. https://arxiv.org/abs/2006.11239. [19] Classifier-Free Diffusion Guidance. https://arxiv.org/abs/2207.12598. [20] Denoising Diffusion Implicit Models. https://arxiv.org/abs/2010.02502. [21] Introduction to Diffusion Models. https://lilianweng.github.io/posts/2021-07-11-diffusion-models/. [22] Mixed Precision Training. https://arxiv.org/abs/1710.03740. [23] FSDP tutorial. https://pytorch.org/tutorials/intermediate/FSDP_tutorial.html. [24] DeepSpeed. https://github.com/microsoft/DeepSpeed. [25] Parallel Sampling of Diffusion Models. https://arxiv.org/abs/2305.16317. [26] Consistency Models. https://arxiv.org/abs/2303.01469. [27] Inception score. https://en.wikipedia.org/wiki/Inception_score. [28] FID calculation. https://en.wikipedia.org/wiki/Fr%C3%A9chet_inception_distance. [29] CLIPScore: A Reference-free Evaluation Metric for Image Captioning. https://arxiv.org/abs/2104.08718. [30] Sora overview. https://openai.com/index/video-generation-models-as-world-simulators/. [31] Imagen Video: High Definition Video Generation with Diffusion Models. https://arxiv.org/abs/2210.02303. [32] Finetune Stable Diffusion Models with DDPO via TRL. https://huggingface.co/blog/trl-ddpo. [33] Kandinsky: an Improved Text-to-Image Synthesis with Image Prior and Latent Diffusion. https://arxiv.org/abs/2310.03502. [34] On the Importance of Noise Scheduling for Diffusion Models. https://arxiv.org/abs/2301.10972. [35] Patch n' Pack: NaViT, a Vision Transformer for any Aspect Ratio and Resolution. https://arxiv.org/abs/2307.06304. [36] InternVL: Scaling up Vision Foundation Models and Aligning for Generic Visual-Linguistic Tasks. https://arxiv.org/abs/2312.14238. [37] BLIP-2: Bootstrapping Language-Image Pre-training with Frozen Image Encoders and Large Language Models. https://arxiv.org/abs/2301.12597. [38] Adding Conditional Control to Text-to-Image Diffusion Models. https://arxiv.org/abs/2302.05543. [39] StyleDrop: Text-to-image generation in any style. https://research.google/blog/styledrop-text-to-image-generation-in-any-style/.
Сноски
- На практике мы предоставляем diffusion модели дополнительные conditioning входные данные, такие как timestep. Это будет обсуждаться подробнее в разделе обучения. ↩
- Для простоты мы опустили входные данные timestep для diffusion модели. Это будет рассмотрено подробнее в разделе обучения. ↩
- Для простоты в качестве входных данных включены только зашумлённое изображение xt и timestep t. Как упоминалось ранее, conditioning сигналы также могут включаться как входные данные. ↩
- Для простоты пространство embedding'ов визуализируется в 2D. В действительности это d-мерное пространство, где d представляет размер embedding'а. ↩
Персонализированная генерация портретных фотографий
Введение
Персонализированные text-to-image (T2I) модели — одно из новых применений генеративного ИИ. Представьте, что вы просите T2I модель сгенерировать изображение вашего друга Джона с prompt'ом «Джон сидит на стуле и читает книгу». Хотя модель, вероятно, создаст изображение сидящего читающего человека, она, скорее всего, не изобразит именно «Джона». Для этого нам нужно персонализировать T2I модель, обучив её понимать интересующий субъект (т.е. Джона).
В этой главе мы рассмотрим, как разработать персонализированную T2I модель, способную генерировать профессиональные портретные фотографии конкретных людей.
Уточнение требований
Ниже приведён типичный диалог между кандидатом и интервьюером:
Кандидат: Предназначены ли сгенерированные портреты прежде всего для деловых профилей, таких как LinkedIn? Интервьюер: Верно.
Кандидат: Предполагаю, что пользователи будут предоставлять несколько своих изображений в разных вариантах — позах, ракурсах — с видимым лицом. Это верно? Сколько изображений мы попросим загрузить? Интервьюер: Да, всё верно. Предположим, от 10 до 20 изображений.
Кандидат: Что если некоторые изображения не подходят — например, слишком тёмные или лицо не видно? Интервьюер: Нам нужно это обнаруживать и уведомлять пользователя о необходимости предоставить лучшие изображения.
Кандидат: Должны ли пользователи иметь возможность указывать атрибуты, такие как причёска, на сгенерированных изображениях? Интервьюер: Для простоты предположим, что управление атрибутами не требуется.
Кандидат: Какое разрешение требуется для портретов? Интервьюер: Система должна поддерживать выходные данные 1024x1024.
Кандидат: Могу ли я предположить, что мы можем начать с предобученной универсальной T2I модели? Интервьюер: Да.
Кандидат: Должны ли пользователи иметь возможность предоставлять текстовые prompt'ы для управления сгенерированными портретами? Интервьюер: Мы предпочитаем сохранить простоту, поэтому предположим, что пользователи не будут предоставлять текстовые prompt'ы.
Кандидат: Сколько портретных изображений должна генерировать система? Интервьюер: 50 изображений.
Кандидат: Какова ожидаемая задержка? Интервьюер: Пользователь предоставляет изображения, и мы уведомляем его по электронной почте, когда изображения готовы. Весь процесс должен занимать менее часа.
Формулировка задачи как ML-задачи
В этом разделе мы используем генерацию портретов в качестве примера для рассмотрения важного аспекта генерации изображений: персонализации. Этот процесс включает адаптацию предобученной T2I модели для изучения нового субъекта — в данном случае лица пользователя.
Определение входа и выхода системы
Вход включает несколько изображений лица пользователя, снятых под разными углами и в разных позах. Выход — профессиональные портреты человека. Эти портреты высокого качества, разнообразны и сохраняют идентичность человека.
Выбор подходящего ML-подхода
Diffusion модели исключительно хорошо генерируют очень детализированные и реалистичные изображения. Распространённая практика для персонализации — начать с T2I модели, предобученной на широком диапазоне изображений, как базовой модели.
Существует два основных подхода к персонализации предобученной T2I модели: с настройкой (tuning-based) и без настройки (tuning-free).
Методы с настройкой проводят fine-tuning T2I модели на наборе эталонных изображений для каждой идентичности. Этот подход интегрирует новую идентичность в модель, позволяя ей генерировать разнообразные изображения, сохраняя идентичность.
С другой стороны, методы без настройки позволяют обойтись без fine-tuning T2I модели для каждой новой идентичности. Вместо этого они один раз проводят fine-tuning предобученной T2I модели вместе с визуальным encoder'ом. После этого обучения визуальный encoder извлекает признаки из нового эталонного изображения и внедряет их в T2I модель. Это позволяет модели генерировать персонализированные изображения без корректировки внутренних весов для каждой идентичности.
Методы без настройки, такие как Imagine Yourself от Meta [1], проще, поскольку требуют обучения только один раз с меньшим количеством параметров, чем обучение всей T2I модели. Необходимо развернуть одну модель, и требуется только одно эталонное изображение, что позволяет одной и той же предобученной модели генерировать персонализированные изображения для нескольких идентичностей. Однако эти методы часто опираются на одно эталонное изображение для захвата черт лица, что может не захватить все детали под разными углами или выражениями. Кроме того, они требуют специальных корректировок, адаптированных к субъекту. Например, [1] использует специфические encoder'ы и предлагает технику для генерации синтетических парных данных.
С другой стороны, методы с настройкой, как правило, захватывают более детальные черты субъекта. Они также более универсальны, обрабатывая более широкий спектр субъектов помимо человеческих лиц. Учитывая эти преимущества, в остальной части главы мы сосредоточимся на методах с настройкой. Если вас интересует дополнительная информация о методах без настройки, обратитесь к [2][3][1].
Для достижения персонализации можно использовать несколько методов с настройкой, каждый из которых предлагает уникальные преимущества и недостатки. Три наиболее распространённых:
- Textual inversion
- DreamBooth
- Low-rank adaptation (LoRA)
Textual inversion
Textual inversion [4] персонализирует T2I модель, вводя новый специальный token, представляющий субъект, и обучая его embedding. В ходе fine-tuning модель обновляет embedding специального token'а, тогда как diffusion модель, text encoder и embedding'и других token'ов остаются неизменными.
После fine-tuning модель генерирует изображения нового субъекта при появлении в prompt'е специального token'а.
Рассмотрим плюсы и минусы textual inversion.
Плюсы:
- Эффективность: Обучение textual inversion включает обучение только нового token embedding'а, что делает процесс лёгким и эффективным.
- Сохранение исходных возможностей модели: Возможности T2I модели сохраняются, поскольку параметры diffusion модели остаются неизменными.
- Минимальные требования к хранилищу: Требуется минимальное хранилище, поскольку для каждой персонализированной модели необходимо хранить только embedding специального token'а.
Минусы:
- Трудность изучения деталей субъекта: Textual inversion часто испытывает трудности с точным изучением новых деталей субъекта из-за ограниченной способности кодировать эти детали. Это ограничение возникает потому, что новый субъект представлен одним token embedding'ом.
Подводя итог: хотя textual inversion является эффективным методом персонализации T2I модели, он часто испытывает трудности с захватом и сохранением всех деталей нового субъекта.
DreamBooth
DreamBooth [5] — популярный метод персонализации, представленный Google в 2023 году. Он проводит fine-tuning предобученной diffusion модели с использованием изображений интересующего субъекта. В отличие от textual inversion, DreamBooth обновляет все параметры diffusion модели в ходе fine-tuning. Это позволяет модели более эффективно захватывать детали нового субъекта.
DreamBooth в основном опирается на две техники для успешного fine-tuning:
- Идентификатор редкого token'а
- Функция loss сохранения предшествующего распределения класса
Идентификатор редкого token'а
Большинство методов персонализации выбирают идентификатор для представления интересующего субъекта. В то время как textual inversion создаёт новый token для этой цели, DreamBooth выбирает идентификатор из существующего словаря token'ов. Авторы обнаружили, что простой выбор любого существующего token'а на практике работает плохо. Давайте рассмотрим, почему простые подходы, такие как выбор обычного или случайного идентификатора, не работают, и почему предпочтителен идентификатор редкого token'а.
Проблемы с обычными или случайными идентификаторами token'ов
Простой способ представить интересующий субъект — выбрать обычное английское слово, такое как «unique» или «special». Этот подход проблематичен, поскольку token обычно имеет устоявшееся значение. Например, если мы выберем «special» для обозначения интересующего субъекта с prompt'ом вроде «a special person sitting», модель может испытывать трудности, потому что «special» уже имеет широкое, устоявшееся значение в контексте языка. Модели приходится отделять этот token от его первоначального значения и обучать для него новое значение, ссылающееся на нового субъекта.
Другой способ представить субъект — случайно комбинировать символы. Этот подход также создаёт проблемы, поскольку токенизатор может обрабатывать каждый символ отдельно, создавая сильные априорные ассоциации для этих символов. Например, если мы выберем «xxy5syt00» для представления субъекта, токенизатор может разбить его на отдельные символы, каждый из которых может иметь существующие ассоциации в модели. Такая фрагментация может привести к тому, что модель будет генерировать выходные данные, на которые влияют значения или паттерны, связанные с этими отдельными символами или подединицами, вместо того чтобы воспринимать идентификатор как уникальную и цельную сущность.
Как работает идентификатор редкого token'а
DreamBooth решает эти проблемы, выбирая редкие token'ы — те, которые редко встречаются в обучающих данных. Представление интересующего субъекта с помощью этих token'ов обеспечивает баланс: они достаточно отличительны, чтобы избежать сильных априорных ассоциаций, но при этом достаточно цельны, чтобы токенизаторы воспринимали их как единицу.
Вот пошаговый процесс формирования идентификатора:
- Идентификация нескольких редких token'ов в словаре: Словарь модели содержит большой набор token'ов, каждый с уникальным ID. «Редкий token» — тот, который редко встречается в обучающих данных. Такие редкие token'ы определяются путём оценки распределения частот token'ов.
- Генерация последовательности редких token'ов: После идентификации редких token'ов мы генерируем последовательность с использованием некоторых из них.
- Формирование идентификатора: Токенизатор преобразует последовательность ID token'ов обратно в соответствующую текстовую форму. Это формирует идентификатор для представления субъекта. «XyZ», «SKS» и «[V]» — возможные примеры идентификатора.
Функция loss сохранения предшествующего распределения класса
DreamBooth проводит fine-tuning всех слоёв diffusion модели. Хотя это улучшает качество сгенерированных изображений, это может привести к переобучению, при котором модель теряет разнообразие в выходных данных. Например, обучение модели на изображениях конкретной собаки может привести к тому, что ей будет сложно генерировать изображения других собак.
Для решения этой проблемы DreamBooth использует функцию loss сохранения предшествующего распределения класса для поддержания общих характеристик класса. Это предотвращает переобучение модели на конкретных примерах и потерю способности генерировать разнообразные изображения, принадлежащие к более широкому классу. Функцию loss DreamBooth мы рассмотрим в разделе обучения.
DreamBooth имеет ряд плюсов и минусов.
Плюсы:
- Эффективное изучение деталей субъекта: Обновление большего числа параметров позволяет модели более точно изучать детали субъекта.
- Требуется меньше изображений: Поскольку обновляется вся diffusion модель, для изучения субъекта требуется меньший набор изображений.
Минусы:
- Высокие требования к хранилищу: После fine-tuning для каждого субъекта вся diffusion модель должна храниться для будущего использования. Это может потребовать нескольких гигабайт на субъект, что дорогостояще и не масштабируется.
- Ресурсоёмкость: Обновление всей diffusion модели требует больше памяти GPU при обучении.
Подводя итог: DreamBooth эффективно изучает детали субъекта, но дорог в обучении и хранении. Напротив, textual inversion эффективен и компактен, но менее результативен. Далее мы рассмотрим LoRA, который предлагает сбалансированный подход.
LoRA
LoRA, представленный Microsoft [6], — это мощный метод для эффективного fine-tuning очень больших моделей. Этот метод был изначально разработан для адаптации больших языковых моделей (LLM) к конкретным задачам, но позднее был принят для других задач, включая персонализацию T2I.
Ключевая мотивация LoRA заключается в том, что fine-tuning всех параметров больших предобученных моделей, таких как GPT-3 [7], занимает много времени и обходится дорого. Вместо этого LoRA адаптирует большую модель к новой задаче, вводя небольшой набор параметров и обновляя только их, тем самым значительно снижая вычислительные затраты.
Ввиду важности темы, давайте подробно рассмотрим математические основы LoRA.
Математика LoRA
В типичном слое нейронной сети матрица весов W∈Rdout ×din W \in \mathbb{R}^{d_{\text {out }} \times d_{\text {in }}}W∈Rdout ×din преобразует входной вектор x∈Rdinx \in \mathbb{R}^{d_{i n}}x∈Rdin в выходной вектор y∈Rdout y \in \mathbb{R}^{d_{\text {out }}}y∈Rdout .
Цель fine-tuning — скорректировать параметры весов WWW для улучшения производительности на конкретной задаче. Вместо прямого изменения WWW, LoRA модифицирует веса, вводя дополнительный низкоранговый компонент ΔW\Delta WΔW, который можно выразить как произведение двух обучаемых низкоранговых матриц:
где:
- A∈Rdout ×drA \in \mathbb{R}^{d_{\text {out }} \times d_r}A∈Rdout ×dr,
- B∈Rdr×dinB \in \mathbb{R}^{d_r \times d_{\mathrm{in}}}B∈Rdr×din,
- din d_{\text {in }}din и dout d_{\text {out }}dout представляют входное и выходное измерения,
- rrr — малое целое число, представляющее ранг, как правило значительно меньше din d_{\text {in }}din и dout d_{\text {out }}dout .
Обучение новых параметров, введённых LoRA, более эффективно, чем fine-tuning полной матрицы. В частности, исходная матрица WWW имеет dout ×din d_{\text {out }} \times d_{\text {in }}dout ×din параметров, тогда как низкоранговая аппроксимация вводит лишь r×(din +dout )r \times\left(d_{\text {in }}+d_{\text {out }}\right)r×(din +dout ) параметров. При малых значениях rrr это может приводить к значительной экономии как в хранилище, так и в вычислениях.
LoRA для персонализации T2I
Для применения LoRA к нашей предобученной T2I модели мы внедряем обучаемые параметры в diffusion модель и обновляем только эти параметры в ходе fine-tuning для изучения новой идентичности. Этот метод требует обучения лишь небольшой доли параметров модели, что значительно быстрее и более экономично по памяти.
Плюсы:
- Сохранение исходных возможностей модели: LoRA сохраняет возможности T2I модели, замораживая исходные параметры модели.
- Снижение потребностей в памяти и вычислениях: LoRA обновляет лишь небольшую долю параметров модели, что делает его более эффективным, чем DreamBooth.
- Минимальные требования к хранилищу: Поскольку исходная модель остаётся неизменной, хранятся только слои LoRA. Это, как правило, составляет всего несколько мегабайт на персонализированную модель, что экономично и масштабируемо.
Минусы:
- Менее эффективное обучение: LoRA менее эффективен, чем DreamBooth, поскольку проводит fine-tuning лишь небольшого числа параметров, ограничивая способность изучать новый субъект.
- Незначительное увеличение времени inference: LoRA незначительно увеличивает время inference из-за дополнительных параметров и вычислений. Однако это зачастую пренебрежимо мало по сравнению с общими преимуществами снижения требований к хранилищу и более быстрой адаптации.
Таблица 1 содержит сравнение трёх методов персонализации с настройкой.
| Textual Inversion | LoRA | DreamBooth | |
|---|---|---|---|
| Learning effectiveness | Low | Moderate | High |
| Required storage | Low | Moderate | High |
| Required training resources | Low | Moderate | High |
| Maintaining the original model's capabilities | Yes | Yes | No |
Таблица 1: Сравнение популярных методов персонализации с настройкой
Какой метод больше подходит для генерации портретов?
Пригодность этих методов зависит от варианта использования и требований системы. Для генерации портретов мы выбираем DreamBooth по трём основным причинам:
- Лучшее сохранение идентичности: DreamBooth наиболее эффективен при сохранении деталей субъекта, обеспечивая лучшее сохранение идентичности.
- Приемлемое время обучения: Согласно [5], fine-tuning diffusion модели с использованием DreamBooth занимает около 15 минут. Это время обучения приемлемо, поскольку мы должны предоставить пользователю сгенерированные изображения в течение часа.
- Не требуется хранение: После генерации портретов нам не нужно сохранять персонализированные модели, поэтому проблемы с хранением при подходе DreamBooth не актуальны.
Подготовка данных
Необходимое количество изображений варьируется в зависимости от метода. Методы без настройки, как правило, требуют только одного изображения, тогда как методы с настройкой, такие как DreamBooth, нуждаются примерно в 10–20 изображениях.
Поскольку мы используем DreamBooth, пользователей просят загрузить 10–20 изображений. Эти изображения могут иметь разное разрешение и соотношение сторон. Для их подготовки к обучению мы выполняем следующие шаги:
- Изменение размера изображений
- Аугментация изображений
- Добавление данных общих лиц
Изменение размера изображений
Diffusion модели обычно требуют фиксированных входных размеров, но загружаемые пользователями изображения часто различаются по размерам. Мы изменяем размер изображений до единых размеров, подходящих для diffusion модели.
Аугментация изображений
T2I модели требуют большого количества изображений для изучения таких концепций, как объекты, идентичности и сцены. Однако для персонализации зачастую не хватает большого набора данных. Мы используем техники аугментации изображений, такие как зеркальное отражение, небольшие повороты и масштабирование, для искусственного расширения набора данных. Этот шаг важен, когда для обучения доступно лишь небольшое количество изображений.
Добавление данных общих лиц
Обучение только на предоставленных изображениях может привести к переобучению модели к конкретной идентичности и забыванию ранее усвоенных знаний. Чтобы предотвратить это, мы объединяем загружаемые пользователями изображения с более крупным, универсальным набором данных лиц. Мы генерируем эти изображения с помощью предобученной diffusion модели с prompt'ами вроде «изображение человека».
Разработка модели
Архитектура
Метод DreamBooth проводит fine-tuning предобученной diffusion модели. Мы используем модель с архитектурой U-Net, предобученную для вывода изображений 1024x1024, аналогично тому, что мы рассматривали в Главе 9. Архитектура остаётся неизменной: серия блоков downsampling, за которыми следует серия блоков upsampling.
Обучение
Для fine-tuning предобученной diffusion модели мы следуем тому же процессу, что и при обучении diffusion модели с нуля:
- Добавление шума: К изображению добавляется шум на основе случайно выбранного timestep'а.
- Подготовка conditioning сигналов: Отдельные encoder'ы подготавливают подпись к изображению и timestep как conditioning сигналы для предсказания шума моделью.
- Предсказание шума: Модель предсказывает шум, который нужно удалить из зашумлённого изображения, используя conditioning сигналы.
Обучающие данные
Обучающие данные состоят из изображений, загружаемых пользователями, и общих изображений лиц, добавленных на этапе подготовки данных. Загружаемые пользователем изображения помечены как «Изображение человека [V]», тогда как общие изображения лиц помечены как «Изображение человека».
ML objective и функция loss
Ключевая трудность персонализации — обеспечить способность модели генерировать как общие категории (например, человеческие лица), так и конкретные экземпляры в этих категориях (например, уникального человека). Для решения этой задачи мы используем две функции loss:
- Функция loss реконструкции
- Функция loss сохранения предшествующего распределения класса
Функция loss реконструкции
Эта функция loss измеряет различия между реконструированным изображением и фактическими изображениями конкретного субъекта. Она помогает модели сохранять идентичность субъекта.
Функция loss сохранения предшествующего распределения класса
Эта функция loss измеряет разницу между сгенерированными изображениями и фактическими изображениями общих лиц. Она гарантирует, что модель сохраняет характеристики класса людей и не переобучается к конкретным идентичностям.
Общая функция loss
Общая функция loss — это взвешенная комбинация функции loss реконструкции и функции loss сохранения предшествующего распределения класса. Формула для общей функции loss может быть выражена как:
α\alphaα и β\betaβ — гиперпараметры, управляющие балансом между сохранением конкретной идентичности и сохранением характеристик человека. ML objective — минимизировать общую функцию loss, что позволяет модели генерировать изображения уникальных идентичностей, сохраняя способность создавать общие изображения человеческих лиц.
Сэмплирование
Процесс сэмплирования для генерации портретов аналогичен генерации T2I, обсуждавшейся в Главе 9, поскольку оба используют diffusion модели. Однако ключевое различие заключается в способе предоставления текстовых prompt'ов. В Главе 9 пользователь предоставлял текстовый prompt. Например, пользователь мог ввести «кошка сидит на стуле», и diffusion модель генерировала изображение, отражающее этот текст.
При генерации портретов пользователи не предоставляют prompt'ы. Вместо этого мы создаём набор вручную составленных prompt'ов. Эти prompt'ы представляют различные профессиональные условия и включают идентификатор, использованный при обучении, чтобы гарантировать, что сгенерированные изображения отражают идентичность пользователя. Некоторые примеры:
- «Профессиональный портрет [V] с улыбкой на фоне простого белого фона.»
- «Крупный план [V] с нейтральным выражением лица, одетого в официальную одежду.»
- «Портрет [V] с мягким освещением, смотрящего немного влево.»
- «Снимок в профиль [V] на фоне размытого уличного пространства.»
- «Профессиональный портрет [V] в деловом костюме с уверенным выражением лица.»
После создания prompt'ов мы сэмплируем одно изображение на каждый prompt из diffusion модели. Мы следуем стандартным шагам сэмплирования diffusion процесса:
- Генерация начального случайного шума.
- Итеративный денойзинг входных данных через несколько шагов с использованием Classifier-free guidance (CFG) [8] на каждом шаге для уменьшения шума и уточнения деталей изображения. Для ознакомления с CFG обратитесь к [8] или Главе 9.
Оценка
Офлайн-метрики оценки
Важно оценивать персонализированные diffusion модели, чтобы убедиться, что наша модель способна сохранять идентичность пользователя на сгенерированных изображениях. Мы оцениваем производительность персонализированной diffusion модели, фокусируясь на трёх основных аспектах:
- Соответствие тексту
- Качество изображений
- Соответствие изображений
Соответствие тексту
Соответствие тексту означает, насколько близко сгенерированные изображения соответствуют текстовым prompt'ам. Распространённой метрикой для его измерения является CLIPScore [9], который гарантирует, что изображения не только высокого качества, но и релевантны входному тексту.
Качество изображений
Как мы рассматривали в предыдущих главах, для измерения качества сгенерированных изображений мы используем распространённые метрики, такие как FID [10] и Inception score [11].
Соответствие изображений
Соответствие изображений, особенно актуальное для персонализированных text-to-image моделей, означает оценку визуального сходства между сгенерированными изображениями и интересующим субъектом. Например, если модели поручено сгенерировать изображение конкретного рюкзака, соответствие изображений измеряет, насколько близко сгенерированное изображение напоминает этот рюкзак.
Широко используемые метрики для измерения визуального сходства между сгенерированным и исходным субъектом:
- CLIP score
- DINO score
- Оценка схожести лиц
CLIP score
Модель CLIP [12] использует два encoder'а — один для изображений и один для текста. Эти encoder'ы обучены гарантировать, что image embedding'и и text embedding'и релевантной пары изображение–текст находятся близко в пространстве embedding'ов.
Для измерения соответствия изображений с помощью CLIP мы отбрасываем text encoder и используем image encoder для генерации embedding'ов как сгенерированных, так и реальных изображений. Затем мы вычисляем косинусное сходство между этими embedding'ами для оценки их схожести. Более высокие оценки указывают на то, что персонализированная diffusion модель создаёт изображения, более визуально похожие на реальные.
DINO score
DINO [13] — метод самообучения без учителя, разработанный Meta. DINO обучается визуальным представлениям изображений без необходимости в размеченных данных. В частности, он использует метод, называемый контрастным обучением [14], при котором модель учится различать схожие и несхожие изображения, организуя их в пространстве embedding'ов — похожие изображения располагаются ближе, несхожие — дальше.
DINO и его более поздние вариации, такие как DINOv2 [15], особенно хороши при захвате сходства между изображениями, поскольку обучены распознавать тонкие различия. Это делает DINO особенно эффективным для измерения соответствия изображений. Сравнивая embedding'и сгенерированного и реального изображений, DINO может оценить, насколько хорошо сгенерированное изображение соответствует реальному.
DINO против CLIP
DINO предпочтителен для сравнения изображений, поскольку обучен захватывать детальные визуальные признаки. Например, два изображения — одно с жёлтой курткой и другое с красной — могут иметь низкую оценку DINO из-за разницы цвета. С другой стороны, CLIP лучше подходит для сравнения изображений с текстом, поскольку обучен сопоставлять описания с визуальным контентом. Те же изображения с разными цветами курток могут иметь высокую оценку CLIP, если оба отражают человека в куртке.
Оценка схожести лиц
Хотя CLIP и DINO измеряют визуальное сходство между сгенерированными и реальными изображениями, они не предназначены для оценки сохранения идентичности. Например, два изображения лиц могут показывать разных людей, но CLIP и DINO всё равно могут давать высокую оценку сходства. Для решения этой проблемы мы используем модель распознавания лиц для сравнения сгенерированных и реальных изображений. Эти модели специализируются на идентификации и измерении схожести лиц, что является ключевым требованием в системе персонализированной генерации портретов.
Сочетание оценок DINO, CLIP и схожести лиц обеспечивает комплексную оценку соответствия изображений в персонализированной diffusion модели.
Онлайн-метрики оценки
Онлайн-оценка критически важна при генерации портретов. Она непосредственно измеряет удовлетворённость пользователей, что жизненно важно, когда пользователи платят за сервис, поскольку более высокая удовлетворённость часто приводит к увеличению доходов. Мы фокусируемся на двух основных метриках:
- Обратная связь от пользователей: Эта метрика непосредственно отражает удовлетворённость пользователей. После получения сгенерированных портретов пользователи оценивают свою удовлетворённость по шкале от 1 до 5. Высокие оценки указывают на то, что портреты соответствуют или превосходят ожидания, тогда как низкие оценки подчёркивают необходимость улучшений.
- Коэффициент конверсии в платный сервис: Эта метрика измеряет процент пользователей, переходящих от проявления интереса к статусу новых платящих клиентов. Она рассчитывается как отношение числа новых платящих клиентов к общему числу пользователей, взаимодействовавших с сервисом (посещение сайта, регистрация для пробного использования или запросы) за определённый период.
Общее проектирование ML-системы
Генерация профессиональных портретов требует большего, чем просто diffusion модель. В этом разделе мы рассматриваем три ключевых конвейера:
- Конвейер данных
- Конвейер обучения
- Конвейер inference
Конвейер данных
Этот конвейер выполняет две задачи:
- Подготовка изображений с интересующим субъектом
- Подготовка изображений общих человеческих лиц
Подготовка изображений с интересующим субъектом
Этот процесс оценивает загружаемые пользователями изображения, чтобы убедиться в соответствии предопределённым стандартам, и подготавливает их для обучения.
В частности, проверяется, что изображения разнообразны и содержат только один объект интереса: лицо пользователя. Для этого мы используем различные эвристики и ML-модели для анализа изображений, проверяя такие факторы, как чёткость, разнообразие ракурсов, выражений лица и наличие лица пользователя. Если какие-либо изображения не соответствуют этим критериям, они отклоняются, и пользователя просят загрузить больше. Это гарантирует, что для fine-tuning diffusion модели используются только высококачественные изображения.
Подготовка изображений общих человеческих лиц
Этот шаг включает подготовку изображений с общими лицами для предотвращения переобучения модели к интересующему субъекту. Мы используем предобученную T2I модель для генерации этих изображений с prompt'ами вроде «человек, сидящий на стуле».
Конвейер обучения
Этот конвейер отвечает за персонализацию предобученной diffusion модели.
Конвейер inference
Конвейер inference отвечает за генерацию портретов пользователя с использованием персонализированной T2I модели. Три основных компонента конвейера inference:
- Генератор изображений
- Сервис оценки качества
- Сервис загрузки
Генератор изображений
Генератор изображений использует персонализированную T2I модель и вручную составленные текстовые prompt'ы для генерации одного изображения на каждый prompt.
Сервис оценки качества
Этот сервис гарантирует, что сгенерированные изображения соответствуют стандартам сохранения идентичности. Он использует предобученную модель распознавания лиц для сравнения сгенерированного портрета с реальными изображениями пользователя. Если сгенерированное изображение не сохраняет идентичность пользователя, сервис отклоняет его и запрашивает у генератора изображений новое изображение с тем же текстовым prompt'ом, но с другим начальным шумом.
Сервис загрузки
Сервис загрузки управляет доставкой сгенерированных изображений пользователю. Он загружает изображения в облачное хранилище, чтобы пользователи могли скачать свои портреты.
Дополнительные темы для обсуждения
Если в конце интервью есть дополнительное время, вот несколько тем для обсуждения:
- Предотвращение катастрофического забывания при fine-tuning [16].
- Детали выбора редкого token'а для нового субъекта и его важность [5].
- Детали функции loss сохранения предшествующего распределения класса [5].
- Решение проблемы сниженного разнообразия выходных данных после fine-tuning [5].
- Поддержка нескольких размеров и соотношений сторон сгенерированных изображений [17].
- Детали методов без настройки, таких как Imagine Yourself от Meta [1].
- Снижение рисков и этических проблем, связанных с генерацией и обнаружением дипфейков [18].
- ML-техники для безопасной обработки персональных данных (PII) с обеспечением конфиденциальности данных [19][20].
Итоги
Справочные материалы
[1] Imagine yourself: Tuning-Free Personalized Image Generation. https://ai.meta.com/research/publications/imagine-yourself-tuning-free-personalized-image-generation/. [2] MoA: Mixture-of-Attention for Subject-Context Disentanglement in Personalized Image Generation. https://arxiv.org/abs/2404.11565. [3] InstantID: Zero-shot Identity-Preserving Generation in Seconds. https://arxiv.org/abs/2401.07519. [4] An Image is Worth One Word: Personalizing Text-to-Image Generation using Textual Inversion. https://textual-inversion.github.io/. [5] DreamBooth: Fine Tuning Text-to-Image Diffusion Models for Subject-Driven Generation. https://arxiv.org/abs/2208.12242. [6] LoRA: Low-Rank Adaptation of Large Language Models. https://arxiv.org/abs/2106.09685. [7] Language Models are Few-Shot Learners. https://arxiv.org/abs/2005.14165. [8] Classifier-Free Diffusion Guidance. https://arxiv.org/abs/2207.12598. [9] CLIPScore: A Reference-free Evaluation Metric for Image Captioning. https://arxiv.org/abs/2104.08718. [10] FID calculation. https://en.wikipedia.org/wiki/Fr%C3%A9chet_inception_distance. [11] Inception score. https://en.wikipedia.org/wiki/Inception_score. [12] Learning Transferable Visual Models From Natural Language Supervision. https://arxiv.org/abs/2103.00020. [13] Emerging Properties in Self-Supervised Vision Transformers. https://arxiv.org/abs/2104.14294. [14] Contrastive Representation Learning. https://lilianweng.github.io/posts/2021-05-31-contrastive/. [15] DINOv2: Learning Robust Visual Features without Supervision. https://arxiv.org/abs/2304.07193. [16] An Empirical Study of Catastrophic Forgetting in Large Language Models During Continual Fine-tuning. https://arxiv.org/abs/2308.08747. [17] SDXL: Improving Latent Diffusion Models for High-Resolution Image Synthesis. https://arxiv.org/abs/2307.01952. [18] Deepfakes, Misinformation, and Disinformation in the Era of Frontier AI, Generative AI, and Large AI Models. https://arxiv.org/abs/2311.17394. [19] Privacy-Preserving Personal Identifiable Information (PII) Label Detection Using Machine Learning. https://ieeexplore.ieee.org/document/10307924. [20] Does fine-tuning GPT-3 with the OpenAI API leak personally-identifiable information? https://arxiv.org/abs/2307.16382.
Генерация видео по тексту
Введение
Генерация видео по тексту -- это ключевое приложение генеративного ИИ, позволяющее создавать видео из текстовых описаний. В этой главе рассматриваются основные компоненты, необходимые для построения модели генерации видео по тексту.
Уточнение требований
Вот типичный диалог между кандидатом и интервьюером.
Кандидат: Какова ожидаемая длительность сгенерированных видео? Интервьюер: Давайте нацелимся на пятисекундные видео.
Кандидат: Какое разрешение видео мы планируем? Интервьюер: Мы должны стремиться к качеству высокой чёткости, чтобы видео подходило для широкого спектра современных платформ и устройств. Давайте нацелимся на разрешение 720p.
Кандидат: 24 кадра в секунду (FPS) -- желаемая частота кадров для сгенерированного видео? Интервьюер: Да.
Кандидат: Какова ожидаемая задержка генерации видео? Интервьюер: Генерация видео требует значительных вычислительных ресурсов. Для начала несколько минут обработки будет приемлемо. В будущих итерациях мы оптимизируем эффективность и скорость.
Кандидат: Должны ли мы сосредоточиться на определённой категории видео? Интервьюер: Нет, система должна генерировать видео различных жанров и тематик.
Кандидат: Должна ли система поддерживать несколько языков для ввода текста, или мы начинаем только с английского? Интервьюер: Давайте начнём с английского.
Кандидат: Должны ли сгенерированные видео включать аудио? Интервьюер: Давайте пока сосредоточимся на беззвучных видео. Аудио может быть рассмотрено как улучшение в будущих итерациях, но на данном этапе это не приоритет.
Кандидат: Каков примерный размер наших обучающих данных? Интервьюер: У нас есть большой набор данных видео, около 100 миллионов разнообразных видео с подписями. Некоторые подписи могут быть зашумлёнными или не на английском языке.
Кандидат: Распространённый подход к построению модели генерации видео по тексту -- расширение предобученной модели генерации изображений по тексту для обработки видео. Есть ли у нас предобученная модель генерации изображений по тексту? Интервьюер: Да, это разумное допущение.
Кандидат: Учитывая высокие вычислительные требования генерации видео, каков наш бюджет вычислительных ресурсов? Интервьюер: Обучение системы генерации видео требует значительных вычислительных ресурсов. У нас есть более 6000 GPU H100 [2], доступных для обучения генерации видео по тексту.
Кандидат: Необходимо ли обеспечить наличие защитных механизмов для предотвращения генерации оскорбительного или вредного видео? Интервьюер: Отличное замечание. Да, нам нужно обеспечить безопасность нашей системы для пользователей.
Формулировка задачи как задачи ML
В этом разделе формулируется задача генерации видео по тексту как задача ML и выделяются необходимые соображения, выходящие за рамки тех, что использовались в Главе 9 для генерации изображений по тексту.
Определение входных и выходных данных системы
Входные данные -- описательный текст, описывающий сцену, действие или сюжет. Выходные данные -- пятисекундное видео в разрешении 720p (1280 x 720), визуально и темпорально соответствующее заданному текстовому промпту.
Например, при текстовом вводе "Собака играет в fetch в парке в солнечный день" система должна сгенерировать видео, изображающее эту сцену, отражая движение собаки, окружение парка и атмосферу солнечного дня.
Выбор подходящего подхода ML
Генерация видео по тексту схожа по своей природе с генерацией изображений по тексту. Обе задачи создают визуальный контент из текстовых описаний. Техники, такие как авторегрессионное моделирование и диффузионные модели, популярные в генерации изображений по тексту, также эффективны для генерации видео по тексту. Как мы уже видели, диффузионные модели демонстрируют высокую производительность в создании детализированных и реалистичных изображений. Поэтому мы выбираем диффузионную модель для разработки нашей системы генерации видео по тексту.
Однако между ними существует принципиальное различие. Для генерации видео модель должна обрабатывать и генерировать последовательность frame, а не одно изображение. Это значительно увеличивает вычислительную нагрузку. Например, генерация пятисекундного видео при 24 FPS означает, что модель должна создать 120 frame. Генерация изображения 512x512 может занять около 1 секунды на высокопроизводительном GPU, таком как NVIDIA H100, но масштабирование до пятисекундного видео 720p потребует значительно больше времени, так как каждый frame 720p содержит примерно в 3,6 раза больше пикселей. В результате генерация пятисекундного видео 720p может занять около семи минут.
Для решения проблемы сложности и вычислительных затрат генерации видео мы используем популярный подход латентной диффузионной модели (LDM). Этот подход был впервые популяризирован статьёй Stable Diffusion [3] и впоследствии применён и использован большинством моделей генерации видео, таких как Sora от OpenAI [1] и Movie Gen от Meta [4]. Рассмотрим этот подход подробнее.
Латентная диффузионная модель (LDM)
Основная идея LDM заключается в том, что диффузионная модель работает в низкоразмерном latent space, а не непосредственно в пространстве пикселей. Диффузионная модель обучается удалять шум из этих низкоразмерных латентных представлений, а не из исходных видеопикселей обучающего набора данных.
LDM в первую очередь полагается на сеть сжатия для преобразования видеопикселей в латентное представление. Рассмотрим сеть сжатия подробнее.
Сеть сжатия
Сеть сжатия -- это нейронная сеть, которая отображает видеопиксели в latent space. Она принимает необработанное видео на вход и выдаёт сжатое латентное представление, уменьшая как количество frame (temporal размерность), так и разрешение (spatial размерности).
Сеть сжатия обычно основана на модели вариационного автоэнкодера (VAE) [5], которая обучается отдельно от диффузионной модели. Визуальный encoder VAE преобразует входное видео в латентное представление, а визуальный decoder реконструирует исходные видеокадры из latent space.
Как LDM решает проблему вычислительной сложности?
LDM требуют меньше вычислительных ресурсов, чем стандартные диффузионные модели, поскольку обработка сжатых представлений дешевле, чем работа с высокоразмерными пикселями. Чтобы понять влияние этого сжатия, рассмотрим пример.
Представьте, что нам нужно видео с 24 FPS, длительностью пять секунд и разрешением 720p. Это означает 120 frame, каждый с 1280x720 пикселями -- значительный объём данных для обработки. Если мы используем сеть сжатия, аналогичную [4], которая уменьшает как temporal, так и spatial разрешение в 8 раз, spatial размерность видео становится 160x90 пикселей, а temporal размерность сокращается до 15 frame.
Это сжатое представление в 512 раз меньше своего эквивалента в пространстве пикселей, что делает обучение LDM в 512 раз эффективнее. Эта эффективность приводит к более быстрому времени генерации и снижению потребления ресурсов, что особенно ценно при работе с видеоданными высокого разрешения.
Как генерировать видео с помощью обученной LDM
Для генерации видео с помощью обученной LDM мы начинаем с чистого шума в latent space. LDM постепенно уточняет его до очищенного от шума латентного представления. Затем визуальный decoder преобразует это латентное представление обратно в пространство пикселей для создания финального видео.
Для этой главы мы выбираем подход LDM для разработки нашей системы генерации видео по тексту, поскольку он эффективен и снижает вычислительную нагрузку. Чтобы узнать больше о LDM, обратитесь к [6].
Подготовка данных
Набор данных для генерации видео по тексту включает 100 миллионов пар текстовых описаний и соответствующих видео. Эти пары охватывают различные тематики и действия, позволяя модели обучаться на разнообразных видео. В этом разделе мы подготавливаем видео и подписи для обучения нашей LDM.
Подготовка видео
Мы фокусируемся на трёх ключевых шагах подготовки видео для обучения:
- Фильтрация неподходящих видео
- Стандартизация видео
- Предварительное вычисление представлений видео в latent space
Фильтрация неподходящих видео
Большие наборы данных часто содержат нежелательный контент. На этом шаге удаляются неподходящие видео, чтобы модель обучалась только на качественных материалах. Типичные шаги включают:
- Удаление низкокачественных или коротких видео: следуя Movie Gen [4], мы удаляем видео с низким разрешением, слишком короткие, замедленные или искажённые артефактами сжатия.
- Удаление дублирующихся видео (дедупликация): мы используем метод дедупликации, такой как [7], для устранения идентичных видео. Это обеспечивает разнообразие обучающих данных и предотвращает чрезмерное представление определённых видео.
- Удаление вредоносных видео: мы используем модели обнаружения вредоносного контента для выявления и удаления видео с откровенным содержанием. Этот шаг критически важен для того, чтобы наша модель генерации видео по тексту не создавала вредоносные видео.
Стандартизация видео
- Корректировка длины видео: мы разрезаем более длинные видео на пятисекундные клипы, чтобы обучающие данные состояли только из видео одинаковой длины.
- Стандартизация частоты кадров: мы перекодируем видео с более высокой частотой кадров до 24 FPS, чтобы все видео имели одинаковую частоту кадров.
- Корректировка размеров видео: мы изменяем размер и обрезаем видео до стандартного размера, например, 1280x720 пикселей.
Предварительное вычисление представлений видео в latent space
Поскольку LDM работает в latent space, ей нужны только латентные представления на вход. Таким образом, каждая итерация обучения обычно требует следующих шагов:
- Извлечение frame из видео в обучающих данных.
- Прохождение этих frame через предобученную сеть сжатия для получения латентных представлений.
- Использование латентных представлений для продолжения обучения диффузионной модели.
Однако эти шаги неэффективны. Извлечение frame и их сжатие для миллионов видео каждый раз при обучении новой модели замедляет обучение диффузии. Вычисление латентных представлений на лету является ресурсоёмким и времязатратным.
Для оптимизации процесса мы предварительно вычисляем латентные представления для всех видео и кэшируем их в хранилище. Во время обучения диффузионная модель напрямую обращается к предвычисленным латентным представлениям без ожидания процессов извлечения frame или сжатия. Этот подход значительно ускоряет процесс обучения диффузии, при этом затраты на хранение остаются управляемыми. Проведём быстрый расчёт для понимания потребности в хранении.
Примерный расчёт: предположим, что каждый видеокадр при сжатии в латентное представление уменьшается в размере в 512 раз. Таким образом, если видео из 1000 frame занимает около 1000 МБ, его латентное представление займёт около 2 МБ. Если мы кэшируем латентные представления для 100 миллионов видео, общий требуемый объём хранения составит около 200 ТБ. Учитывая современные возможности хранения, это вполне управляемо, особенно по сравнению со значительной экономией времени при обучении.
Подготовка подписей
Важно иметь качественные, согласованные подписи. Некоторые подписи, вероятно, будут отсутствовать или будут нерелевантными. Типичные шаги подготовки подписей:
- Обработка отсутствующих или неанглоязычных подписей: для видео без подписей или с подписями на другом языке мы используем модели, такие как LLaMa3-Video [8] или LLaVA [9], для автоматической генерации описательных подписей.
- Переподписывание: мы улучшаем существующие подписи с помощью предобученных моделей подписывания видео, таких как LLaMa3-Video или LLaVA, для генерации более длинных и детальных версий. Команда Sora [1] показала, что этот процесс необходим для повышения качества и соответствия тексту.
- Предварительное вычисление embedding подписей: обучение диффузионной модели требует embedding подписей для обусловливания. Мы используем текстовый encoder для предварительного вычисления embedding подписей, ускоряя обучение LDM.
Разработка модели
Архитектура
При выборе архитектуры для диффузионной модели генерации видео по тексту у нас есть два основных варианта: U-Net и DiT. Мы рассмотрим каждый из них и определим дополнительные слои, необходимые для расширения их на обработку видео.
U-Net для видео
Кратко рассмотрим архитектуру U-Net, прежде чем расширять её для обработки видео. Как мы рассмотрели в Главе 9, архитектура U-Net состоит из серии блоков понижения разрешения (downsampling), за которой следует серия блоков повышения разрешения (upsampling). Каждый блок понижения разрешения включает 2D-свёртки для обработки и обновления признаков изображения и слой cross-attention для обновления признаков путём внимания к текстовому промпту.
Однако эти слои в основном фокусируются на захвате взаимосвязей между пикселями в пределах одного изображения. Эта конструкция создаёт проблему для видео, где поддержание temporal согласованности критически важно для плавного движения и непрерывности между frame. Текущие слои работают spatial в пределах отдельных frame, а не между ними.
Для решения этого недостатка мы модифицируем архитектуру U-Net, чтобы понимать взаимосвязи между frame. В частности, мы внедряем два часто используемых temporal слоя:
- Temporal attention
- Temporal свёртка
Кратко рассмотрим каждый слой.
Temporal attention: temporal attention использует механизм attention между frame. Каждый признак обновляется путём внимания к релевантным признакам других frame. Рисунок 12 показывает, как определённый признак во frame 2 обновляется путём внимания к признакам других frame.
- Temporal свёртка: temporal свёртка означает применение оператора свёртки к 3D-сегменту данных для захвата temporal размерности. Рисунок 13 иллюстрирует 2D и 3D temporal свёртки.
Подводя итог, для расширения архитектуры U-Net на обработку видео мы можем чередовать слои temporal свёртки и temporal attention в каждом блоке downsampling и upsampling. Эти слои позволяют архитектуре U-Net моделировать движение во входных видео и генерировать последовательность frame, которые temporal согласованы. Чтобы узнать больше о том, как эти слои могут чередоваться, обратитесь к [10].
DiT для видео
В отличие от U-Net, основанного преимущественно на свёртках, DiT в основном опирается на архитектуру Transformer. Как показано на Рисунке 14, DiT состоит из четырёх основных компонентов:
- Patchify
- Positional encoding
- Transformer
- Unpatchify
Рассмотрим каждый компонент и поймём его назначение.
Patchify
Этот компонент преобразует вход в последовательность векторов embedding. Сначала он делит вход на меньшие фрагменты (patches) фиксированного размера. Затем каждый фрагмент выравнивается для формирования последовательности векторов. Выровненные фрагменты затем преобразуются в embedding фрагментов с помощью слоя проекции. Этот шаг критически важен для согласования размера embedding каждого выровненного фрагмента с размером скрытого состояния Transformer.
Процесс patchify аналогичен для входных изображений и видео. Для изображений вход делится на 2D-фрагменты фиксированного размера. Для видео вход делится на 3D-фрагменты.
Positional encoding
Компонент positional encoding создаёт embedding для каждой позиции в исходной последовательности. Эти embedding предоставляют Transformer информацию о расположении каждого фрагмента в исходном входе.
Как мы видели в Главе 2, существуют разные способы кодирования позиций: некоторые методы используют фиксированное positional encoding во время обучения, в то время как другие делают positional encoding обучаемым. Также существуют разные способы назначения позиций каждому фрагменту. Например, можно присвоить каждому фрагменту одно число для обозначения его места в последовательности или использовать 3D-координаты (2D для изображений) для указания местоположения каждого фрагмента в пространстве и времени.
Не существует единственно лучшего способа positional encoding. Часто необходимо проводить эксперименты, чтобы найти подход, наиболее эффективный для данных и задачи. В этой главе мы следуем OpenSora [11] и используем RoPE [12] positional encoding. Чтобы узнать больше о positional encoding в моделях генерации видео по тексту, обратитесь к [4].
Transformer
Transformer обрабатывает последовательность embedding и другие обусловливающие сигналы, такие как текстовый промпт, для прогнозирования шума для каждого фрагмента.
Unpatchify
Unpatchify преобразует предсказанные векторы шума обратно к исходным размерностям входа. Он включает LayerNorm для нормализации, линейный слой для корректировки длины вектора и операцию изменения формы (reshape) для формирования финального выхода.
U-Net vs. DiT
Обе архитектуры, U-Net и DiT, доказали свою эффективность для генерации видео по тексту. Архитектура U-Net существует дольше и была тщательно протестирована. Популярные модели генерации видео по тексту на основе U-Net включают Stable Video Diffusion [13] от Stability AI и EMU video [14] от Meta.
Архитектура DiT более новая и продемонстрировала большой потенциал с превосходными результатами. DiT лучше работает при увеличении объёма данных и вычислительных мощностей благодаря масштабируемой природе Transformer. Кроме того, она имеет гибкую архитектуру, что облегчает адаптацию к видео и другим модальностям ввода. Movie Gen от Meta и Sora от OpenAI -- популярные модели, основанные на архитектуре DiT.
В этой главе мы следуем Sora и выбираем архитектуру DiT.
Обучение
Обучение видеодиффузионной модели очень похоже на обучение диффузионной модели для изображений. Во время обучения мы добавляем шум к исходному видео, моделируя прямой процесс, и обучаем модель предсказывать добавленный шум. Три конкретных шага, выполняемых за одну итерацию обучения:
- Добавление шума: случайно выбирается временной шаг для определения уровня добавления шума. Выбранный временной шаг используется для добавления шума к входному видео.
- Прогнозирование шума: модель DiT получает зашумлённое видео на вход и прогнозирует добавленный шум на основе обусловливающих сигналов, таких как текстовый промпт и выбранный временной шаг.
- Вычисление потерь: потеря вычисляется путём сравнения предсказанного шума с фактическим шумом.
Для более детального обзора обучения диффузии обратитесь к Главе 9.
Целевая функция ML и функция потерь
Основная функция потерь -- это потеря реконструкции, вычисляемая с использованием формулы среднеквадратичной ошибки (MSE). Эта потеря измеряет разницу между предсказанным шумом и фактическим шумом, побуждая модель точно предсказывать добавленный шум. Целевая функция ML -- минимизировать потерю реконструкции, что ведёт к точной реконструкции видео.
Исследователи экспериментировали с добавлением других функций потерь для улучшения производительности генерации видео по тексту. Чтобы узнать больше, обратитесь к [4].
Проблемы обучения видеодиффузионных моделей
Обучение модели DiT для генерации видео по тексту связано с рядом проблем и проектных решений. В этом разделе рассматриваются две важные проблемы:
- Нехватка крупномасштабных данных видео-текст
- Вычислительная стоимость генерации видео высокого разрешения
Нехватка крупномасштабных обучающих данных видео-текст
Обучение больших моделей требует большого объёма данных. В отличие от обучения моделей генерации изображений по тексту, где доступно огромное количество пар изображение-текст, парные данные видео-текст ограничены. Этот дефицит создаёт проблему при обучении эффективных моделей генерации видео.
Существуют две распространённые стратегии решения проблемы нехватки крупномасштабных данных:
- Обучение модели DiT на данных изображений и видео одновременно: эта стратегия рассматривает каждое изображение как одно-frame видео, позволяя модели обучаться как на парах изображение-текст, так и на парах видео-текст.
- Предобучение модели DiT на данных изображений: эта стратегия сначала предобучает модель DiT на парах изображение-текст для использования обширных данных изображений и построения прочной визуальной основы. Затем предобученная модель дообучается на парах видео-текст для генерации видео.
Обе стратегии используют сотни миллионов пар изображение-текст при обучении, позволяя модели DiT обучаться как на изображениях, так и на видео. Для простоты мы выбираем первую стратегию, так как она требует только одного этапа обучения. Однако обе стратегии могут быть эффективны на практике.
Вычислительная стоимость генерации видео высокого разрешения
Как обсуждалось ранее, обработка и генерация видео дороже, чем изображений. Это в первую очередь связано с тем, что видео обычно содержит сотни frame, что делает процесс медленнее и дороже. Генерация видео высокого разрешения, такого как 720p или 1080p, добавляет дополнительную сложность.
Вот несколько распространённых стратегий снижения вычислительных затрат при обучении моделей генерации видео высокого разрешения:
- Использование подхода на основе LDM: вместо обучения модели DiT непосредственно в пространстве пикселей мы используем сеть сжатия для преобразования видео из пространства пикселей в низкоразмерное latent space. Обучение диффузионной модели в этом latent space снижает вычислительную нагрузку.
- Предварительное вычисление видеопредставлений: предварительно вычисляя видеопредставления в latent space перед обучением, мы избегаем повторяющихся вычислений во время обучения. Использование этих кэшированных данных ускоряет процесс обучения.
- Использование модели spatial super-resolution: как предложено Google в "Imagen video" [15], мы используем отдельно обученную модель для увеличения разрешения сгенерированных видео. Модель DiT генерирует видео в более низком разрешении, которое затем улучшается до желаемого разрешения моделью spatial super-resolution. Например, модель DiT может генерировать видео в 720p, а модель spatial super-resolution может затем масштабировать их до 1080p или 4K.
- Использование модели temporal super-resolution: как предложено в [15], мы используем модель для увеличения temporal разрешения путём интерполяции между frame. Например, если видео должно быть пять секунд при 24 FPS (то есть 120 frame всего), модель DiT может сгенерировать его при 12 FPS (60 frame), а модель temporal super-resolution затем интерполирует до 24 FPS.
- Использование более эффективных архитектур: мы можем принять эффективную реализацию механизма attention [16] для снижения вычислительной нагрузки при обучении. Кроме того, такие техники, как Mixture of Experts (MoE) [17], могут использоваться для ускорения процесса обучения.
- Использование распределённого обучения: мы используем техники распределённого обучения, такие как тензорный параллелизм, для параллелизации обучения на нескольких устройствах. Разделяя модель, данные или и то и другое между различными устройствами, мы можем значительно ускорить обучение и более эффективно обрабатывать большие наборы видеоданных. Этот подход особенно полезен для генерации видео высокого разрешения, где требования к памяти и вычислениям существенны. Для обзора распределённого обучения обратитесь к Главе 1.
Сэмплирование
Процесс сэмплирования в диффузионных моделях начинается со случайного шума, и модель итеративно очищает сэмпл от шума, пока не будет получено полностью очищенное видеопредставление в latent space. Для более подробной информации о сэмплировании в диффузионных моделях обратитесь к Главе 9.
Оценка
Офлайн-метрики оценки
Единый бенчмарк критически важен для оценки моделей генерации видео. VBench [18] и Movie Gen Bench [19] обеспечивают это, предоставляя отобранный набор промптов, предназначенных для тестирования различных аспектов генерации видео, таких как когерентность движения, temporal согласованность и сложность сцены. Мы можем использовать эти бенчмарки для измерения того, насколько хорошо модель создаёт реалистичные и плавные видео, фокусируясь на качестве видео, точности движений и переходах между сценами. Рассмотрим как автоматические метрики, так и человеческую оценку, сосредоточившись на трёх ключевых областях:
- Качество frame
- Temporal согласованность
- Соответствие видео тексту
Качество frame
Качество frame означает измерение качества каждого frame независимо. Для измерения этого качества мы используем FID [20] и Inception score (IS) [21], которые широко применяются для изображений. Общее качество вычисляется путём усреднения оценок FID и IS по всем frame. Также могут использоваться другие метрики, такие как LPIPS [22] и KID [23].
Хотя FID и IS измеряют качество отдельных frame, они не учитывают temporal согласованность в сгенерированных видео. Например, видео может иметь высококачественные frame, но при этом не иметь плавных переходов, что даёт высокий FID score без визуальной когерентности. Рассмотрим temporal согласованность и распространённые метрики для её измерения.
Temporal согласованность
Temporal согласованность означает, насколько плавно визуальное содержимое переходит от одного frame к следующему. Оценка temporal согласованности важна для обеспечения естественного потока сгенерированного видео. Распространённая метрика для измерения temporal согласованности -- это Frechet Video Distance (FVD).
FVD
FVD [24], являющийся расширением FID, оценивает как визуальное качество, так и temporal согласованность видео. Он сравнивает статистическое распределение сгенерированных видео с реальными видео в пространстве embedding.
Вот пошаговое руководство по вычислению FVD score:
- Генерация видео: мы начинаем с генерации большого набора видео с помощью модели, которую хотим оценить. Эти видео будут сравниваться с набором реальных видео для оценки их качества и согласованности.
- Извлечение признаков: мы пропускаем каждое видео (как сгенерированное, так и реальное) через предобученную модель I3D [25] и извлекаем признаки из определённого слоя. Модель I3D расширяет архитектуру Inception v3 [26] на последовательные данные путём обучения на распознавании действий.
- Вычисление среднего и ковариации: мы вычисляем среднее и ковариацию извлечённых признаков отдельно для сгенерированных и реальных видео. Эти статистические меры суммируют распределение признаков для обоих наборов видео.
- Вычисление расстояния Фреше: мы вычисляем FVD score как расстояние Фреше между средним и ковариацией сгенерированных и реальных видео. Расстояние Фреше измеряет, насколько близки два распределения.
Более низкий FVD score указывает на большее сходство между распределениями, означая, что сгенерированные видео более реалистичны и temporal согласованы.
Соответствие видео тексту
Соответствие видео тексту означает, насколько точно сгенерированное видео отражает текстовое описание, которым оно было обусловлено.
Широко используемая метрика для измерения соответствия видео тексту -- это CLIP similarity score, вычисляемый следующим образом:
- Извлечение признаков на уровне frame: мы пропускаем каждый frame видео через предобученный CLIP image encoder для получения визуальных признаков. Текст кодируется с помощью текстового encoder для получения текстовых признаков.
- Вычисление сходства: для каждого frame мы вычисляем косинусное сходство между его визуальными признаками и текстовыми признаками. Эта оценка показывает, насколько хорошо содержимое frame соответствует тексту.
- Агрегация покадровых оценок сходства: мы агрегируем эти оценки сходства для получения единой оценки, представляющей общее соответствие видео тексту. Агрегация может выполняться путём усреднения, взятия максимальной оценки или с использованием других статистических методов.
Высокий CLIP similarity score указывает на то, что сгенерированные видео соответствуют своим текстовым описаниям.
Человеческая оценка
Наряду с описанными автоматическими метриками, человеческая оценка по-прежнему жизненно важна для оценки генеративных моделей, поскольку она обеспечивает субъективную оценку, дополняющую автоматические метрики.
Для человеческой оценки мы генерируем видео из тестовых промптов с помощью двух различных моделей. Затем мы представляем пары видео, по одному от каждой модели, аннотаторам-людям. Они выбирают лучшее видео, оценивая соответствие видео тексту, качество видео и temporal согласованность. Этот процесс позволяет сравнить две модели и определить, какая из них работает лучше.
Онлайн-метрики оценки
Онлайн-метрики оценки для моделей генерации видео по тексту аналогичны метрикам для моделей генерации изображений по тексту. Важные метрики включают:
- Кликабельность (click-through rate)
- Время, проведённое на странице
- Обратная связь от пользователей
- Конверсия
Эти метрики помогают оценить вовлечённость пользователей, удовлетворённость и общую производительность модели в продакшене.
Общий дизайн системы ML
В этом разделе мы рассмотрим целостный дизайн системы генерации видео по тексту. В частности, мы рассмотрим следующие конвейеры:
- Конвейер данных
- Конвейер обучения
- Конвейер инференса
Конвейер данных
Конвейер данных подготавливает обучающие данные путём фильтрации непригодных изображений и видео, их стандартизации, а также предварительного вычисления и сохранения латентных представлений. Он обеспечивает релевантность и детальность подписей путём переподписывания и использования предобученного текстового encoder для предварительного вычисления и сохранения embedding подписей.
Конвейер обучения
Конвейер обучения обучает модель, используя обучающие данные, подготовленные конвейером данных.
Конвейер инференса
Конвейер инференса обрабатывает запросы пользователей в реальном времени для генерации видео из текстовых промптов. Как показано на Рисунке 24, он имеет несколько критически важных компонентов, обеспечивающих качество и безопасность системы.
Большинство компонентов аналогичны тем, что мы рассмотрели в Главе 9 для генерации изображений по тексту. Уникальные компоненты для генерации видео по тексту:
- Визуальный decoder
- Temporal super-resolution
Визуальный decoder
LDM генерирует выход в latent space, а не в пространстве пикселей. Затем визуальный decoder использует сеть сжатия для преобразования этого латентного представления обратно в пространство пикселей.
Temporal super-resolution
Этот компонент интерполирует между сгенерированными frame, обеспечивая более плавное движение в видео.
Дополнительные темы для обсуждения
Если в конце интервью останется дополнительное время, можно обсудить следующие темы:
- Обеспечение гибкости сэмплирования для различных длительностей, разрешений и соотношений сторон [1].
- Расширение модели генерации видео по тексту на смежные приложения, такие как inpainting, outpainting, стилизация видео-в-видео, интерполяция frame, super-resolution и анимация изображений (изображение-в-видео) [10].
- Поддержка управления сгенерированными видео, например, уровень желаемого движения и тип движения (камера vs. объект) [27].
- Использование техник прогрессивной дистилляции для снижения вычислительных требований обучения [28].
- Детали моделей spatial и temporal super-resolution [15].
- Детали модели переподписывания [9][8].
- Различные планировщики шума [29].
- Техники аугментации шума с обусловливанием [30].
- Персонализация модели генерации видео по тексту для конкретного субъекта [31].
- ControlNet для моделей генерации видео по тексту [32].
- Детали метода Stable Cascade [33].
- Детали визуальной сети сжатия [13].
Резюме
Справочные материалы
[1] Video generation models as world simulators. https://openai.com/index/video-generation-models-as-world-simulators/. [2] H100 Tensor Core GPU. https://www.nvidia.com/en-us/data-center/h100/. [3] High-Resolution Image Synthesis with Latent Diffusion Models. https://arxiv.org/abs/2112.10752. [4] Meta Movie Gen. https://ai.meta.com/research/movie-gen/. [5] Auto-Encoding Variational Bayes. https://arxiv.org/abs/1312.6114. [6] The Illustrated Stable Diffusion. https://jalammar.github.io/illustrated-stable-diffusion/. [7] On the De-duplication of LAION-2B. https://arxiv.org/abs/2303.12733. [8] The Llama 3 Herd of Models. https://arxiv.org/abs/2407.21783. [9] LLaVA-NeXT: A Strong Zero-shot Video Understanding Model. https://llava-vl.github.io/blog/2024-04-30-llava-next-video/. [10] Lumiere: A Space-Time Diffusion Model for Video Generation. https://arxiv.org/abs/2401.12945. [11] OpenSora Technical Report. https://github.com/hpcaitech/Open-Sora/blob/main/docs/report_02.md. [12] RoFormer: Enhanced Transformer with Rotary Position Embedding. https://arxiv.org/abs/2104.09864. [13] Stable Video Diffusion: Scaling Latent Video Diffusion Models to Large Datasets. https://arxiv.org/abs/2311.15127. [14] Emu Video: Factorizing Text-to-Video Generation by Explicit Image Conditioning. https://arxiv.org/abs/2311.10709. [15] Imagen Video: High Definition Video Generation with Diffusion Models. https://arxiv.org/abs/2210.02303. [16] HyperAttention: Long-context Attention in Near-Linear Time. https://arxiv.org/abs/2310.05869. [17] Mixture of Experts Explained. https://huggingface.co/blog/moe. [18] VBench: Comprehensive Benchmark Suite for Video Generative Models. https://vchitect.github.io/VBench-project/. [19] Movie Gen Bench. https://github.com/facebookresearch/MovieGenBench. [20] FID calculation. https://en.wikipedia.org/wiki/Fr%C3%A9chet_inception_distance. [21] Inception score. https://en.wikipedia.org/wiki/Inception_score. [22] The Unreasonable Effectiveness of Deep Features as a Perceptual Metric. https://arxiv.org/abs/1801.03924. [23] Demystifying MMD GANs. https://arxiv.org/abs/1801.01401. [24] Towards Accurate Generative Models of Video: A New Metric & Challenges. https://arxiv.org/abs/1812.01717. [25] Quo Vadis, Action Recognition? A New Model and the Kinetics Dataset. https://arxiv.org/abs/1705.07750. [26] Rethinking the Inception Architecture for Computer Vision. https://arxiv.org/abs/1512.00567. [27] Moonshot: Towards Controllable Video Generation and Editing with Multimodal Conditions. https://arxiv.org/abs/2401.01827. [28] Progressive Distillation for Fast Sampling of Diffusion Models. https://arxiv.org/abs/2202.00512. [29] Schedulers. https://huggingface.co/docs/diffusers/v0.9.0/en/api/schedulers. [30] Photorealistic Text-to-Image Diffusion Models with Deep Language Understanding. https://arxiv.org/abs/2205.11487. [31] CustomVideo: Customizing Text-to-Video Generation with Multiple Subjects. https://arxiv.org/abs/2401.09962. [32] Control-A-Video: Controllable Text-to-Video Generation with Diffusion Models. https://controlavideo.github.io/. [33] Introducing Stable Cascade. https://stability.ai/news/introducing-stable-cascade.