Відновлення з урахуванням квантування перевершує точність 16-бітних LLM
Метод Quantization-Aware Healing (QAH) відновлює стиснуті 4-бітні моделі через дистиляцію безпосередньо з оригінального повнорозмірного вчителя. Застосований до GPT-OSS 120B, стиснутої до 60B в MXFP4, він перевершив власну 16-бітну версію bfloat16 у 7 з 9 бенчмарків.

Вплив: Високий
Чому це важливо
Інженери можуть розгортати 4-бітні моделі половинного розміру, які перевершують за точністю та швидкістю свої 16-бітні чекпоїнти.
TL;DR
- 01QAH здійснює дистиляцію 4-бітних моделей безпосередньо з повнорозмірного вчителя за допомогою KL-дивергенції.
- 02Моделі QAH перевершують власні 16-бітні чекпоїнти у математиці (+5.6 на AIME 2025) та складному аналізі.
- 03QAH досягає оптимуму у 7 разів швидше за QAT (100 кроків проти 700) і запобігає деградації при довгому навчанні.
Ключові факти
- GPT-OSS 60B MXFP4 LiveCodeBench
- 66.5 (проти 16-біт 66.0)
- Приріст AIME 2025 порівняно з 16-біт
- +5.6 пунктів
- Приріст AA-LCR порівняно з 16-біт
- +7.4 пунктів
- Кроки збіжності (QAH проти QAT)
- 100 кроків проти 700 кроків
Подолання бар'єра квантування
Зазвичай квантування моделей до 4 біт призводить до деградації логічного мислення. Застосування Quantization-Aware Healing (QAH) до моделі GPT-OSS 120B, стиснутої до 60B параметрів у форматі MXFP4, дозволило перевершити показники відповідного 16-бітного чекпоїнта bfloat16 у 7 з 9 бенчмарків.
Результати бенчмарків та стабільність навчання
- Контекстне мислення (AA-LCR): +7.4 пункту порівняно з
bfloat16. - Математика (AIME 2025): +5.6 пункту порівняно з
bfloat16. - Генерація коду (LiveCodeBench):
66.5балів проти66.0у повнорозмірної моделі-вчителя 120B. - GPQA Diamond:
67.4проти69.0у вчителя.
Порівняльний аналіз стабільності на моделі GPT-OSS 9B в MXFP4 показав, що QAH досягає пікового результату (54.9) за 100 кроків, тоді як QAT вимагає 700 кроків (54.6). Після 1 200 кроків QAT втрачає майже 19 пунктів через перевизначення, у той час як QAH залишається в межах 2 пунктів від піку.
Спробуй за 2 хвилини
# Example concept for chunked KL-divergence loss calculation
import torch
import torch.nn.functional as F
def chunked_kl_loss(student_logits, teacher_logits, chunk_size=1024):
loss = 0.0
total_tokens = student_logits.size(1)
for i in range(0, total_tokens, chunk_size):
s_chunk = F.log_softmax(student_logits[:, i:i+chunk_size, :], dim=-1)
t_chunk = F.softmax(teacher_logits[:, i:i+chunk_size, :], dim=-1)
loss += F.kl_div(s_chunk, t_chunk, reduction='batchmean')
return loss / (total_tokens / chunk_size)python
✓ Коли використовувати
- Під час квантування великих моделей (30B+) до 4 біт для високопродуктивного розгортання на GPU.
- Якщо стандартне навчання QAT виявляється нестабільним або призводить до деградації результатів.
✕ Коли НЕ варто
- Якщо логіти оригінальної моделі-вчителя недоступні для офлайн-кешування.
- Для малих моделей менше 3B параметрів, де структурне обрізання призводить до незворотної втрати ємності.
Що зробити сьогодні
- Розглянути пайплайни QAH з KL-дивергенцією під час квантування локальних моделей до 4 біт.
- Протестувати формат MXFP4 для локальних агентів та автономних середовищ виконання.
Джерела