Как использовать фреймворк Perseus для решения задач. Open source.. Open source. Машинное обучение.. Open source. Машинное обучение. персонализация.. Open source. Машинное обучение. персонализация. рекомендательные системы.. Open source. Машинное обучение. персонализация. рекомендательные системы. трансформеры.. Open source. Машинное обучение. персонализация. рекомендательные системы. трансформеры. туториал.. Open source. Машинное обучение. персонализация. рекомендательные системы. трансформеры. туториал. фреймворк.
Как использовать фреймворк Perseus для решения задач - 1

Привет! Я Аня Никифорова, ML-разработчик по направлению рекомендательных систем в Т-Банке. Этим летом на Turbo ML Conf 2026 мы представили фреймворк Perseus. Он подходит для работы с гетерогенными последовательностями событий пользователей, хорошо масштабируется под различные задачи и из коробки поддерживает кросс-доменные сценарии. 

В силу своей гибкости фреймворк может показаться сложным, поэтому мы решили поделиться подробным туториалом, где по шагам рассказываем, как использовать Perseus для своих задач. Посмотрим, как с помощью Perseus реализовать кандидатогенерацию, ранжирование, классификацию и регрессию поверх истории действий пользователей в сервисах: как подготовить данные, какие команды вызвать и как оценить результат. На примере кандидатогенерации подробно рассмотрим, как добавлять в модель дополнительные фичи, обогащать ее событиями и как настраивать модель, меняя лишь пару строк в конфиге.

План туториала

Код туториала опубликован, — можно изучать и пользоваться.

Устройство Perseus

Устройство Perseus

Будем работать с датасетом T-ECD, он был опубликован нами в сентябре 2025 и недавно представлен на конференции KDD-2026 (A*) в Южной Корее. T-ECD основан на данных сервисов, где Perseus уже доказал свою эффективность в продакшене. Речь идет о повышенном кэшбеке, Шопинге и Супермаркетах, которые можно найти в разделе «Город» мобильного приложения Т-Банка.

Раздел «Город» в мобильном приложении Т-Банка

Раздел «Город» в мобильном приложении Т-Банка

Датасет хорошо иллюстрирует, что такое экосистема: у нас есть разнообразные сервисы, в которых клиенты совершают действия, и некоторые клиенты пользуются сразу несколькими сервисами. Более того, события, различные по своей природе, могут содержать признаки, указывающие на одну сущность. Например, бренд товара фигурирует и при покупке в магазине, и при заказе на сайте. Всего в датасете представлено пять доменов, и Perseus позволяет легко использовать все многообразие экосистемных данных для улучшения качества на целевой задаче в конкретном домене.

Количество интеракций в различных доменах T-ECD

Количество интеракций в различных доменах T-ECD

План туториала:

  1. Показать, как Perseus выглядит с точки зрения ML-разработчика и что нужно знать, чтобы работать с фреймворком.

  2. Подготовить данные для дальнейших экспериментов.

  3. Посмотреть, как решать задачу кандидатогенерации.

  4. Построить базовую модель исключительно на последовательности item_id.

  5. Обогатить модель дополнительными фичами, событиями и настроить архитектурные компоненты модели (поменять тип бэкбона и пулинга).

  6. Построить пайплайн ранжирования с помощью Perseus.

  7. Собрать пайплайн классификации.

  8. Обучить регрессию.

В каждой из задач обязательно сравнимся с бейзлайнами.

Что важно знать про работу с Perseus

Зона ответственности ML-инженера при работе с фреймворком ограничивается четырьмя шагами:

  1. Загрузкой списка событий, которые модель сможет использовать для обучения, в хранилище (Event Hub).

  2. Подготовкой представления данных (timestamp, client_id, target) в формате, требуемом для конкретной задачи: кандидатогенерации, ранжирования, классификации или регрессии.

  3. Созданием YAML-конфига, задающего параметры модели.

  4. Запуском команд обучения и инференса.

Все остальные операции автоматически выполняются фреймворком.

Мы собрали глоссарий, чтобы все термины воспринимались в нужном контексте. 

Данные:

  • Событие — факт действия клиента в определенный момент (timestamp, client_id). У события всегда есть тип и опционально атрибуты, например id или бренд товара, с которым пользователь провзаимодействовал. События могут быть любыми — покупка товара, прослушивание музыки, обращение в поддержку — и не обязаны иметь одинаковую схему. 

  • Event Hub — единое хранилище событий, из которого Perseus собирает историю клиента. Если нужного события в нем нет, ML-разработчик добавляет его туда сам. Event Hub достаточно собрать единожды, а затем обращаться к нему при различных задачах. Тем не менее Event Hub не является неизменным и его в любой момент можно обогащать новыми событиями.

  • Признак — то, что модель учитывает. Признак может лежать в событии, контексте или артефактах, и в конфиге для каждого признака это указывается явно (located_in).

  • Энкодер превращает признак в эмбеддинг: id для категориальных, ple для числовых, bag-of-words для текстовых. Один энкодер можно переиспользовать для нескольких признаков. Например, бренд из разных доменов попадет в общее пространство.

Устройство базиса:

  • Базис — постановка ML-задачи, которая задается как датасет объектов, на которых модель обучается и инференсится. Для обучения базис нужно разделить на train- и test-фолды, лучше всего по времени. По test-фолду фреймворк отслеживает метрики от эпохи к эпохе.

  • Объект базиса (сэмпл) — пара (client_id, timestamp): клиент в фиксированный момент. Именно для него модель формирует эмбеддинг и делает предсказание.

  • Контекст — признаки уровня объекта базиса, то есть те, что относятся ко всему сэмплу целиком, а не к отдельному событию. Например, соцдем-кластер клиента или флаг, является ли дата праздничным днем. В контексте можно указать только те признаки, которые также будут доступны на инференсе.

  • Таргет — правильный ответ для объекта базиса. Его вид зависит от задачи: список айтемов для кандидатогенерации, список айтемов с релевантностями для ранжирования, метка класса для классификации, число для регрессии.

  • Артефакты — дополнительная информация о таргете, которую неудобно хранить в самих сэмплах. Для кандидатогенерации и ранжирования это таблица айтемов (и их признаков), у классификации и регрессии артефактов нет.

  • Группы — срезы, в которых дополнительно (помимо overall) считаются метрики. Актуальны только для test-фолда.

Архитектура модели:

  • Бэкбон сводит историю событий и контекст в один эмбеддинг клиента. Сначала event_aggregator векторизует каждое событие, context_aggregator — контекст, а затем history_aggregator обрабатывает полученную последовательность и выдает итоговый вектор. history_aggregator — это ядро бэкбона, именно он отвечает за sequence modeling. Доступные варианты: modern_bert (дефолт), bert, ligr, danet, hstu и mamba. Почти все они основаны на трансформерах, не считая mamba.

  • Голова — то, что считается поверх эмбеддинга клиента для получения предсказания. Обучается end-to-end вместе с бэкбоном.

Пайплайн:

  • Конфиг — один YAML-файл, в котором описано все перечисленное: задача и метрики, используемые события и их атрибуты, признаки и энкодеры, бэкбон и гиперпараметры обучения и инференса.

  • Workdir — рабочая директория, с которой работают команды фреймворка. ML-разработчик кладет в нее базис и конфиг, а Perseus складывает туда все, что считает (ее структура — в конце раздела). Важная оговорка, что Event Hub и Workdir — разные сущности, которые не обязаны физически соседствовать.

Perseus для каждого сэмпла базиса из Event Hub собирает предшествующий ему набор событий клиента. Сиквенс событий и контекст пропускаются через энкодеры и бэкбон. Полученное скрытое состояние проходит через голову, которая формирует предсказание, это предсказание сравнивается с таргетом, и считается лосс. Event Hub связывается с базисом через (timestamp, client_id).

Хранилище событий — Event Hub

Event Hub — единое хранилище событий, из которых Perseus собирает сиквенс для клиента. Путь до Event Hub нужно прописать в переменные окружения. 

Каждый тип событий в Event Hub хранится в отдельной папке, а сами события сгруппированы по дням и партиционированы для удобства обращения к ним фреймворка.

Структура Event Hub

Структура Event Hub

Для работы с Event Hub удобно использовать следующие команды:

python -m perseus event-hub add-events events.pq --name transactions-purchase  # добавить события из parquet-файла
python -m perseus event-hub list-events  # показать доступные события (диапазон дат + атрибуты)
python -m perseus event-hub delete-events --name transactions-purchase --date-from 2025-01-01 --date-to 2025-06-30  # удалить события за период

Постановка задачи

Тип решаемой задачи выбирает ML-разработчик. Задача задает вид таргета и артефактов, набор доступных метрик и голову. Несколько примеров, чтобы понять, какую постановку ML-задачи выбрать:

  1. Бизнес хочет рекомендовать клиентам те товары, которые они с наибольшей вероятностью купят в супермаркете. Это классическая задача рекомендаций. Здесь важно не просто предсказать вероятность покупки, а предложить клиенту ограниченный набор наиболее релевантных товаров. Такую задачу можно сформулировать как задачу кандидатогенерации — отбора кандидатов. 

  2. Есть готовый пул товаров, которые нужно упорядочить в ленте так, чтобы на самых верхних позициях оказались те, по которым клиент с наибольшей вероятностью совершит покупку. В отличие от предыдущего примера здесь важен не факт попадания товара в подборку, а именно порядок: чем выше релевантный товар, тем лучше и ошибки на первых позициях критичнее, чем на последних. Это классическая задача ранжирования.

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

  4. Нужно предсказать, сколько клиент потратит в следующем месяце по всем своим счетам. Здесь целевая переменная принимает непрерывные значения, поэтому постановка задачи — регрессия.

Perseus работает с представлением пользователя, поэтому задачи регрессии и классификации должны быть связаны с пользователями. Так, с помощью Perseus нельзя предсказать цену товара, но можно предсказать суммарные траты пользователя в следующем месяце.

Еще одно решение, которое принимает ML-разработчик, — насколько дробным брать timestamp в базисе. От этого зависит и сама постановка, и количество обучающих сэмплов. Можно предсказывать следующее видео, которое лайкнет пользователь (timestamp с точностью до наносекунд), а можно — видео, которые пользователь лайкнет в течение следующего дня (timestamp, округленный до даты). 

Общая рекомендация — группировать базис в соответствии с тем, как часто обновляются рекомендации в сервисе: если раз в сутки, имеет смысл округлять timestamp до даты. То же самое касается задержек при поставке продакшен-данных. Если в сервисе данные приходят с задержкой в час, правильно будет вычесть этот час из timestamp базиса — иначе на обучении модель будет видеть историю, которой в момент предсказания в продакшене еще не окажется.

Запуск пайплайна

Определившись с постановкой задачи, ML-разработчик должен подготовить базис (basis/) и конфиг (config.yaml) в рабочей директории. А дальше все делается командами фреймворка. Обучение:

python -m perseus train prepare-dataset --workdir <workdir>  # собрать датасет из базиса и событий
accelerate launch -m perseus train fit-model --workdir <workdir>  # обучить модель, результат — в checkpoint/

В итоге рабочая директория выглядит так. ML-разработчик готовит только basis/ и config.yaml, все остальное появляется само по мере вызова команд:

workdir/
├── basis/                      # ML-разработчик
│   ├── train/                  #   фолд для обучения
│   │   ├── samples.pq          #     объекты базиса с таргетом
│   │   └── artifacts/items.pq  #     айтемы и их признаки
│   ├── test/                   #   фолд для валидации, структура та же
│   │   ├── samples.pq
│   │   └── artifacts/items.pq
│   ├── samples.pq              #   базис для инференса
│   └── artifacts/items.pq      #   айтемы для инференса
├── config.yaml                 # ML-разработчик
├── dataset/                    # Perseus, prepare-dataset: датасет для обучения
└── checkpoint/                 # Perseus, prepare-dataset и fit-model: препроцессоры, а после завершения обучения — веса модели

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

python -m perseus inference distribute-samples --workdir <workdir>  # разложить базис по партициям
python -m perseus inference make-items-embeddings --workdir <workdir>  # посчитать эмбеддинги айтемов (необязательный шаг)
python -m perseus inference make-backbone-embeddings --workdir <workdir>  # посчитать эмбеддинги клиентов
python -m perseus inference make-head-predictions --workdir <workdir>  # получить предсказания в predictions/

Рабочая директория после инференса выглядит так:

workdir/
├── basis/                      # ML-разработчик
│   ├── samples.pq              #   базис для инференса
│   └── artifacts/items.pq      #   айтемы для инференса
├── checkpoint/                 # Perseus, prepare-dataset и fit-model: препроцессоры и веса модели
├── samples/                    # Perseus, distribute-samples: базис для инференса, разложенный по партициям
├── embeddings/                 # Perseus, make-backbone-embeddings: эмбеддинги клиентов
├── items/                      # Perseus, make-items-embeddings: эмбеддинги айтемов
└── predictions/                # Perseus, make-head-predictions: итоговые предсказания

Артефакты (artifacts/items.pq) нужны только для кандидатогенерации и ранжирования — у классификации и регрессии в базисе лежат только samples.pq. Метрики из конфига Perseus считает сам, но только на этапе обучения и только на test-фолде. Метрики на инференсе, в том числе сравнение с бейзлайном, — уже ответственность ML-разработчика. 

Важно для честного сравнения: предсказания получают только те объекты базиса, по которым есть хотя бы одно событие раньше timestamp сэмпла, иначе сэмпл выпадает и строк в predictions/ оказывается меньше, чем в базисе. Поэтому в туториале мы считаем бейзлайн на полном базисе и им же заполняем пропуски в предсказаниях модели.

Эксперименты на T-ECD

Пройдем по всем четырем типам задач. Порядок действий в каждой из них одинаковый: соберем базис, посчитаем бейзлайн, обучим модель, проинференсим ее и сравним метрики на одном и том же inference-базисе. Все эксперименты мы запускали на одной H100, линейный прогон занимает около 8 часов.

Работать будем с малой версией T-ECD, домен Marketplace: события четырех типов (view, click, like, clickout), справочник товаров с брендом и ценой и справочник пользователей с соцдем-кластером. Добавим данные из доменов Retail и Offers.

Listing
snapshot_download(
    repo_id="t-tech/T-ECD",
    repo_type="dataset",
    allow_patterns="dataset/small/marketplace/",
    local_dir=download_dir
)
snapshot_download(
    repo_id="t-tech/T-ECD",
    repo_type="dataset",
    allow_patterns="dataset/small/users.pq",
    local_dir=download_dir
)

snapshot_download(
    repo_id="t-tech/T-ECD",
    repo_type="dataset",
    allow_patterns="dataset/small/retail/",
    local_dir=download_dir
)

snapshot_download(
    repo_id="t-tech/T-ECD",
    repo_type="dataset",
    allow_patterns="dataset/small/offers/",
    local_dir=download_dir
)

EVENTS = pl.read_parquet(download_dir / "dataset/small/marketplace/events")
EVENTS = EVENTS.select(
    pl.col("action_type").alias("event"),
    pl.col("timestamp").dt.total_microseconds().cast(pl.Datetime("us")).dt.truncate("1s").cast(pl.Datetime("ns")).alias("timestamp"),
    pl.col("user_id").cast(pl.String).alias("client_id"),
    "item_id",
    "subdomain",
).unique()

CLIENTS = pl.read_parquet(download_dir / "dataset/small/users.pq", columns=["user_id", "socdem_cluster"])
CLIENTS = CLIENTS.with_columns(pl.col("user_id").cast(pl.String).alias("client_id")).drop("user_id")

ITEMS = pl.read_parquet(download_dir / "dataset/small/marketplace/items.pq", columns=["item_id", "brand_id", "price"])

retail_events = (
    pl.scan_parquet(download_dir / "dataset/small/retail/events")
    .filter(pl.col("action_type").eq("order"))
    .select(
        pl.col("action_type").alias("event"),
        pl.col("timestamp").dt.total_microseconds().cast(pl.Datetime("us")).dt.truncate("1s").cast(pl.Datetime("ns")).alias("timestamp"),
        pl.col("user_id").cast(pl.String).alias("client_id")
    ).unique()
).collect()

offers_events = (
    pl.scan_parquet(download_dir / "dataset/small/offers/events")
    .filter(~pl.col("action_type").eq("view"))
    .select(
        pl.col("action_type").alias("event"),
        pl.col("timestamp").dt.total_microseconds().cast(pl.Datetime("us")).dt.truncate("1s").cast(pl.Datetime("ns")).alias("timestamp"),
        pl.col("user_id").cast(pl.String).alias("client_id"),
        pl.col("item_id")
    ).unique()
    .join(
        pl.scan_parquet(download_dir / "dataset/small/offers/items.pq").select("item_id", "brand_id"), 
        on="item_id", 
        how="left"
    )
    .drop("item_id")
).collect()

Создание Event Hub

Сначала подготовим события. В нашем случае Event Hub — локальная директория, путь до которой указывается в env-файле. Каждое событие обязано содержать timestamp (тип ns) и client_id (строка), а все остальное — необязательные атрибуты, которые дальше можно будет использовать как признаки: для Marketplace это item_id и subdomain (рекомендательная поверхность, где было совершено событие). 

Каждый тип события загружается отдельной командой под своим именем — так в Event Hub появятся marketplace-view, marketplace-click, marketplace-like и marketplace-clickout. 

Listing
env = {**os.environ, "INTERNAL_STORAGE_EVENT_HUB": "/event-hub/"}

for (event,), group in EVENTS.group_by("event"):
    event_name = f"marketplace-{event}"
    group = group.drop("event")

    with tempfile.TemporaryDirectory() as tmp:
        staging = Path(tmp) / "events.pq"
        group.write_parquet(staging)
        subprocess.run(
            [
                "uv", "run", "python", "-m", "perseus", 
                "event-hub", "add-events", 
                str(staging), 
                "--name", event_name, 
                "--source", "event_hub"
            ],
            env=env,
            check=True
        )
Как выглядят события внутри Event Hub 

Как выглядят события внутри Event Hub 

Аналогично добавим события Offers (offers-click, offers-clickout, offers-like) и Retail (retail-order). 

Listing
for (event,), group in retail_events.group_by("event"):
    event_name = f"retail-{event}"
    group = group.drop("event")

    with tempfile.TemporaryDirectory() as tmp:
        staging = Path(tmp) / "events.pq"
        group.write_parquet(staging)
        subprocess.run(
            [
                "uv", "run", "python", "-m", "perseus", 
                "event-hub", "add-events", 
                str(staging), 
                "--name", event_name, 
                "--source", "event_hub"
            ],
            env=env,
            check=True
        )

for (event,), group in offers_events.group_by("event"):
    event_name = f"offers-{event}"
    group = group.drop("event")

    with tempfile.TemporaryDirectory() as tmp:
        staging = Path(tmp) / "events.pq"
        group.write_parquet(staging)
        subprocess.run(
            [
                "uv", "run", "python", "-m", "perseus", 
                "event-hub", "add-events", 
                str(staging), 
                "--name", event_name, 
                "--source", "event_hub"
            ],
            env=env,
            check=True
        )

Соцдем-кластер клиента мы положим в контекст базиса, а признаки товаров — в артефакты. Существует возможность также приджойнить эти признаки к событиям Event Hub, чтобы они обрабатывались энкодером на уровне события.

Кандидатогенерация

Будем предсказывать, на какие товары пользователь наиболее вероятно кликнет в Marketplace на следующий день. В качестве метрик возьмем Recall@100, NDCG@100 и Coverage@100.

Базису для кандидатогенерации нужны артефакты — таблица items.pq с обязательной колонкой item. Остальные ее колонки можно использовать как признаки айтема. В train/samples.pq таргет — список айтемов, с которыми клиент провзаимодействовал полезным для бизнеса образом. Негативы фреймворк сгенерирует сам во время обучения. В test/samples.pq у каждого айтема в таргете дополнительно указывается релевантность (везде 1, так как товары не различаются по уровню релевантности) — она используется при расчете метрики NDCG. В артефактах каждого фолда должны быть все айтемы, встречающиеся в его таргете.

Соберем базис из событий кликов за последние 120 дней. timestamp округлим до даты и сдвинем на −12 часов (ограничение T-ECD). Таргет сэмпла — все товары, на которые клиент кликнул в этот день. Разделим базис на train и test в отношении 80/20 по времени. В тестовой части проставим всем айтемам релевантность 1, так как товары не отличаются друг от друга по степени полезности. В контекст положим соцдем-кластер клиента, в артефакты — айтемы из train-части вместе с брендом. Тестовые артефакты продублируют тренировочные: модель завязана на id товара и не сможет рекомендовать то, что не видела при обучении.

Listing
samples = pl.read_parquet("/event-hub/marketplace-click")
start_date = samples["date"].max() - timedelta(days=120)
samples = (
    samples
    .filter(pl.col("timestamp") >= start_date)
    .with_columns((pl.col("date").cast(pl.Datetime) - timedelta(hours=12)).cast(pl.Datetime("ns")).alias("timestamp"))
    .drop("date")
    .group_by("timestamp", "client_id").agg(pl.col("item_id").unique().alias("target"))
    .sort("timestamp")
)

train_ratio = 0.8
split_idx = int(len(samples) * train_ratio)
split_timestamp = samples["timestamp"][split_idx]
train_samples = samples.filter(pl.col("timestamp") < split_timestamp)
test_samples = samples.filter(pl.col("timestamp") >= split_timestamp)
test_samples = (
    test_samples
    .with_columns(
        pl.col("target").list.eval(
            pl.struct([
                pl.element().alias("item"),
                pl.lit(1).alias("relevance")
            ])
        ).alias("target")
    )
)

train_samples = train_samples.join(CLIENTS, on="client_id", how="left")
train_samples = train_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
test_samples = test_samples.join(CLIENTS, on="client_id", how="left")
test_samples = test_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
train_samples.write_parquet("../retrieval/basis/train/samples.pq")
test_samples.write_parquet("../retrieval/basis/test/samples.pq")

artifacts = pl.DataFrame(train_samples["target"].explode().unique())
artifacts = artifacts.join(ITEMS, left_on="target", right_on="item_id", how="left").rename({"target": "item"}).select("item", "brand_id")
artifacts.write_parquet("../retrieval/basis/train/artifacts/items.pq")
artifacts.write_parquet("../retrieval/basis/test/artifacts/items.pq")

Получается такой train/samples.pq:

Как использовать фреймворк Perseus для решения задач - 7

В test/samples.pq к каждому айтему таргета добавляется релевантность:

Как использовать фреймворк Perseus для решения задач - 8

А artifacts/items.pq, одинаковый для обоих фолдов, выглядит так: 

Как использовать фреймворк Perseus для решения задач - 9

Базис для инференса — копия тестовой части вместе с артефактами. Подготовим его один раз, до бейзлайна и первого обучения. Базис при этом не меняется, поэтому метрики всех вариантов модели останутся сравнимыми между собой и с бейзлайном.

Listing
inference_basis = pl.read_parquet("../retrieval/basis/test/samples.pq")
inference_basis.write_parquet("../retrieval/basis/samples.pq")
inference_artifacts = pl.read_parquet("../retrieval/basis/test/artifacts/items.pq")
inference_artifacts.write_parquet("../retrieval/basis/artifacts/items.pq")
uv run python -m perseus inference distribute-samples --workdir ../retrieval

В качестве бейзлайна возьмем топ-100 самых популярных айтемов из train-части базиса и порекомендуем их всем пользователям. Скор айтема — его позиция в топе, так что порядок внутри рекомендаций тоже определен.

Listing
toppop_items = (
    train_samples.explode("target")["target"]
    .value_counts().sort("count", descending=True).head(100)["target"].to_list()
)
toppop_prediction = [
    {"item": item, "score": float(len(toppop_items) - rank)}
    for rank, item in enumerate(toppop_items)
]

Посчитаем бейзлайн сразу, до обучения модели: на инференсе сэмплы без истории отбрасываются, поэтому предсказания модели мы потом приджойним к полному базису и заполним пропуски бейзлайном. Так обе оценки окажутся на одном и том же наборе сэмплов. Дальше во всех задачах будем поступать точно так же:

Вариант

Recall@100

NDCG@100

Coverage@100

Топ популярных (бейзлайн)

0,1275

0,0430

0,0005

Первая модель будет использовать единственный признак — последовательность item_id. В конфиге опишем задачу и метрики, целевое событие, энкодер айтема, бэкбон (ModernBERT на 4 слоя, сумму как агрегатор событий) и параметры обучения и инференса. 

config1.yaml
task:
  type: retrieval
  metrics:
    recall@100:
      type: recall_at_k
      params:
        k: 100
    ndcg@100:
      type: ndcg_at_k
      params:
        k: 100
    coverage@100:
      type: coverage_at_k
      params:
        k: 100

events:
  marketplace-click:
    attributes:
      item_id:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  item:
    located_in:
      artifacts: true
    encoder: item

backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: recall@100
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  predict_kwargs:
    k: 1000

Дальше достаточно двух команд: собрать датасет и обучить модель: 

!uv run python -m perseus train prepare-dataset --workdir ../retrieval
!uv run accelerate launch -m perseus train fit-model --workdir ../retrieval

Весь процесс обучения будет автоматически документироваться в виде текстовых логов. 

Логи в процессе обучения Perseus

Логи в процессе обучения Perseus

При настроенном ClearML увидим следующую картину.

Прогресс обучения в ClearML

Прогресс обучения в ClearML

После обучения проинференсим модель и посчитаем по ее предсказаниям те же метрики, что и для бейзлайна:

uv run python -m perseus inference make-backbone-embeddings --workdir ../retrieval
uv run python -m perseus inference make-head-predictions --workdir ../retrieval

Вариант

Recall@100

NDCG@100

Coverage@100

Топ популярных (бейзлайн)

0,1275

0,0430

0,0005

Perseus, только item_id

0,1518

0,0545

0,0056

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

uv run python -m perseus inference make-items-embeddings --workdir ../retrieval

В Marketplace представлен не только item_id. Добавим в модель подраздел сервиса (subdomain) из событий, бренд товара из артефактов и соцдем-кластер клиента из контекста. Новой подготовки данных не потребуется: все это мы сохранили еще на этапе сбора Event Hub и базиса, поэтому достаточно дописать признаки в конфиг, указав для каждого, где он расположен и каким энкодером кодируется. 

config2.yaml
task:
  type: retrieval
  metrics:
    recall@100:
      type: recall_at_k
      params:
        k: 100
    ndcg@100:
      type: ndcg_at_k
      params:
        k: 100
    coverage@100:
      type: coverage_at_k
      params:
        k: 100

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  item:
    located_in:
      artifacts: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      artifacts: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster
  
backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: recall@100
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003
      
inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  predict_kwargs:
    k: 1000

Датасет пересоберем, модель обучим и проинференсим заново:

uv run python -m perseus inference make-backbone-embeddings --workdir ../retrieval
uv run python -m perseus inference make-head-predictions --workdir ../retrieval
uv run python -m perseus inference make-items-embeddings --workdir ../retrieval

Вариант

Recall@100

NDCG@100

Coverage@100

Perseus, только item_id

0,1518

0,0545

0,0056

Perseus + признаки

0,1553

0,0572

0,0099

Добавление признаков позволило немного улучшить качество модели. Теперь обогатим модель событиями. Помимо кликов добавим лайки и кликауты Marketplace, а также события соседних доменов: заказы в Retail и клики, лайки и кликауты в Offers. Дополним конфиг.

Обратим внимание на две вещи. У событий Offers нет item_id, зато есть бренд — тот же признак, что и в артефактах Marketplace, поэтому кодировать его будем общим энкодером и информация из другого домена попадет в то же пространство. А еще событий стало значительно больше, а длина истории ограничена, поэтому зададим целевому событию более высокий приоритет. Иначе клики по товарам вытеснятся из последовательности остальными событиями.

config3.yaml
task:
  type: retrieval
  metrics:
    recall@100:
      type: recall_at_k
      params:
        k: 100
    ndcg@100:
      type: ndcg_at_k
      params:
        k: 100
    coverage@100:
      type: coverage_at_k
      params:
        k: 100

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 2
  marketplace-like:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  marketplace-clickout:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  retail-order:
    max_duration_per_sequence: 365d
  offers-click:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-like:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-clickout:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  item:
    located_in:
      artifacts: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      event: true
      artifacts: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster

backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: recall@100
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  predict_kwargs:
    k: 1000

Снова соберем датасет, обучим и проинференсим модель с помощью уже знакомых команд:

uv run python -m perseus inference make-backbone-embeddings --workdir ../retrieval
uv run python -m perseus inference make-head-predictions --workdir ../retrieval
uv run python -m perseus inference make-items-embeddings --workdir ../retrieval

Вариант

Recall@100

NDCG@100

Coverage@100

Perseus + признаки

0,1553

0,0572

0,0099

Perseus + события

0,1638

0,0634

0,0046

Видим приросты в качестве относительно предыдущей версии модели. В этом и есть основная сила Perseus: он позволяет учитывать в пользовательской истории события из разных доменов с разными схемами. Более того, из обширного Event Hub можно подключать только нужный набор событий, тем самым обучая модели на разных срезах данных без необходимости их перезаписи. 

Наконец, изменим архитектуру: заменим ModernBERT на HSTU (он учитывает не только порядок событий, но и время между ними), а сумму в агрегаторе событий — на взвешенную сумму. Для этого достаточно поменять пару строк в конфиге. 

config4.yaml
task:
  type: retrieval
  metrics:
    recall@100:
      type: recall_at_k
      params:
        k: 100
    ndcg@100:
      type: ndcg_at_k
      params:
        k: 100
    coverage@100:
      type: coverage_at_k
      params:
        k: 100

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 2
  marketplace-like:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  marketplace-clickout:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  retail-order:
    max_duration_per_sequence: 365d
  offers-click:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-like:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-clickout:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  item:
    located_in:
      artifacts: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      event: true
      artifacts: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster

backbone:
  dim: 256
  history_aggregator:
    type: hstu
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: weighted_sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: recall@100
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  predict_kwargs:
    k: 1000

Данные при этом не изменились, поэтому пересобирать датасет не нужно — сразу запускаем обучение:

cp ../retrieval/config.yaml ../retrieval/checkpoint/config.yaml
uv run accelerate launch -m perseus train fit-model --workdir ../retrieval
uv run python -m perseus inference make-backbone-embeddings --workdir ../retrieval
uv run python -m perseus inference make-head-predictions --workdir ../retrieval

Видим заметный прирост метрик. 

Вариант

Recall@100

NDCG@100

Coverage@100

Perseus + события

0,1638

0,0634

0,0046

Perseus + архитектура

0,1812

0,0705

0,0542

Нужно отметить, что под разные задачи могут подходить разные комбинации гиперпараметров, поэтому необходимо экспериментировать.

Соберем все метрики в одну таблицу.

Вариант

Recall@100

NDCG@100

Coverage@100

Топ популярных (бейзлайн)

0,1275

0,0430

0,0005

Perseus, только item_id

0,1518

0,0545

0,0056

Perseus + признаки

0,1553

0,0572

0,0099

Perseus + события

0,1638

0,0634

0,0046

Perseus + архитектура

0,1812

0,0705

0,0542

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

Ранжирование

Следующая задача — переупорядочить готовый пул кандидатов. Будем считать, что с точки зрения бизнеса события ранжируются как clickout > like > click > view. В качестве метрик возьмем NDCG@20 и MRR@20.

В Perseus ранжирование реализовано как предсказание вероятностей целевых событий (multi-label classification). По взвешенной сумме этих вероятностей затем можно проранжировать объекты. Веса задаются априорно, а не выучиваются моделью, что позволяет ML-разработчику приоритизировать то или иное событие в зависимости от целей бизнеса.

Базис. Ранжированию, как и кандидатогенерации, нужны артефакты с айтемами, а таргет устроен сложнее. Для каждого айтема в нем хранится не только сам айтем, но и набор булевых меток — какими именно целевыми действиями клиент с ним провзаимодействовал. В тестовой части к меткам добавляется релевантность, а в самих сэмплах появляется обязательная колонка items — тот самый пул кандидатов, который модель будет скорить и переупорядочивать.

Соберем базис из всех четырех типов событий Marketplace за те же 120 дней и с тем же округлением timestamp до даты. Метками будут четыре флага по типам событий, а релевантностью — 0 для просмотра, 1 для клика, 2 для лайка и 3 для кликаута. Оставим только те сэмплы, в которых встречается больше одного уровня релевантности: если все айтемы одинаково хороши, упорядочивать нечего и метрика по такому сэмплу неинформативна. Далее так же, как и в кандидатогенерации: сплит 80/20 по времени, соцдем-кластер в контекст, айтемы с брендом в артефакты.

Listing
samples = pl.concat([
    pl.read_parquet("/event-hub/marketplace-clickout").with_columns(pl.lit("clickout").alias("event_type")),
    pl.read_parquet("/event-hub/marketplace-click").with_columns(pl.lit("click").alias("event_type")),
    pl.read_parquet("/event-hub/marketplace-like").with_columns(pl.lit("like").alias("event_type")),
    pl.read_parquet("/event-hub/marketplace-view").with_columns(pl.lit("view").alias("event_type")),
])

start_date = samples["date"].max() - timedelta(days=120)
samples = (
    samples
    .filter(pl.col("timestamp") >= start_date)
    .with_columns((pl.col("date").cast(pl.Datetime) - timedelta(hours=12)).cast(pl.Datetime("ns")).alias("timestamp")).drop("date")
)
samples = samples.with_columns(
    pl.struct([
        pl.col("item_id").alias("item"),
        pl.struct([
            (pl.col("event_type") == "view").alias("view"),
            (pl.col("event_type") == "like").alias("like"),
            (pl.col("event_type") == "click").alias("click"),
            (pl.col("event_type") == "clickout").alias("clickout"),
        ]).alias("labels"),
    ]).alias("target")
).drop("item_id", "event_type")
samples = (
    samples
    .group_by("timestamp", "client_id")
    .agg(pl.col("target").unique())
    .filter(
        pl.col("target").list.eval(
            pl.when(pl.element().struct.field("labels").struct.field("view")).then(0)
            .when(pl.element().struct.field("labels").struct.field("click")).then(1)
            .when(pl.element().struct.field("labels").struct.field("like")).then(2)
            .when(pl.element().struct.field("labels").struct.field("clickout")).then(3)
            .otherwise(-1)
        ).list.unique().list.len() > 1
    )
    .sort("timestamp")
)

train_ratio = 0.8
split_idx = int(len(samples) * train_ratio)
split_timestamp = samples["timestamp"][split_idx]

train_samples = samples.filter(pl.col("timestamp") < split_timestamp)
test_samples = samples.filter(pl.col("timestamp") >= split_timestamp)
test_samples = (
    test_samples
    .explode("target")
    .with_columns([
        pl.col("target").struct.field("item").alias("item"),
        pl.col("target").struct.field("labels").struct.field("view").alias("view"),
        pl.col("target").struct.field("labels").struct.field("like").alias("like"),
        pl.col("target").struct.field("labels").struct.field("click").alias("click"),
        pl.col("target").struct.field("labels").struct.field("clickout").alias("clickout"),
    ])
    .with_columns(
        pl.when(pl.col("view")).then(0)
         .when(pl.col("click")).then(1)
         .when(pl.col("like")).then(2)
         .when(pl.col("clickout")).then(3)
         .otherwise(-1)
         .alias("relevance")
    )
    .group_by("timestamp", "client_id", "item")
    .agg([
        pl.col("view").any(),
        pl.col("like").any(),
        pl.col("click").any(),
        pl.col("clickout").any(),
        pl.col("relevance").max()
    ])
    .with_columns(
        pl.struct([
            "item",
            pl.struct(["view", "like", "click", "clickout"]).alias("labels"),
            "relevance"
        ]).alias("target")
    )
    .group_by("timestamp", "client_id")
    .agg([
        pl.col("target"),
        pl.col("item").alias("items")
    ])
)

train_samples = train_samples.join(CLIENTS, on="client_id", how="left")
train_samples = train_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
test_samples = test_samples.join(CLIENTS, on="client_id", how="left")
test_samples = test_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
train_samples.write_parquet("../ranking/basis/train/samples.pq")
test_samples.write_parquet("../ranking/basis/test/samples.pq")

artifacts = pl.DataFrame(
    train_samples.select(pl.col("target").list.eval(pl.element().struct.field("item")).explode()).unique()
).rename({"target": "item"})
artifacts = artifacts.with_columns(pl.col("item")).join(ITEMS, left_on="item", right_on="item_id", how="left").select("item", "brand_id")
artifacts.write_parquet("../ranking/basis/train/artifacts/items.pq")
artifacts.write_parquet("../ranking/basis/test/artifacts/items.pq")

train/samples.pq выглядит так:

Как использовать фреймворк Perseus для решения задач - 12

В test/samples.pq к каждому айтему таргета добавляется релевантность, а рядом появляется колонка items — пул кандидатов (labels для краткости свернуты, в файле они такие же, как в train):

Как использовать фреймворк Perseus для решения задач - 13

Артефакты те же, что и в кандидатогенерации, — item и бренд:

Как использовать фреймворк Perseus для решения задач - 14

Пулом кандидатов на инференсе будет колонка items из тестовой части — собственный набор айтемов каждого сэмпла, тот же, на котором модель валидировалась в процессе обучения. В продакшене пул приходил бы от кандидатогенератора, например от retrieval-модели из предыдущего раздела. В туториале мы оцениваем качество переупорядочивания в чистом виде, поэтому в пул кладем айтемы, которые пользователь в этот день видел.

В качестве бейзлайна отранжируем пул по популярности айтема в train-части базиса, считая ее по позитивным событиям: клику, лайку и кликауту.

Listing
item_to_popularity = dict(
    train_samples.select(pl.col("target").explode()).unnest("target").unnest("labels")
    .filter(pl.col("click") | pl.col("like") | pl.col("clickout"))
    ["item"].value_counts().iter_rows()
)

inference_basis = pl.read_parquet("../ranking/basis/samples.pq")
inference_artifacts = pl.read_parquet("../ranking/basis/artifacts/items.pq")

inference_basis = inference_basis.with_columns(
    baseline_prediction=pl.col("items").list.eval(
        pl.struct(
            item=pl.element(),
            probas=pl.struct(**{label: pl.lit(0.0, pl.Float32) for label in ["view", "click", "like", "clickout"]}),
            score=pl.element().replace_strict(item_to_popularity, default=0).cast(pl.Float32),
        )
    )
)

Модель предсказывает вероятность каждого типа события для каждого кандидата, а пул упорядочивается по одному скору — взвешенной сумме этих вероятностей. По умолчанию веса равны, то есть просмотр вносит в скор такой же вклад, что и кликаут. Через label_to_weight зададим веса равными релевантностям: тогда скор — это ожидаемая релевантность и просмотры ее не завышают. 

config5.yaml
task:
  type: ranking
  head:
    label_to_weight:
      view: 0
      click: 1
      like: 2
      clickout: 3
  metrics:
    ndcg@20:
      type: ndcg_at_k
      params:
        k: 20
    mrr@20:
      type: mrr_at_k
      params:
        k: 20

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  marketplace-like:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 2
  marketplace-clickout:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 3
  marketplace-view:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 0

max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  item:
    located_in:
      artifacts: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      artifacts: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster

backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: ndcg@20
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true

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

uv run python -m perseus inference distribute-samples --workdir ../ranking
uv run python -m perseus inference make-backbone-embeddings --workdir ../ranking
uv run python -m perseus inference make-head-predictions --workdir ../ranking

Модель возвращает вероятности для каждого типа событий и их взвешенную сумму.

Предсказания Perseus

Предсказания Perseus

Вариант

NDCG@20

MRR@20

По популярности (бейзлайн)

0,4812

0,4357

Perseus

0,4984

0,4578

Perseus выигрывает у бейзлана, хотя и с меньшим отрывом, чем было в случае с кандидатогенерацией.

Классификация

Будем предсказывать, совершит ли пользователь хотя бы одно активное действие (клик, лайк или кликаут) в Marketplace в течение 7 дней после даты сэмпла. В качестве метрики возьмем ROC-AUC.

Базис для классификации устроен максимально просто: таргет — строка с названием класса, дополнительных колонок и артефактов не требуется. Соберем базис из дней, в которые пользователь был активен в Marketplace, за те же последние 120 дней и с тем же округлением timestamp. 

Таргет посчитаем по окну (t, t + 7 дней], то есть строго в будущем относительно сэмпла: visit, если активность в окне была, и no_visit иначе. Последние 7 дней выборки отбросим: для них окно неполное и таргет оказался бы занижен. Затем, как и раньше, разделим базис 80/20 по времени и положим соцдем-кластер в контекст. 

Listing
samples = pl.concat([
    pl.read_parquet("/event-hub/marketplace-click"),
    pl.read_parquet("/event-hub/marketplace-like"),
    pl.read_parquet("/event-hub/marketplace-clickout"),
]).select("date", "timestamp", "client_id")

start_date = samples["date"].max() - timedelta(days=120)
samples = (
    samples
    .filter(pl.col("timestamp") >= start_date)
    .with_columns((pl.col("date").cast(pl.Datetime) - timedelta(hours=12)).cast(pl.Datetime("ns")).alias("timestamp")).drop("date")
    .unique(["timestamp", "client_id"])
    .sort("client_id", "timestamp")
)

horizon = timedelta(days=7)
next_week_visits = samples.rolling(
    index_column="timestamp",
    period="7d",
    offset="0d",
    closed="right",
    group_by="client_id",
).agg(pl.len().alias("num_visits"))
samples = (
    samples
    .filter(pl.col("timestamp") < pl.col("timestamp").max() - horizon)
    .join(next_week_visits, on=["client_id", "timestamp"], how="left")
    .select(
        "timestamp",
        "client_id",
        pl.when(pl.col("num_visits") > 0).then(pl.lit("visit")).otherwise(pl.lit("no_visit")).alias("target"),
    )
    .sort("timestamp")
)

train_ratio = 0.8
split_idx = int(len(samples) * train_ratio)
split_timestamp = samples["timestamp"][split_idx]
train_samples = samples.filter(pl.col("timestamp") < split_timestamp)
test_samples = samples.filter(pl.col("timestamp") >= split_timestamp)

train_samples = train_samples.join(CLIENTS, on="client_id", how="left")
train_samples = train_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
test_samples = test_samples.join(CLIENTS, on="client_id", how="left")
test_samples = test_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
train_samples.write_parquet("../classification/basis/train/samples.pq")
test_samples.write_parquet("../classification/basis/test/samples.pq")

train/samples.pq и test/samples.pq устроены одинаково:

Как использовать фреймворк Perseus для решения задач - 16

В качестве бейзлайна воспользуемся следующим правилом: будем предсказывать позитивную метку, если у клиента было хотя бы одно положительное взаимодействие за 30-дневный период, предшествующий timestamp-у сэмпла.

Listing
visit_rate = (train_samples["target"] == "visit").mean()
baseline_prediction = pl.struct(
  no_visit=pl.lit(1 - visit_rate, pl.Float32),
  visit=pl.lit(visit_rate, pl.Float32),
)

inference_basis = pl.read_parquet("../classification/basis/samples.pq")

activity = pl.concat([train_samples, test_samples]).select("client_id", "timestamp").sort("client_id", "timestamp")

visited_before = activity.join(
  activity.rolling(
      index_column="timestamp",
      period="7d",
      offset="-30d",
      closed="left",
      group_by="client_id",
  ).agg(pl.len().alias("num_prior_visits")),
  on=["client_id", "timestamp"],
  how="left",
).with_columns(
  (pl.col("num_prior_visits").fill_null(0) > 0).cast(pl.Float32).alias("visit")
).select("client_id", "timestamp", "visit")

Конфиг получится проще, чем в предыдущих задачах: артефакты не нужны, модель по эмбеддингу пользователя сразу предсказывает распределение по классам. Для ROC-AUC в бинарном случае необходимо указать pos_label. 

config6.yaml
task:
  type: classification
  metrics:
    roc_auc:
      params:
        pos_label: visit

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 2
  marketplace-like:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  marketplace-clickout:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  retail-order:
    max_duration_per_sequence: 365d
  offers-click:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-like:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-clickout:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      event: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster

backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: roc_auc
    mode: max
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true

Команды обучения и инференса остаются такими же, как при кандидатогенерации и ранжировании:

uv run python -m perseus train prepare-dataset --workdir ../classification
uv run accelerate launch -m perseus train fit-model --workdir ../classification
uv run python -m perseus inference distribute-samples --workdir ../classification
uv run python -m perseus inference make-backbone-embeddings --workdir ../classification
uv run python -m perseus inference make-head-predictions --workdir ../classification

Вариант

ROC-AUC

Правило (бейзлайн)

0,52

Perseus

0,59

Регрессия

Последняя задача — регрессия. Будем предсказывать суммарную стоимость товаров, с которыми пользователь позитивно провзаимодействует в Marketplace в течение следующего месяца. Просмотры в нее не входят, так как просмотр — это не позитивное взаимодействие. В качестве метрик возьмем MAE и RMSE.

Соберем базис аналогично тому, как делали для классификации. Но теперь нам нужны сами товары и их стоимость, поэтому дедуплицировать события до пар (дата, пользователь) будем только после джойна со справочником товаров. Свернем события в стоимость корзины за день и просуммируем ее в окне (t, t + 30 дней]. Последний месяц выборки отбросим как неполный. 

Listing
events = pl.concat([
    pl.read_parquet("/event-hub/marketplace-click"),
    pl.read_parquet("/event-hub/marketplace-like"),
    pl.read_parquet("/event-hub/marketplace-clickout"),
]).select("date", "timestamp", "client_id", "item_id")

start_date = events["date"].max() - timedelta(days=120)
events = (
    events
    .filter(pl.col("timestamp") >= start_date)
    .with_columns((pl.col("date").cast(pl.Datetime) - timedelta(hours=12)).cast(pl.Datetime("ns")).alias("timestamp")).drop("date")
    .unique(["timestamp", "client_id", "item_id"])
    .join(ITEMS.select("item_id", pl.col("price").cast(pl.Float64)), on="item_id", how="left")
)

horizon = timedelta(days=30)
samples = (
    events
    .group_by("timestamp", "client_id")
    .agg(pl.col("price").sum().alias("daily_spend"))
    .sort("client_id", "timestamp")
)
next_month_spend = samples.rolling(
    index_column="timestamp",
    period="30d",
    offset="0d",
    closed="right",
    group_by="client_id",
).agg(pl.col("daily_spend").sum().alias("target"))
samples = (
    samples
    .filter(pl.col("timestamp") < pl.col("timestamp").max() - horizon)
    .join(next_month_spend, on=["client_id", "timestamp"], how="left")
    .select("timestamp", "client_id", pl.col("target").fill_null(0.0))
    .sort("timestamp")
)

train_ratio = 0.8
split_idx = int(len(samples) * train_ratio)
split_timestamp = samples["timestamp"][split_idx]
train_samples = samples.filter(pl.col("timestamp") < split_timestamp)
test_samples = samples.filter(pl.col("timestamp") >= split_timestamp)

train_samples = train_samples.join(CLIENTS, on="client_id", how="left")
train_samples = train_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
test_samples = test_samples.join(CLIENTS, on="client_id", how="left")
test_samples = test_samples.with_columns(
    pl.struct(["socdem_cluster"]).alias("context")
).drop(["socdem_cluster"])
train_samples.write_parquet("../regression/basis/train/samples.pq")
test_samples.write_parquet("../regression/basis/test/samples.pq")

Базис получается такой же формы, что и в классификации, только в target-число:

Как использовать фреймворк Perseus для решения задач - 17

В качестве бейзлайна возьмем константное предсказание — среднее по train-части базиса.

Listing
mean_target = train_samples["target"].mean()
baseline_prediction = pl.lit(mean_target, pl.Float32)

Конфиг почти повторяет конфиг классификации: меняются тип задачи, метрики и голова. Таргет получается скошенным, поэтому отнормируем его для обучения — применим MinMax Scaling, указав соответствующую строчку в конфиге.

config7.yaml
task:
  type: regression
  metrics:
    mae:
    rmse:
  preprocessor:
    scaler: minmax

events:
  marketplace-click:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 2
  marketplace-like:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  marketplace-clickout:
    attributes:
      item_id:
      subdomain:
    max_duration_per_sequence: 365d
    priority: 1
  retail-order:
    max_duration_per_sequence: 365d
  offers-click:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-like:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
  offers-clickout:
    attributes:
      brand_id:
    max_duration_per_sequence: 365d
max_events_per_sequence: 512

encoders:
  item:
    type: id
  subdomain:
    type: id
  brand_id:
    type: id
  socdem_cluster:
    type: id

features:
  item_id:
    located_in:
      event: true
    encoder: item
  subdomain:
    located_in:
      event: true
    encoder: subdomain
  brand_id:
    located_in:
      event: true
    encoder: brand_id
  socdem_cluster:
    located_in:
      context: true
    encoder: socdem_cluster

backbone:
  dim: 256
  history_aggregator:
    type: modern_bert
    params:
      num_layers: 4
      num_heads: 4
      dropout: 0.1
  event_aggregator:
    type: sum
  context_aggregator:
    type: identity

training:
  num_epochs: 10
  log_every_n_train_steps: 1000
  dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
    shuffle: true
  test_dataloader:
    batch_size: 32
    num_workers: 2
    pin_memory: true
  early_stopping:
    metric: rmse
    mode: min
    patience: 3
  optimizer:
    params:
      lr: 0.0003

inference:
  backbone_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true
  head_dataloader:
    batch_size: 64
    num_workers: 0
    pin_memory: true

Команды те же, что и в предыдущих задачах:

uv run python -m perseus train prepare-dataset --workdir ../regression
uv run accelerate launch -m perseus train fit-model --workdir ../regression
uv run python -m perseus inference distribute-samples --workdir ../regression
uv run python -m perseus inference make-backbone-embeddings --workdir ../regression
uv run python -m perseus inference make-head-predictions --workdir ../regression

Вариант

MAE

RMSE

Среднее (бейзлайн)

7,7591

13,9357

Perseus

7,1497

13,5175

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

Заключение

Мы рассмотрели, как работать с фреймворком Perseus. На данных датасета T-ECD, хорошо отражающих сложность и многогранность реальной системы, мы разобрали четыре сценария: кандидатогенерацию, ранжирование, классификацию и регрессию. В каждом случае строили модель, сравнивали с бейзлайном и смотрели, как меняется качество.

Туториал иллюстрирует гибкость фреймворка. Добавить фичи? Подключить события из соседнего домена? Попробовать другой бэкбон или тип пулинга? Достаточно поменять пару строк в YAML-конфиге — никакого переписывания кода с нуля.

Мы постарались показать Perseus с практической стороны, без лишней теории. Конечно, чтобы освоиться, потребуется разобраться в форматах данных и структуре конфигов, но, надеюсь, наш туториал станет хорошей точкой входа. Мы уверены, что Perseus стоит того, чтобы потратить на него время, особенно если вы работаете с мультидоменными данными.

Ждем ваших впечатлений, комментариев и вопросов!

Полезные ссылки: 

Автор: sonyaleaf

Источник