🔬 Метод
За основу берут метод из H-Net, но моделируют на уровне токенов, а не бит.
Dynamic Large Concept Model (DLCM) работает следующим образом:
1️⃣ Токенизируем текст каким-то токенизатором
2️⃣ Прогоняем текст через некий энкодер
3️⃣ Затем считаем похожесть (косинусное расстояние) между прошлой query и текущим ключом. Если расстояние больше заданного порога, то начинаем следующий токен, иначе сливаем текущий токен с тем что есть. Получаем таким образом укороченную последовательность.
4️⃣ Прогоняем эту укороченную последовательность через основную модель.
5️⃣ Декодируем обратно в исходное пространство токенов через cross-attention на исходную последовательность токенов.
При обучении задается желаемое сжатие R (сколько в среднем исходных токенов сжимаются в латентный токен). Для поддержания целевой степени сжатия добавляется вспомогательный лосс, который способствует тому, чтобы в среднем R токенов сливалось в один токен.
На обучении разбиение на токены сэмплируют из распределения Бернулли для exploration, на инференсе разбивают по порогу 0.5.
Из деталей реализации стоит отметить следующее. Так как при подаче батча с фиксированной длиной исходных токенов число латентных может разниться, в данной работе реплицируют их на этапе обучения. Flex Attention с паддингами выглядит естественным решением, но оказывается, что это работает медленнее (в 1.4-1.7 раз), чем Flash Attention c репликацией.
🧪 Эксперименты
Метод валидируют, обучая семейство моделей Llama-like архитектуры, используя токенизатор DeepSeek.
Для подбора оптимальных гипепараметров используют \muP параметризацию, как для энкодера, так и основной модели. Параметры настраивают на маленькой 87M модели и масштабируют на большие.
Кроме того, в данной статье предлагают scaling law лосса в зависимости от степени сжатия R, и доли параметров, приходящихся на энкодер P. Оказывается, что R=4 более менее оптимальный выбор с точки зрения соотношения качество/скорость.
Для оценки качества берут выборку из 12 бенчмарков из lm-eval-harness. DLCM дает прирост почти на всех бенчах и в среднем 2-3% качества по сравнению с стандартной токенизацией. Основной прирост на задачах, требующих reasoning,
Глобальная регуляризация (приведение среднего сжатия к R) лучше, чем на уровне отдельного предложения.
💡 Выводы
Неплохой результат с очевидной практической пользой - солидной экономией вычислений за счет более коротких последовательностей. Интересно, будет ли данное направление дальше развиваться и увидим ли мы SOTA-level LLM c отходом от стандартной токенизации?
Post #614
2K
- 👍 7
- 🔥 2
- 🙏 2