8000
Skip to content

Latest commit

 

History

13 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 

Repository files navigation

SASRec on Yandex Yambda

Open In Colab License: MIT PyTorch Dataset: Yambda NDCG

Decoder-only трансформер для sequential recommendation, causal self-attention реализован вручную. Аудио-эмбеддинги треков квантованы в semantic IDs через RQ-VAE: как замена item-эмбеддингов они проигрывают, как дополнение поднимают NDCG@10 на 39%. Ablation из шести моделей на Yandex Yambda, всё в бесплатном Colab T4.

TL;DR

Рекомендательная модель обычно заводит обучаемый вектор на каждый трек: 181 тысяча треков, 11.6 миллиона параметров только в таблице эмбеддингов. У такой модели есть слепое пятно: про новый трек она не знает ничего, потому что его вектор ещё не обучен.

Идея semantic IDs в том, чтобы вывести представление трека из его звучания, а не из истории прослушиваний. Аудио-эмбеддинг сжимается в кортеж из нескольких небольших целых чисел, вроде (142, 87, 231), словарь схлопывается со 181K айтемов до 1026 токенов, и рекомендация превращается в генерацию: модель предсказывает следующий трек код за кодом, как языковая модель порождает текст.

Результат: заменить item-эмбеддинги кодами не получается (NDCG@10 падает с 0.413 до 0.181, ниже даже popularity-baseline). Но если использовать коды вместе с item-эмбеддингами, качество растёт на 39.2% относительно базовой модели (0.413 → 0.575), и прирост подтверждается на отложенной выборке (+40.2%).

Что здесь можно посмотреть помимо метрик:

  • отрицательный результат не оставлен без объяснения: per-position perplexity показывает, что из четырёх кодов из истории предсказуем только первый (49.7 из 257 возможных), а второй и третий модель угадывает почти случайно (152.2 и 201.3 из 256), и это следствие того, что RQ-VAE оптимизировался на реконструкцию аудио, а не на предсказуемость из поведения;
  • поймана ошибка в метрике: примитивный last-item baseline показывал HitRate@10 = 1.0000, потому что при равных скорах ранг считался в его пользу; после честного деления ничьих он упал до 0.0145;
  • качество токенизации отделено от качества рекомендаций: коды осмысленно кодируют звучание (косинус внутри кластера 0.81–0.89 против 0.10–0.20 до чужих треков), и это не мешает модели на одних кодах проигрывать;
  • аудио-эмбеддинги для 181K треков извлечены из файла на 13.8 GB потоковым чтением за 3.2 минуты без скачивания, через footer parquet и построчную фильтрацию;
  • начальный loss сверен с теоретическим в двух разных постановках (1.3904 против 1.3863 для BCE с негативом, 6.9373 против 6.9334 для полного softmax), что отлавливает ошибки инициализации до обучения;
  • реализован constrained beam search по префиксному дереву: генерация идёт только по путям, соответствующим реальным трекам.

Данные музыкальные, но ничего специфически музыкального в постановке нет: нужна лишь последовательность взаимодействий и контентное представление объекта.

Задача

Дана история пользователя $[i_1, \dots, i_n]$, упорядоченная по времени последовательность треков. Требуется предсказать $i_{n+1}$.

Это next-token prediction, только вместо слов идентификаторы треков. Отсюда архитектура GPT: decoder-only трансформер с causal self-attention, positional embeddings и предсказанием каждого следующего элемента за один проход.

Формально модель выдаёт представление состояния пользователя после $t$-го события:

$$h_t = f_\theta(i_1, \dots, i_t) \in \mathbb{R}^{d}$$

а скор трека $j$ как кандидата на позицию $t+1$ это скалярное произведение с его представлением:

$$s(j \mid h_t) = h_t^\top v_j$$

Всё дальнейшее про то, как получается $v_j$. Трансформер не меняется ни разу.

Что где в ноутбуке

Компонент Что применяется Секция
Causal self-attention Q/K/V вручную, две маски, обработка полностью замаскированных строк 5
Leave-last-out split Разбиение по времени пер-юзер, разбор возможной утечки 3
Sampled evaluation Протокол 1 против 100, честная обработка ничьих в ранге 4
Потоковое чтение parquet Footer, отбор колонок, iter_batches по файлу 13.8 GB 7
RQ-VAE Остаточное квантование, straight-through, k-means инициализация 8
Generative retrieval Semantic IDs как последовательность токенов, полный softmax 9
Constrained beam search Генерация по префиксному дереву валидных путей 10
Ablation Шесть моделей на одном протоколе с одним seed 11, 12

Данные

Yandex Yambda, подмножество likes из 50M-версии.

Параметр Значение
Событий 881 456
Пользователей (после фильтра < 5 событий) 6 951
Уникальных треков 180 942
Длина истории: медиана / 95-й перцентиль / максимум 44 / 398 / 8 697
Окно контекста MAX_LEN 100 треков
Аудио-эмбеддинги 128-мерные, покрытие 94.2%

Выбраны лайки, а не прослушивания: listens в этой версии содержит 46.5M событий, что избыточно для бесплатного Colab, а сигнал там шумный (автоплей, фон). Лайк это явное действие, поэтому последовательности чище.

Датасет полностью анонимизирован: только числовые идентификаторы, без названий треков и имён исполнителей. Единственная контентная информация это аудио-эмбеддинги, полученные свёрточной сетью на спектрограммах.

Про пропуски. У 10 474 треков (5.8%) аудио-эмбеддинга нет, но на них приходится лишь 3.7% взаимодействий, а среди ста самых популярных треков без аудио остался один. Пропуски смещены в хвост, поэтому им выдан выделенный первый код 256 со значением «звучание неизвестно»: модель учится опираться для них только на поведенческий сигнал.

Архитектура

Вектор трека

Здесь и проходит граница между конфигурациями:

$$v_j = \underbrace{e_j}_{\text{обучаемый вектор}} + \underbrace{\sum_{p=1}^{4} C_p\left[k_p(j)\right]}_{\text{эмбеддинги аудио-кодов}}$$

Оба слагаемых включаются флагами, так что три конфигурации (только $e_j$, только коды, оба вместе) получаются из одного класса. Архитектура, loss, процедура обучения и протокол оценки при этом совпадают до строчки.

Residual-квантование

Аудио-эмбеддинг $x \in \mathbb{R}^{128}$ проходит через энкодер в $z \in \mathbb{R}^{32}$, дальше квантуется последовательно:

$$r_0 = z, \qquad k_p = \arg\min_{c} \left\lVert r_{p-1} - C_p[c] \right\rVert, \qquad r_p = r_{p-1} - C_p[k_p]$$

$$\hat{z} = \sum_{p=1}^{3} C_p[k_p]$$

Каждый следующий кодбук приближает то, что не смог передать предыдущий. Отсюда иерархия: первый код отвечает за грубую область пространства звучаний, следующие уточняют.

Операция $\arg\min$ недифференцируема, поэтому используется straight-through estimator: на forward идёт квантованное значение, на backward градиент проходит насквозь, z_q = z + (z_q - z).detach().

Три уровня по 256 кодов дают около 16.7M различимых комбинаций при словаре всего в 768 обучаемых векторов. Четвёртый код добавлен как разводящий: внутри группы коллизий трекам присваивается порядковый номер, семантики он не несёт.

Функция потерь

Полный softmax по 181K айтемов не помещается в T4: матрица логитов для батча 128 и окна 100 весит около 9 GB. Поэтому negative sampling, один случайный негатив на позицию, BCE. Отдельного выходного слоя нет, для скоринга переиспользуется то же представление, что на входе (weight tying).

В генеративной постановке словарь схлопывается до 1026 токенов, и полный softmax становится дешёвым. Это главный технический выигрыш токенизации, независимо от её влияния на метрики.

Результаты

Сводная таблица (test, 6 951 пользователей)

Протокол: 1 настоящий следующий трек против 100 случайных негативов, один seed для всех моделей, ничьи в ранге делятся поровну.

8000
Модель Представление трека NDCG@10 HitRate@10
last-item последний трек истории 0.0145 0.0145
popularity частота в train 0.4061 0.6179
SASRec на semantic IDs 4 кода как последовательность токенов 0.1809 0.3371
SASRec на аудио-кодах те же 4 кода как признаки трека 0.4024 0.6621
SASRec на item-эмбеддингах обучаемый вектор 0.4130 0.6255
SASRec гибрид вектор плюс коды 0.5750 0.7911

Гибрид превосходит модель на item-эмбеддингах на 39.2% по NDCG@10 и на 26.5% по HitRate@10.

У конфигурации на аудио-кодах метрики расходятся: по NDCG она чуть ниже popularity (0.4024 против 0.4061), по HitRate заметно выше (0.6621 против 0.6179). Аудио-семантика затаскивает нужный трек в топ-10, но не выводит его на первые позиции.

Сравнение моделей

Проверка на отложенной выборке

Гибрид имеет больше параметров, чем обе изолированные конфигурации, поэтому прирост перепроверен на val.

Модель NDCG test NDCG val разница
item embeddings 0.4130 0.4258 +0.0129
audio codes 0.4024 0.4193 +0.0170
hybrid 0.5750 0.5969 +0.0219

Сдвиг одинаковый у всех трёх моделей и по знаку, и по величине. Val-цель стоит ближе к концу обучающей истории и потому предсказывается легче: на переобучение это не похоже. Прирост гибрида держится на обеих выборках, +39.2% на test и +40.2% на val.

Разрыв между val и test растёт с числом параметров: 0.0129 у item-эмбеддингов, 0.0219 у гибрида.

Вклад уровней квантования

Кодов учтено в скоре NDCG@10 HitRate@10
первый 0.0802 0.2060
первые два 0.0955 0.2228
первые три 0.1279 0.2643
все четыре 0.1809 0.3371

Прирост монотонный, и его даёт даже разводящий код, лишённый семантики. Шума коды, таким образом, не добавляют.

Размер моделей

Модель Параметров Из них в таблице представлений
SASRec на item-эмбеддингах 11 686 848 11 580 352 (99.1%)
SASRec на semantic IDs 258 306 65 664 (25.4%)

Генеративная модель легче в 45 раз. У первой 99% весов уходит на запоминание отдельных треков, у второй большую часть занимает сам трансформер: то, что раньше лежало в параметрах, теперь закодировано в структуре кодов.

Ключевой вывод: почему semantic IDs не работают как замена

Разрыв между 0.181 и 0.413 объясняется тем, как они используются.

Шаг 1. Где узкое место. В генеративной постановке модель предсказывает четыре токена на трек, и скор кандидата это сумма логарифмов вероятностей всех четырёх кодов. Достаточно одного плохо предсказуемого уровня, чтобы просел весь кортеж.

Шаг 2. Измерение. Перплексия посчитана отдельно по каждой позиции кода:

Позиция Perplexity Размер кодбука Насколько лучше случайного
код 1 (грубая семантика) 49.7 257 в 5.2 раза
код 2 (уточнение) 152.2 256 в 1.7 раза
код 3 (уточнение) 201.3 256 в 1.3 раза
код 4 (разводящий) 2.2 256 вырожден: у 82% треков значение 0

Содержательно выучен только первый уровень. Второй и третий модель угадывает почти случайно, а низкая перплексия четвёртого это перекос распределения.

Предсказуемость уровней квантования

Шаг 3. Почему так. RQ-VAE обучался на реконструкцию аудио: его функция потерь это MSE между исходным вектором и восстановленным. Первый уровень при этом кодирует грубую область звучания, и она выводится из поведения (слушал электронику, дальше вероятна электроника). Нижние уровни кодируют мелкие акустические детали остатка, которых в истории прослушиваний нет в принципе. Токенизация оптимизирована под одну задачу, а используется в другой.

Шаг 4. Чего это не означает. Непредсказуемость мешает ровно тогда, когда коды нужно угадать. Если подать те же коды как признаки трека, модель их не предсказывает, а считывает с кандидата, и слабая предсказуемость перестаёт мешать. Конфигурация на аудио-кодах даёт 0.4024 против 0.4130 у обучаемых эмбеддингов, почти догоняя их при 45-кратной разнице в числе параметров.

На этом же держится и гибрид: item-эмбеддинг хранит поведенческую историю трека, коды описывают его звучание, и вместе они дают 0.5750, больше суммы вкладов по отдельности.

Токенизация осмысленна, и это отдельный вопрос

Провал модели на кодах легко списать на плохую токенизацию. Дело не в ней.

UMAP аудио-эмбеддингов

Цвет это первый код RQ-VAE, серым показаны треки вне десятки крупнейших кластеров (около 95% выборки). Визуальная плотность подтверждается численно в исходном 128-мерном пространстве:

Метрика Значение
Косинус внутри кластера 0.809 – 0.887
Косинус до чужих треков 0.103 – 0.197
Косинус между центроидами ближайшей пары кодов 0.825
Медиана по всем парам топ-10 0.096

Разброс между внутрикластерным и межкластерным косинусом почти на порядок, так что разделение не сводится к артефакту проекции. Перекрытия, где они есть, отражают иерархию: первый уровень нарезает непрерывное пространство звучаний, соседние области неизбежно граничат, и у ближайшей пары кодов косинус между центроидами (0.825) на порядок выше медианного (0.096).

Качество реконструкции 0.9272 по косинусу (медиана 0.9355) означает, что три числа сохраняют почти всю информацию о звучании, четвёртый семантический уровень не нужен.

Про проекцию: t-SNE с евклидовой метрикой на этих данных рвал кластеры на части, потому что векторы нормированы и лежат на сфере, а правильная мера близости здесь угловая. UMAP с metric='cosine' даёт цельную картину.

Инженерные детали, на которых легко ошибиться

Извлечение 93 MB из файла на 13.8 GB. Аудио-эмбеддинги Yambda лежат одним parquet на 7.72M треков. Скачивать его на Colab невозможно, а нужны эмбеддинги только для 181K треков, то есть около 2% файла. Помогают три свойства колоночного формата. Footer с метаданными читается одним HTTP-range-запросом, без тела файла. Колонки читаются независимо, поэтому берётся только embed (6.0 GB) вместо normalized_embed (8.1 GB), а нормализация делается локально одной строкой. iter_batches тянет файл кусками, и фильтрация идёт на лету, так что в памяти оседают только нужные векторы. 3.2 минуты потокового чтения против скачивания 13.8 GB.

Ловушка с NaN в attention. Если у примера вся последовательность это паддинг (в батче такое встречается), у соответствующей строки запрещены все ключи, softmax по сплошным $-\infty$ даёт $0/0$, и одна строка через residual-связи отравляет весь батч. Без torch.nan_to_num после softmax обучение падает в NaN на первом же шаге.

Ничьи в ранжировании. Наивный подсчёт «сколько кандидатов строго выше настоящего» при равенстве скоров отдаёт первое место настоящему треку даром. Для непрерывных скоров это ничего не меняет, но last-item baseline выдаёт всего два значения, 0 или 1, и потому показывал HitRate@10 = 1.0000. После деления ничьих поровну (ранг = «строго выше» плюс половина от числа равных) он упал до 0.0145.

Калибровка начального loss. Приём из разбора Карпатого: посчитать, каким должен быть loss у необученной модели, и сверить с фактическим. Для BCE с одним негативом это $2\ln 2 \approx 1.3863$, фактически 1.3904. Для полного softmax по 1026 токенам это $\ln 1026 \approx 6.9334$, фактически 6.9373. Расхождение здесь означало бы поломанную инициализацию или неверное маскирование паддинга, и ловится это до первой эпохи.

Мёртвые коды в RQ-VAE. При случайной инициализации часть векторов кодбука никогда не оказывается ближайшей ни к чему и перестаёт обучаться, поэтому кодбуки прогреваются k-means на реальных остатках. По ходу обучения доля живых кодов первого уровня падает со 100% до 57% и возвращается к 90%: энкодер уходит от инициализации, старая геометрия рассыпается, кодбуки пересобираются под новую.

Ограничения

  • Sampled evaluation оптимистичен. Ранжирование против 100 негативов проще, чем против каталога в 181K. Числа сопоставимы между моделями, но не с метриками из статей, считанными по другому протоколу.
  • Негативы не исключают историю пользователя. В оригинальной статье их сэмплируют из треков, с которыми пользователь не взаимодействовал, здесь равномерно по каталогу. Смещение работает против моделей, метрики скорее занижены.
  • Один прогон, один seed. Доверительных интервалов нет; различия в третьем знаке ничего не значат.
  • Коллизии кодов остаются высокими. 28.5% треков делят кортеж из трёх семантических кодов с кем-то ещё, в крупнейшей группе 242 трека. Разводящий код решает различимость, но семантики не несёт, что и видно по его перплексии.
  • Кодовая схема зависит от состояния глобального RNG. При прогоне сверху вниз результат воспроизводится, но переобучение RQ-VAE после других операций с генератором сдвинет и коллизии, и размер словаря токенов.
  • Обучение только на likes, 881K событий вместо 46.5M в listens. Упирается во время прогона, не в код: смена подмножества это один аргумент загрузки.
  • Окно 100 треков. Внимание стоит $O(L^2)$, а в генеративной постановке трек занимает 4 позиции. Три четверти пользователей помещаются целиком, у остальных история обрезается.

Что можно сделать дальше

Обучать RQ-VAE на предсказуемость, а не на реконструкцию. Если оптимизировать квантователь совместно с рекомендательной моделью, нижние уровни начнут кодировать выводимое из поведения, а не мелкие акустические детали. Постановка при этом становится сквозной вместо двухэтапной.

Подключить artist_id и album_id. Они доступны в Yambda и почти наверняка добавят сигнала, особенно для тех 5.8% треков, у которых нет аудио и которые сейчас обслуживаются заглушкой.

Больше негативов на позицию. Сейчас один, как в оригинальной статье. In-batch negatives или sampled softmax с $k \gg 1$ обычно поднимают качество без архитектурных изменений.

Прогон на listens. 46.5M событий и словарь около миллиона айтемов, где маленький выходной слой генеративной постановки наконец имеет смысл.

Структура и запуск

SASRec.ipynb        # весь проект, 12 секций
README.md
LICENSE
.gitignore
assets/
    umap.png        # кластеры кодов RQ-VAE (секция 8)
    ablation.png    # сводное сравнение моделей (секция 12)
    perplexity.png  # предсказуемость уровней квантования (секция 12)

Запуск. Откройте ноутбук в Colab по бейджу выше, включите GPU (Runtime → Change runtime type → T4 GPU) и запустите Run all. Данные скачиваются автоматически с HuggingFace, регистрация не нужна.

Окружение. Бесплатного Colab с T4 (16 GB) достаточно. Дополнительно ставятся datasets, pyarrow, umap-learn, всё остальное предустановлено.

Время. Полный прогон около 15 минут: извлечение аудио-эмбеддингов (~3 мин), RQ-VAE (~1 мин), четыре обучения (~3 мин суммарно), оценка semantic IDs (~3 мин), остальное быстро.

Чекпойнты. Пишутся на Google Drive после каждой эпохи вместе с состоянием оптимизатора, поэтому обрыв сессии Colab не требует начинать заново. Разбиение данных, извлечённые эмбеддинги и коды кэшируются там же. Если RQ-VAE переобучался, старые веса окажутся несовместимы с новой схемой кодов: для этого случая в секции Setup есть флаг RESET_CHECKPOINTS.

Литература

About

SASRec с нуля на PyTorch (causal self-attention вручную) на Yandex Yambda. Аудио-эмбеддинги квантованы в semantic IDs через RQ-VAE: как замена item-эмбеддингов проигрывают, как дополнение дают +39% к NDCG@10 (0.575 vs 0.413). Ablation из шести моделей

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Contributors

Languages

0