TGStat
TGStat
Введите текст для поиска
Расширенный поиск каналов
  • flag Russian
    Язык сайта
    flag Russian flag English flag Uzbek
  • Вход на сайт
  • Каталог
    Каталог каналов и чатов Поиск каналов
    Добавить канал/чат
  • Рейтинги
    Рейтинг каналов Рейтинг чатов Рейтинг публикаций
    Рейтинги брендов и персон
  • Аналитика
  • Поиск по публикациям
  • Мониторинг Telegram
Егорка думает, что

26 May, 13:16

Открыть в Telegram Поделиться Пожаловаться

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

В вычислительной схеме трансформер куча мест, где мы сравниваем циферки сильно разного масштаба, в матрицах весов, в атеншене, в РМСнорме. В числах с плавающей точкой ограниченный размер мантиссы, в bf16 у нас 7 бит, в fp16 - 10. Это значит, что число в таком формате может хранить не больше определенного количества значих цифр. Придумаем ситуацию, с 3 корзинами с яблоками. Если у нас есть две корзинки, в которых у нас по 0.01 яблоку, ну то есть по очень маленькому кусочку, у нас теперь два кусочка яблока. В третьей корзинке, что нам дают, лежит 10 тысяч яблок, и теперь у нас 10 тысяч яблок и еще два кусочка. Но если бы мы изначально получили 10 тысяч яблок, эти два кусочка, согласитесь, были бы уже не так важны. 10000 - это 5 значащих цифр только для целой части, если мы прибавим еще 0.01, бит для новой части просто не хватит, они потеряются в округлении, то есть в bf16 10000 + 0.01 ~ 10000.0, при этом если сначала сложить все мелкие части, то они соберутся в что-то достаточно крупное, чтобы выжить рядом с большим числом. Это все называется catastrophic cancellation и это причина, почему в вычислениях с плавающей точкой (a + b) + c != a + (b + c). Стандарт IEEE 754, почитайте.

Собсна, вот вам пример:

import torch

a = torch.tensor(10000.0, dtype=torch.float16)
b = torch.tensor(0.01, dtype=torch.float16)
c = torch.tensor(0.01, dtype=torch.float16)

print((a + b) + c) # tensor(10000., dtype=torch.float16)
print(a + (b + c)) # tensor(10000., dtype=torch.float16)

a = torch.tensor(1000.0, dtype=torch.bfloat16)
b = torch.tensor(0.1, dtype=torch.bfloat16)
c = torch.tensor(0.2, dtype=torch.bfloat16)

print((a + b) + c) # tensor(1000., dtype=torch.bfloat16)
print(b + c + a) # tensor(1000.2500, dtype=torch.bfloat16)


А теперь имаджинируйте, что циферки не три, а 4096, собранных в вектор - как раз размер хиддена в какой-нибудь лламе 3. И нужно нам в рмснорме посчитать среднее квадратов по этому вектору: сложить 4096 циферки в произвольном порядке. На видеокарточках мы это делаем редукцией, делим вектор на блоки, суммируем блоки параллельно, затем суммируем суммы блоков между собой. Разбиение на блоки зависит от размера матча и количества активных SM (Streaming Multiprocessor, мини-процессор внутри видеокарточки). На малых батчах занимать, скажем, 132/132 SM на карточке это тупо, соответственно схема разбиения меняется, получаем другой порядок сложений.

Разница в конкретных числах будет мизерной, типа 1e-3 относительной ошибки на уровне одного оператора. Вот только таких операторов типа сотни на один форвард. Получаем аккумулирующуюся ошибку и достаточное изменение логитов для того, чтобы аргмакс возвращал уже другой токен.

Почему батчи разные при одинаковых запросах? Потому что в реальном проде есть continuous batching. Человеков, общающихся с моделькой, много, запросы от них приходят асинхронно, размер батчей постоянно гуляет. Зафиксировать батчи - значит, падать запросы до единого размера, значит терять вычисления, либо ждать накопления матча, то есть ловить латенси.

Детерминизм очень нужен, но в достаточно редких сценариях. Помимо тестов, CI и всяких аудитов, у нас есть новомодный on-policy RL, GRPO всякий и всё в таком духе. Суть его в том, что модель сама генерирует ответы, эти ответы получают ревард, модель учится делать лучше. On-policy в названии как раз и означает, что градиент мы считаем по тем токенам, которые модель сгенерировала прямо сейчас. Собственно, трейнинг луп видит токены, которые моделька сгенерила на батчсайзе 47, считает их вероятность при батчсайзе 8, получает другие логиты, вот оно уже и не on-policy. PPO-клиппинг частично спасает от большого дрейфа, но не от маленькой системной аккумуляции.

217 1 2
Каталог
Каталог каналов и чатов Подборки каналов Поиск каналов Добавить канал/чат
Рейтинги
Рейтинг каналов Telegram Рейтинг чатов Telegram Рейтинг публикаций Рейтинги брендов и персон
API
API статистики API поиска публикаций API Callback
Наши каналы
@TGStat @TGStat_Chat @telepulse @TGStatAPI
Почитать
Академия TGStat Исследование Telegram 2019 Исследование Telegram 2021 Исследование Telegram 2023
Контакты
Справочный центр Поддержка Почта Вакансии
Всякая всячина
Пользовательское соглашение Политика конфиденциальности Публичная оферта
Наши боты
@TGStat_Bot @SearcheeBot @TGAlertsBot @tg_analytics_bot @TGStatChatBot