LoRA-дообучение Llama 3.1 8B на одной RTX 5090 за полдня

Форматирование датасета, настройки QLoRA, которые помещаются в 32 GB, сколько на самом деле занимает эпоха и как оценить результат, прежде чем запускать адаптер в работу.

GCGPU-команда CheapServАвтор 9 мин чтения
Видеокарта, над которой парит небольшая сеть, похожая на мозг
На этой странице8
  1. Что помещается в 32 GB
  2. Окружение
  3. Данные: то, что определяет результат
  4. Скрипт обучения
  5. Сколько это занимает
  6. Оцените модель до слияния
  7. Слияние и запуск
  8. Сколько это стоит

Раньше для дообучения требовался узел с восемью картами. Для большинства практических задач адаптации модели на 8B одной RTX 5090 с 32 GB VRAM достаточно, и вся работа укладывается в полдня. В этом руководстве мы берём Llama 3.1 8B Instruct, инструктивный датасет на 10 000 примеров и конфигурацию QLoRA, которая обучается примерно по 40 минут на эпоху, а затем показываем, как оценить адаптер, прежде чем сливать его с моделью и запускать в работу.

Что помещается в 32 GB

Полное дообучение 8B параметров в BF16 требует памяти под веса, градиенты и состояния оптимизатора: примерно 16 + 16 + 64 GB. Это не помещается. LoRA вместо весов обучает небольшие матрицы адаптера, а QLoRA дополнительно хранит замороженную базовую модель в 4-битном представлении. Тогда картина по памяти на 5090 выглядит так:

КомпонентQLoRA 4-bit, 8B, контекст 4k, батч 4
Базовые веса (NF4)~5.5 GB
Параметры LoRA + оптимизатор (ранг 32, все линейные слои)~0.8 GB
Активации при gradient checkpointing~9 GB
Контекст CUDA, кэш, фрагментация~3 GB
Итого~19 GB

Остаётся запас, чтобы увеличить размер батча или контекст до 8k. Обычная LoRA в BF16 (без квантования) тоже помещается — примерно в 26 GB — и обучается примерно на 30% быстрее; применяйте её, если примеры в ваших данных короче 4k токенов.

Окружение

Начните с шаблона LLaMA-Factory или Unsloth либо соберите окружение сами на базе Ubuntu 24.04 с CUDA 12.8:

uv venv /opt/ft && source /opt/ft/bin/activate
uv pip install torch==2.5.1 --index-url https://download.pytorch.org/whl/cu124
uv pip install "transformers==4.46.*" "peft==0.13.*" "trl==0.12.*" bitsandbytes datasets accelerate
huggingface-cli download meta-llama/Llama-3.1-8B-Instruct --local-dir /data/models/llama-8b

Данные: то, что определяет результат

Большинство неудачных дообучений упирается в данные. Оформляйте каждый пример в виде чат-шаблона, который модель уже знает, — с системным сообщением, если вы используете его в продакшене, — и выдерживайте единый стиль ответов: модель учит формат не меньше, чем содержание.

{"messages": [
  {"role": "system", "content": "You are a support assistant for Acme Cloud."},
  {"role": "user", "content": "How do I rotate my API key?"},
  {"role": "assistant", "content": "Open Account → API keys, click Rotate next to the key…"}
]}

Десяти тысяч примеров такого вида хватает с запасом для адаптации стиля и предметной области. Отложите 500 примеров для оценки до начала обучения и не заглядывайте в них при подборе гиперпараметров. Удаляйте дубликаты почти одинаковых промптов: дубликаты делают кривую loss красивой, а модель — хуже.

Скрипт обучения

from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
from peft import LoraConfig
from trl import SFTTrainer, SFTConfig
import torch

base = "/data/models/llama-8b"
tok = AutoTokenizer.from_pretrained(base)
model = AutoModelForCausalLM.from_pretrained(
    base, torch_dtype=torch.bfloat16,
    quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
                                           bnb_4bit_compute_dtype=torch.bfloat16))
ds = load_dataset("json", data_files={"train": "train.jsonl", "eval": "eval.jsonl"})
cfg = SFTConfig(output_dir="/data/runs/support-v1", num_train_epochs=2,
    per_device_train_batch_size=4, gradient_accumulation_steps=4, learning_rate=2e-4,
    lr_scheduler_type="cosine", warmup_ratio=0.03, bf16=True, gradient_checkpointing=True,
    max_seq_length=4096, logging_steps=10, eval_strategy="steps", eval_steps=100, save_steps=100)
lora = LoraConfig(r=32, lora_alpha=64, lora_dropout=0.05, task_type="CAUSAL_LM",
    target_modules=["q_proj","k_proj","v_proj","o_proj","gate_proj","up_proj","down_proj"])
SFTTrainer(model=model, args=cfg, peft_config=lora, train_dataset=ds["train"],
           eval_dataset=ds["eval"], tokenizer=tok).train()

Решения, которые стоит пояснить: ранг 32 на всех линейных слоях — оптимальное значение для 8B в наших тестах; learning rate 2e-4 — стандарт для LoRA и был бы слишком высоким для полного дообучения; а эффективный батч 16 с косинусным затуханием на протяжении двух эпох — это точка, в которой eval loss обычно достигает минимума для датасетов такого размера.

Сколько это занимает

КонфигурацияТокенов/с10 тыс. примеров × 2 эпохи (в среднем 600 токенов)
RTX 5090, QLoRA 4-bit, с checkpointing~5 000~40 мин на эпоху, 80 мин всего
RTX 5090, LoRA BF16, с checkpointing~6 500~31 мин на эпоху
RTX 4090, QLoRA 4-bit~3 100~65 мин на эпоху
H100 SXM, LoRA BF16, без checkpointing~14 000~14 мин на эпоху

В первые минуты следите за nvidia-smi: загрузка SM должна держаться выше 90%. Если это не так, узкое место — загрузчик данных (dataloader); токенизируйте датасет заранее или увеличьте число воркеров.

Оцените модель до слияния

Снижение eval loss доказывает лишь то, что модель запомнила ваш формат, но не то, что она стала лучше. Прогоните отложенные промпты через базовую модель и через адаптер и сравните ответы: в идеале по рубрике и со второй моделью в роли судьи, а можно и силами человека на выборке из пятидесяти примеров. Три вида сбоев проявляются именно здесь и больше нигде:

  • Переобучение на формулировки: каждый ответ начинается одинаково. Снизьте число эпох до одной или уменьшите ранг.
  • Забывание: модель хуже отвечает на общие вопросы, с которыми раньше справлялась. Добавьте в обучающую выборку от 10 до 20% общих инструктивных данных.
  • Исчезновение отказов: если в ваших данных модель ни разу ни от чего не отказывается, она перестаёт отказываться. Включите примеры того поведения, которое хотите сохранить.

Слияние и запуск

from peft import PeftModel
m = AutoModelForCausalLM.from_pretrained(base, torch_dtype=torch.bfloat16)
m = PeftModel.from_pretrained(m, "/data/runs/support-v1/checkpoint-1250").merge_and_unload()
m.save_pretrained("/data/models/support-v1"); tok.save_pretrained("/data/models/support-v1")

Объединённая модель запускается в vLLM точно так же, как базовая; руководство по vLLM подходит без изменений, а модель на 8B в BF16 на той же 5090 выдаёт около 90 токенов в секунду на запрос. Сохраните и чекпойнт адаптера: он занимает 300 MB и позволяет позже заново слить его с более новой базовой моделью.

Сколько это стоит

Полдня на тарифе с RTX 5090, включая загрузки и три прогона обучения при подборе параметров, — это лишь несколько долларов из месячной цены. Та же работа на почасовых облачных GPU для одного прогона обойдётся сопоставимо, а вот итерации, которых требует настоящее дообучение, — намного дороже; именно месячная карта, на которой можно оставить модели, делает вторую и третью попытки дешёвыми.

GC
GPU-команда CheapServ

Бенчмарки, образы и шаблоны для парка GPU-серверов.

Разверните первый сервер меньше чем за минуту.

Пополните баланс от $25 в BTC, ETH, XMR или USDT. Баланс не сгорает, а неиспользованные средства можно вернуть.

Создать аккаунт