qwen3-0.6b-4b-adapter / TRAINING.ru.md
recoilme's picture
TRAINING.ru.md: подробный русский разбор обучения адаптера
c1ca515
|
Raw History Blame Contribute Delete
30.2 kB

Как обучен адаптер Qwen3-0.6B → Qwen3-4B: подробный разбор

Это технический разбор адаптера adapter_v14_bal.safetensors — 220M-маппинга из пространства текстового энкодера Qwen3-0.6B в пространство текстового энкодера FLUX.2-klein-4B (Qwen3-4B). Английская карточка модели — в README.md, здесь то же самое, но с цифрами замеров, историей решений и разбором того, что не сработало.


1. Зачем вообще адаптер

Учитель в нашей дистилляции — FLUX.2-klein-4B. Его DiT обучен на текстовом условии, которое считает его собственный энкодер: Qwen3-4B, 36 слоёв, hidden 2560, склейка слоёв 9, 18, 27 — это joint_attention_dim = 7680, и ровно эти три слоя берёт сам пайплайн diffusers.

Наш micro-студент имеет дешёвый энкодер — Qwen3-0.6B (28 слоёв, hidden 1024), потому что в micro-модели весь бюджет уходит в DiT, а не в текст.

Возникает развилка при дистилляции:

  • кормить учителя его родным условием (нужен Qwen3-4B: +7.5 ГБ VRAM и ~0.3 с на шаг);
  • либо кормить учителя нашим условием, но переведённым в его пространство.

Второй путь и есть адаптер: y = A(наш_ТЕ(x)), где A учится по тексту, без картинок, без VAE и без диффузии. Цель — не «похоже на подсказку», а в точности то, что дал бы родной энкодер учителя на том же тексте. Тогда учитель ведёт студента в том режиме, в каком он работал всегда, а студент дистиллируется от условия, которое сам и умеет считать.

Ключевое упрощение: задача — чистая регрессия «эмбеддинги в эмбеддинги». Токенизаторы у Qwen3-0.6B и Qwen3-4B идентичны (vocab 151669, тот же chat_template, тот же padding_side=right), поэтому позиции ученика и учителя совпадают ровно, без выравнивания. У референсной работы (recoilme/sdxs, qwen_adapter_2b_4b) токенизаторы разные, и там позиции частично разъезжаются — у нас этой проблемы нет.


2. Как он устроен

текст → Qwen3-0.6B (заморожен) → hidden states слоёв 2,9,14,18,23,27 → склейка (B, 256, 6144)
      → адаптер → (B, 256, 7680) → DiT klein

y_i = MLP(x_i) + Attention(x)_i
      │           │
      │           └─ residual-ветка: 2 блока self-attention, d=1024, 8 голов, 39.6M
      └─ поточечная MLP: 6144 → 8192 → 8192 → 7680, GELU(tanh), 180.4M

Итого 220.0M параметров. Порядок склейки — послойно-майорный, [слой2 | слой9 | … | слой27], ровно так же, как klein собирает свои три слоя (stack(dim=1) → permute).

Применяется адаптер только к условию учителя. Наш DiT по-прежнему получает свой 4D-контекст (B, L, 6, 1024) и обрабатывает его сам (TextFusion + txtmlp) — адаптер существует для взаимодействия с учителем, а не для студента.

Почему шесть слоёв, а не три

Соблазн взять симметрию: klein использует три слоя из 36 (25%, 50%, 75%) — значит, и нам взять три. Мы так и начали, и это оказалось хуже по всем замерам:

вход адаптера параметров отн. ошибка
3 слоя → 3072, MLP 3072→4096→7680, полный промпт 44M 36.03%
то же, но --drop-first 5 и 5 датасетов × 3000 44M 32.04%
3 слоя → 3072, MLP 3072→6144→6144→7680 104M 29.92%
6 слоёв → 6144, MLP 6144→8192→8192→7680 + ln-per-layer 180M 25.43%

Причин у выигрыша шесть-слоёв две, и обе измеримые.

Первая — нормы слоёв растут как 13 → 557. Норма выхода слоя Qwen3-0.6B по глубине отличается в десятки раз. Если склеить слои как есть, вход адаптера почти целиком состоит из глубоких слоёв, а ранний слой даёт 0.05% — то есть его как будто нет. Именно поэтому шесть слоёв идут вместе с --norm ln-per-layer (у каждого среза свои статистики, без обучаемых параметров): после неё все шесть слоёв сопоставимы по масштабу и каждый начинает работать. Отдельного замера «шесть слоёв без нормировки» у нас нет — но при таком разбросе норм он выродился бы в «три глубоких слоя», то есть в 104M-строку таблицы выше.

Вторая — глубина чтения. Выход слоя 9 у 36-слойной модели — это абстракция, снятая после девяти блоков. У 28-слойной модели такой же по смыслу абстракции соответствует не слой 7 (та же доля глубины), а другой слой — и «правильная» глубина чтения зависит от токена: для поверхностных вещей (регистр, отдельный тег, эмодзи) информативнее ранние слои, для смысла — поздние. Отдав адаптеру несколько глубин, мы позволяем ему выбирать чтение самому, вместо того чтобы угадывать единственный «правильный» слой. Эмпирика это подтверждает: набор 2,9,14,18,23,27 — не равномерная сетка (9 и 14 стоят ближе друг к другу), он выбирался по замерам, а не по формуле.

Дополнительный аргумент: ровно этот набор — то, что ест наш micro-DiT (TextFusion → txtmlp). Условие адаптера и условие студента — это одни и те же фичи, разница только в цели (пространство 4B против своего). Для дистилляции это удобно: студент учится на том же представлении текста, что и потом использует сам.

Почему не взять все 28 слоёв: во-первых, вход первого линейного слоя вырос бы до 28 672 (первый слой стал бы 235M параметров сам по себе), во-вторых, соседние слои сильно коррелированы, в-третьих, студенту нужны именно эти шесть. Это выбор, а не измеренный оптимум — сравнения «6 против 12 против 28» у нас нет.

Почему MLP, а не что-то другое

Рассматривались четыре варианта:

  1. Дообучить сам Qwen3-0.6B, чтобы он имитировал выход 4B. Дороже (трогаем 596M), ломает модель для любых других применений и требует учителя «в цикле» при каждом использовании — а нам нужен маленький замороженный артефакт.
  2. Линейная проекция 6144 → 7680 (47M). Не проверяли отдельно, но вся лестница ёмкости показывает, что одной матрицы тут мало: области значений у энкодеров разной глубины и ширины, соответствие нелинейное.
  3. MLP — выбрали его. Дешевле всех в обучении, учится по тексту без картинок, работает побатчно и параллельно по токенам.
  4. Transformer/Perceiver — оставили как ступень развития (см. п. 4).

Форма MLP — 6144 → 8192 → 8192 → 7680, три линейных слоя с двумя GELU(tanh): лестница ёмкости 44M → 104M → 180M даёт монотонный выигрыш (32.0% → 29.9% → 25.4%). В этой лестнице менялись две вещи сразу — ширина и число входных слоёв (3072 против 6144), поэтому вклад именно ширины мы не выделяли: выигрыш дают и ёмкость, и более широкий вход. Ясно только, что дальше 180M упирается уже не в параметры, а в данные.

Почему одной поточечной MLP мало

Поточечная MLP обрабатывает каждый токен независимо — контекста соседей у неё нет. А смысл в тексте распределён по токенам: a bulldog токенизатор режет как ['a', 'Ġbulld', 'og'], a rottweiler — как ['ro', 'tt', 'we', 'iler'], striped bowtie — как ['strip', 'ed', 'Ġbow', 'tie']. Собрать «бульдога» из Ġbulld и og поточечное отображение не может: у него нет способа посмотреть на соседа.

Это видно и в цифрах: если токен совпадает с целым словом, перенос почти идеальный (a red fox — 23.8%, an orange cat — 24.0%), а на обрубках ошибка резко выше (a bulldog — 26.6%, a rottweiler — 33.1%).

Было две попытки добавить контекст.

Попытка 1 — окно соседей (v10). Первый слой получает не один токен, а срез [x_{i-1}, x_i, x_{i+1}] (281.1M параметров). Соседи действительно обучились: нормы блоков ‖W_лев‖ = 15.9, ‖W_прав‖ = 15.0 против ‖W_центр‖ = 124.0 — то есть модель честно использует контекст. Но выигрыш ровный ~1 п.п. на всех промптах, включая цельные слова, а худший обрубок og не сдвинулся вообще (46.1% и до, и после). Вывод: локальное окно не решает задачу — «бульдог» собирается не из соседних токенов, а требует внимания ко всему предложению.

Попытка 2 — residual-ветка внимания (v11, в v14 она же). Поверх MLP добавлен самостоятельный модуль: Attention(x) = out(блоки(norm(x))), где out инициализирован нулями. Это принципиально: на старте обучения адаптер побитово равен старой поточечной MLP, то есть новая архитектура не может испортить уже выученное, а обучение начинается как точное дообучение. Дальше ветка сама набирает вес.

Параметры выбраны скромно: 2 блока, d=1024, 8 голов, 39.6M — всего 18% от общего размера. Одного блока, скорее всего, хватило бы, но два дали лучший результат из измеренного, а цена в 40M при бюджете 220M незначительна. Аттеншн без причинности: текст не автогенеративный, ограничивать токен предыдущими незачем.

Лосс и метрика

Лосс: MSE по токенам с пер-токенным весом

вес_i = 1 / max(‖y_i‖², 300²),   паддинг: 0.1

Без весов лосс на 97% состоит из ошибки пяти служебных токенов шаблона (attention-sink): их норма ~6000 против ~180 у смысловых. Модель честно оптимизировала бы эти пять токенов, а промпт не учила бы вовсе. Поэтому:

  • --norm-floor 300 — порог, выше которого токен считается «крупным» и его вклад нормируется по норме;
  • --drop-first 5 — первые пять токенов шаблона не идут ни в адаптер, ни в цель (срезка делается после адаптера, иначе позиционные ветки видели бы 251 токен на обучении и 256 на инференсе);
  • --pad-weight 0.1 — паддинг-позиции участвуют, но с малым весом.

Метрика: относительная ошибка по токену ‖pred − y‖ / ‖y‖, отдельно по трём группам:

группа что это токенов средняя ‖y‖ отн. ошибка
обычные смысловые токены промпта (‖y‖ < 300) 1499 (16 капшенов) 180.7 21.3%
крупные attention-sink (‖y‖ ≥ 300), после drop-first их почти нет 0–1 ~6000 —
паддинг позиции после конца промпта — в последовательности их ~63% 2517 143.5 60.0%

Паддинг не падает и не должен: содержания промпта в нём нет, вес в лоссе 0.1, а цель там — обобщённый «режим ассистента» учителя с почти постоянной нормой (разброс по позициям ±18 при средней 143.5). Смотреть надо только на «обычные»; ноль на паддинге означал бы зубрёжку константы.

Отдельно про грабли: element-wise MSE и R² здесь нельзя использовать. В референсной работе приведён R² = 0.9995 — но он измерен без маски, и эти 0.9995 на 99.6% состоят из подгонки тех самых пяти синк-токенов. Мы через это прошли и получили ровно такую же «отличную» цифру при реальной относительной ошибке 36–38%.


3. Как обучали

Команда финального прогона:

python3 train_adapter.py --ds-path datasets \
    --max-per-group 3000 --ds-skip testg_320_640_temp --seed 42 \
    --drop-first 5 --attention 2 --attn-dim 1024 --attn-heads 8 \
    --empty-ratio 0.05 --epochs 8 --batch-size 32 --log-every 100 --eval-every 1 \
    --student-layers 2,9,14,18,23,27 --hidden 8192 --proj-layers 3 --norm ln-per-layer \
    --lr 1e-4 --attn-lr 2e-4 --lr-schedule cos --epoch-checkpoints \
    --resume adapter_v12_full.pt --adapter-out adapter_v14_bal.pt

Без кэша энкодеров. Оба энкодера живут онлайн в VRAM (учитель 7.5 ГБ в bf16), кэша эмбеддингов нет. Это дороже по памяти (пик 18.5 ГиБ), зато любой датасет и любой текст тренируются сразу, а не требуют предварительного прогона. Один шаг — 0.68 с (batch 32, длина 256), эпоха — 568 шагов, ~7 минут.

Данные и их состав. Первый прогон был на 13 163 капшенах, потом 64 450, 95 992, 16 157 — и каждый раз цифра улучшалась на десятые доли. Главное открытие оказалось не в объёме, а в составе:

  • обучение на полном пуле (323 173 капшена, 86% которого — один прозаический датасет) даёт модель, которая уезжает в прозу: короткие промпты ухудшаются на 2–3 п.п., и на картинках это видно (бульдог становится всё менее похож);
  • поэтому v14 учили на сбалансированном пуле: по 3000 капшенов на группу источников (--max-per-group группирует по верхней папке датасета), включая danbooru-теги с вычищенной служебной разметкой *tags*;
  • v14 (19 157 капшенов) обошёл v12 (323 173) по всем колонкам. Состав важнее размера.

Пустые промпты обязательны. --empty-ratio 0.05 — 5% обучающих примеров с пустым промптом. Это нужно не для качества условия, а для CFG: у klein-base ветка uncond считается по пустому промпту через тот же адаптер, и если её не учить, при guidance_scale=4.0 картинка портится. Наглядно: v11 (без пустых) — ошибка uncond 20.2%, v12/v14 (с пустыми) — 3.0% / 2.3%. Разница в 8 раз — и именно она решает, будет ли CFG чистым.

Расписание. Обучение с нуля — 3e-4 с косинусом, дообучение — 1e-4 (MLP) и 2e-4 (ветка внимания) с косинусом до нуля. Ветка внимания берёт свой LR: у неё нулевая инициализация, и ей нужен чуть больший шаг, чтобы набрать вес за разумное время.

Резюм и падения. Прогон v14 один раз умер посреди эпохи (оборванный процесс, без трейсбека). Чекпоинт .last.pt пишется каждую эпоху и содержит веса и оптимизатор и номер эпохи, поэтому продолжение — это одна строка --resume ...last.pt при --epochs больше пройденного. Срезы по эпохам (--epoch-checkpoints) пишутся отдельно, чтобы выбирать не «последнюю», а лучшую эпоху.


4. Что получилось

Версии (одни и те же данные: val 627 капшенов + шесть коротких промптов)

чекпоинт bulldog rottweiler striped red fox orange cat теги пустой негатив val 627
v8 — поточечная 180M 29.18 35.80 29.37 26.92 26.69 23.09 4.64 37.40 24.36
v11 — +внимание, без пустых 27.17 34.18 27.06 24.60 24.11 22.40 20.23 36.85 23.73
v12 — 2 эпохи, полный пул 323k 29.52 35.82 30.02 26.99 27.29 23.70 2.99 38.58 23.83
v14 — 8 эпох, баланс + теги 26.56 33.12 25.17 23.77 23.95 20.19 2.27 35.59 21.66

Исторические прогоны (v5–v7) учились на других val-наборах и в эту таблицу не включены: их цифры 25.43 / 24.96 / 24.69% несопоставимы с «627» строками.

Лестница эпох v14 (тот же прогон, те же данные)

эпоха bulldog striped red fox orange теги пустой val 627
ep0 27.79 26.92 25.12 25.34 21.92 3.50 22.87
ep1 27.20 26.17 24.53 24.70 20.94 3.23 22.55
ep2 27.04 25.81 24.32 24.62 21.02 2.93 22.21
ep3 26.77 25.41 24.11 24.25 20.46 2.70 22.01
ep4 26.58 25.25 23.80 24.03 20.42 2.57 21.85
ep5 26.61 25.30 23.87 24.12 20.32 2.40 21.73
ep6 26.56 25.19 23.81 24.00 20.21 2.29 21.67
ep7 26.56 25.17 23.77 23.95 20.19 2.27 21.66

Кривая монотонна, но к концу насыщается: ep5 → ep7 добавляют по 0.03–0.06 п.п. Восемь эпох на сбалансированном пуле — примерно там, где дальнейшее обучение перестаёт платить за себя. Побочный приятный эффект: uncond падает с 3.50% до 2.27% — CFG становится чище от эпохи к эпохе, ничего специально для этого делать не нужно.

Что переносится, а что нет (по картинкам)

Панели «родной Qwen3-4B | Qwen3-0.6B + адаптер», дистиллированная klein, 4 шага, один сид (media/ в репозитории):

  • переносится: сцена, свет, композиция, поза, одежда, стиль, рендер текста — надпись «Qwen3-0.6b» на табличке верна; аниме и danbooru-теги (кошачьи уши, бант, хвост, клыки, поза) работают; прозаический портрет женщины совпадает с родным почти неотличимо;
  • не переносится: идентичность объекта. «Бульдог» выходит терьером, «ротвейлер» — тоже не ротвейлером. Это лучшая из наших цифр по бульдогу (26.56 против 29.18 у поточечной), и всё равно порода не та.

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


5. Что не сработало (и почему это полезно знать)

попытка результат
окно соседей k=1 вместо внимания (v10) соседи обучились (‖W‖ 15.9/15.0 против 124.0 у центра), выигрыш ровный ~1 п.п., худший обрубок не сдвинулся
негативные промпты в обучении (v9, --neg-ratio) ошибка негативов упала 37.4% → 15.7%, но появились синева и пятна на картинках при одном и том же сиде; откатили
обучение на полном пуле 323k (v12) val улучшился, короткие промпты ухудшились — модель уехала в прозу
валидация по прозаическому val скрыла провал на коротких промптах (проза доминировала в цифре)
element-wise MSE / R² как метрика измеряет пять синк-токенов: «R² 0.9995» при реальной ошибке 36–38%
первые версии без --drop-first 97% лосса — служебные токены шаблона
шесть слоёв без ln-per-layer ранний слой даёт 0.05% входа, т.е. шести слоёв как будто нет

6. Инженерные грабли

  • fp8 для энкодера учителя. Официальный fp8-файл klein diffusers не читает (chunk expects at least a 1-dimensional tensor на scale-тензорах). Работает QuantoConfig(weights="float8") через optimum-quanto (устаревший, но живой); torchao-вариант с текущими torch молча не срабатывает.
  • is_distilled в klein. Дистиллированная klein молча отключает CFG (do_classifier_free_guidance = guidance_scale > 1 and not is_distilled), а в пайплайне нет negative_prompt — только negative_prompt_embeds. Для экспериментов с CFG нужен неизбелённый близнец FLUX.2-klein-base-4B; при загрузке из локальной папки model_index.json надо править на is_distilled: false, иначе CFG отключится незаметно.
  • Два флага --len. Длина условия 256, но первые пять токенов шаблона — синк; срезка обязана быть в одном и том же месте на обучении и инференсе, иначе позиционные ветки видят разное число токенов.
  • Скорость. Дистиллированная klein fp8, 4 шага — 1–2 с на картинку 768×1280; base fp8, 50 шагов с CFG 4 — 18–30 с; форвард энкодера учителя — ~0.3 с.

7. Открытые вопросы

  1. Идентичность объекта — главный провал. Данные и эпохи его не лечат (v6: +51k капшенов и +1 эпоха дали ту же цифру и ту же картинку). Следующие нерадикальные ступени: дилатация окна соседей, второй уровень смешивания, token-mixing MLP (ранг 64, +0.2M).
  2. Нет критерия приёмки по существу. Правильный тест — сравнить поле скоростей klein-DiT с родным условием и через адаптер (одни латенты, один t) и решать по этой цифре, а не по проценту ошибки условия.
  3. Нужен ли адаптер вообще. Если в конкретном пайплайне можно позволить себе Qwen3-4B (+7.5 ГБ VRAM, ~0.3 с/шаг), адаптер не нужен: кормите учителя родным энкодером, а наш 0.6B оставьте студенту. Адаптер — это размен «качество против VRAM и скорости».
  4. Сравнение 6 слоёв против 12 и 28 — не измерено, мы остановились на рабочем наборе.

8. Как повторить

# 1. Обучение с нуля (сбалансированный пул + теги, ветка внимания с нуля)
python3 src/train_adapter.py --ds-path datasets --max-per-group 3000 --seed 42 \
    --drop-first 5 --attention 2 --attn-dim 1024 --attn-heads 8 --empty-ratio 0.05 \
    --epochs 8 --batch-size 32 --student-layers 2,9,14,18,23,27 \
    --hidden 8192 --proj-layers 3 --norm ln-per-layer --lr 3e-4 --attn-lr 2e-4 \
    --lr-schedule cos --epoch-checkpoints --adapter-out adapter_v14_bal.pt

# 2. Экспорт в safetensors (схема пишется в метаданные файла, self-check внутри)
python3 src/export_safetensors.py --in adapter_v14_bal.pt --attn-heads 8 \
    --student-layers 2,9,14,18,23,27 --drop-first 5 --len 256

# 3. Использование: подменяем у пайплайна только расчёт условия
python3 example.py "A red fox walks through a snowy forest at dusk." fox.png

Требования: GPU с ~24 ГБ (оба энкодера онлайн + адаптер + активации, пик обучения 18.5 ГиБ), diffusers с поддержкой Flux2KleinPipeline, transformers с Qwen3, и для fp8-учителя — optimum-quanto.