Навчання з урахуванням квантування
Навчання з урахуванням квантування (QAT) доналаштовує модель PyTorch, одночасно моделюючи числові ефекти виконання INT8. Використовуйте його, коли квантування після навчання спричиняє неприйнятну втрату точності й модель можна повторно навчити на репрезентативних даних.
SiMa QAT призначено для роботи в наявному проєкті навчан ня PyTorch. Воно не замінює набір даних, аугментації, функцію втрат, оптимізатор або метрику перевірки. Ці частини початкового проєкту важливо зберегти, оскільки зазвичай саме вони забезпечують точність моделі з рухомою комою.
За можливості починайте з попередньо навченої контрольної точки з рухомою комою. Навчання з випадкової ініціалізації також підтримується, але зазвичай потребує значно більше часу й даних.
Як працює QAT
Підготовка додає до моделі спостерігачі й операції псевдоквантування. Спостерігачі вимірюють діапазони активацій, а псевдоквантування округлює та обмежує значення під час прямого проходу, наближуючи виконання INT8. Тензори, градієнти й оновлення оптимізатора залишаються у форматі з рухомою комою, тому навчання може адаптувати ваги моделі до ефектів квантування.
Робочий процес:
- Підготуйте модель PyTorch у режимі Eager до QAT.
- Прогрійте спостерігачі у звичайному циклі навчання.
- Зафіксуйте діапазони активацій і сумісні із SiMa масштабні коефіцієнти ваг.
- Відновіть точність, продовживши навчання із зафіксованими масштабними коефіцієнтами.
- Завершіть підготовку моделі до виконання.
- Експортуйте стандартну модель ONNX opset-17 із вузлами
QuantizeLinearіDequantizeLinear(QDQ).
Встановлення
Пакет QAT wheel потребує Python 3.10 або новішої версії та PyTorch 2.8.x. Встановіть його в середовище, яке вже містить залежності для навчання моделі.
Завантажте пакет QAT за допомогою sima-cli:
sima-cli neat install qat
Команда завантажує wheel та встановлює або оновлює навичку агента кодування QAT для Codex і Claude. Вона не змінює активне середовище Python. Активуйте середовище навчання та встановіть завантажений wheel:
python -m pip install ./sima_qat-*.whl
python -c "import torch, sima_qat; print(torch.__version__, sima_qat.__file__)"
Додавання QAT до проєкту навчання
Наведені нижче кроки утворюють єдиний робочий процес. Адаптуйте модель, дані, оптимізатор, функцію втрат і виклики перевірки до наявного проєкту навчання.
1. Підготовка моделі
Підготуйте модель до створення оптимізатора. Підготовка повертає ізольований граф QAT і не змінює та не переміщує початкову модель або приклади вхідних даних. Кортеж вхідних даних має відповідати позиційним аргументам, типам даних і формам моделі.
import torch
from sima_qat import (
sima_export_onnx,
sima_finalize_qat_model,
sima_freeze_qat,
sima_prepare_qat_model,
)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Recreate the model and load the floating-point checkpoint.
source_model = MyModel()
source_model.load_state_dict(torch.load("model-fp32.pt", map_location="cpu"))
source_model.train()
example_inputs = (torch.randn(1, 3, 224, 224),)
qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
)
optimizer = torch.optim.AdamW(qat_model.parameters(), lr=1e-5)
criterion = torch.nn.CrossEntropyLoss()
Створюйте оптимізатор із qat_model, а не з source_model, оскільки навчання
оновлює підготовлений граф.
2. Навчання, фіксація та відновлення
Спочатку навчайте модель у звичайному режимі, щоб спостерігачі виміряли репрезентативні діапазони активацій. Після прогрівання зафіксуйте параметри квантування, а потім продовжте навчання, щоб модель відновила точність із зафіксованими сітками квантування.
Під час перевірки залишайте псевдоквантування ввімкненим, але тимчасово вимикайте спостерігачі, щоб валідаційні дані не змінювали їхні діапазони. Самі лише eval() та inference_mode() не зупиняють спостерігачі. У finally відновлюйте їхні попередні стани, зокрема вимкнений стан після фіксації.
from torch.ao.quantization import disable_observer
from torch.ao.quantization.fake_quantize import FakeQuantizeBase
freeze_epoch = 2
num_epochs = 4
for epoch in range(num_epochs):
qat_model.train()
# Reserve one or more later epochs for recovery training.
if epoch == freeze_epoch:
sima_freeze_qat(qat_model)
for images, labels in train_loader:
images = images.to(device)
labels = labels.to(device)
optimizer.zero_grad(set_to_none=True)
predictions = qat_model(images)
loss = criterion(predictions, labels)
loss.backward()
optimizer.step()
observer_states = [
(module, module.observer_enabled.clone())
for module in qat_model.modules()
if isinstance(module, FakeQuantizeBase)
]
try:
qat_model.apply(disable_observer)
validate(qat_model, validation_loader, device)
finally:
for module, enabled in observer_states:
module.observer_enabled.copy_(enabled)
Епоха фіксації залежить від моделі. Корисною відправною точкою є прогрівання спостерігачів протягом більшої частини короткого доналаштування та резервування принаймні однієї останньої епохи для відновлення. Відстежуйте точність перевірки до й після фіксації. Якщо точність різко падає, виконайте фіксацію раніше та збільште тривалість відновлювального навчання.
3. Збереження та продовження навчання
Збережіть підготовлену модель до завершення, щоб навчання можна було продовжити. Контрольна точка має містити стани моделі QAT та оптимізатора.
from pathlib import Path
checkpoint_dir = Path("checkpoints")
checkpoint_dir.mkdir(parents=True, exist_ok=True)
torch.save(
{
"epoch": epoch,
"model": qat_model.state_dict(),
"optimizer": optimizer.state_dict(),
},
checkpoint_dir / f"qat-{epoch:02d}.pt",
)
Щоб продовжити навчання, відтворіть і підготуйте ту саму модель з тим самим прикладом вхідних даних і контрактом пакетної обробки, а потім завантажте збережені стани:
checkpoint = torch.load("checkpoints/qat-03.pt", map_location="cpu")
source_model = MyModel()
source_model.load_state_dict(torch.load("model-fp32.pt", map_location="cpu"))
qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
)
optimizer = torch.optim.AdamW(qat_model.parameters(), lr=1e-5)
qat_model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
start_epoch = checkpoint["epoch"] + 1
Контрольна точка зберігає стан спостерігачів, стан фіксації та точні параметри квантування, які згодом використовуються під час завершення й експорту ONNX.
4. Завершення та експорт
Завершення створює модель лише для виконання. Перемістіть навчену модель QAT і приклади вхідних даних на CPU, а потім експортуйте завершену модель як opset-17 QDQ ONNX.
final_model = sima_finalize_qat_model(qat_model.cpu())
export_inputs = tuple(value.cpu() for value in example_inputs)
sima_export_onnx(
final_model,
export_inputs,
"model.qdq.onnx",
input_names=["images"],
output_names=["predictions"],
device="cpu",
)
Розмір пакета
Під час підготовки перший вимір тензора (вимір батчу) за замовчуванням залишається
динамічним. Завдяки цьому той самий підготовлений граф можна
навчати зі звичайними розмірами пакетів завантажувача даних, обробляти короткий
останній пакет і експортувати з конкретним розміром пакета, переданим до
sima_export_onnx.
Для більшості моделей параметр пакетної обробки не потрібен. Установлюйте
dynamic_batch=False лише тоді, коли модель навмисно потребує точного розміру
пакета з прикладу:
qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
dynamic_batch=False,
)
Фіксована пакетна обробка доречна, коли модель перевіряє розмір пакета або розгалужує виконання за ним, використовує рекурентний стан фіксованого розміру чи поєднує пакет із напрямком, каналом або іншою геометрією компонування. Динамічне захоплення порівнює отриманий результат із початковою моделлю на наданому прикладі та завершується помилкою з порадою вимкнути динамічний пакет, якщо поведінка змінюється.
Перевірка та компіляція
Виміряйте цільову метрику для початкової моделі з рухомою комою, підготовленої моделі до й після фіксації, завершеної моделі PyTorch і моделі ONNX. Це дає змогу визначити, на якому етапі життєвого циклу виникла регресія.
Перед компіляцією перевірте експортований файл:
python - <<'PY'
import onnx
model = onnx.load("model.qdq.onnx")
onnx.checker.check_model(model)
print("ONNX model is valid")
PY
Порівняйте завершену модель PyTorch з ONNX Runtime на репрезентативних зразках для перевірки. На межах квантування можливі невеликі поелементні в ідмінності, тому використовуйте допуски, доречні для вихідних даних, і підтвердьте фактичну метрику точності моделі.
Підтримуються поширені зважені операції, активації, нормалізація, об'єднання,
редукції та операції форми й компонування. Вихідні індекси ArgMax і TopK
залишаються цілими числами. PReLU, ConvTranspose, Embedding/Gather, GridSample,
ReduceMin і CumSum можна навчати, але в цьому випуску вони не мають анотацій QAT.
Експортована модель QDQ ONNX є точкою передавання до Model Compiler. Імпорт, розподіл, оптимізація та призначення апаратних засобів є окремими кроками компіляції.
Приклади, які можна запускати
Каталог examples у репозиторії містить невеликий робочий процес MNIST для CPU, доналаштування попередньо навченого класифікатора ImageNet і робочий процес YOLO26n на чистому PyTorch із продовженням із контрольної точки та експортом.