Fine-Tuning LLMs

Course Content

Fine-Tuning LLMs

6 sections · 52 lessons

For 4-bit QLoRA, what does a typical BitsAndBytesConfig look like and why?


Where 14 GB goes when QLoRA trains an 8B modelLinear layersin NF4 — 3.5 GBEmbeddings andhead, bf16 — 2 GBLoRA weights,grads, Adam — 1 GBActivations,checkpointed — resttopbottomThe same run with a bf16 base needs 16 GB for weights alone.
Quantizing the frozen base is the whole trick — the part being trained was already tiny.

What you need to know

The config, line by line

Python
# transformers 5.x, peft 0.2x, trl 1.x, bitsandbytes 0.4x+import torchfrom transformers import AutoModelForCausalLM, BitsAndBytesConfigfrom peft import LoraConfig, prepare_model_for_kbit_trainingbnb = BitsAndBytesConfig(    load_in_4bit=True,    bnb_4bit_quant_type="nf4",    bnb_4bit_use_double_quant=True,    bnb_4bit_compute_dtype=torch.bfloat16,)model = AutoModelForCausalLM.from_pretrained(    "meta-llama/Llama-3.1-8B-Instruct", quantization_config=bnb, dtype=torch.bfloat16)model = prepare_model_for_kbit_training(model)   # also turns on gradient checkpointinglora = LoraConfig(r=16, lora_alpha=32, lora_dropout=0.05,                  target_modules="all-linear", task_type="CAUSAL_LM")# then: SFTTrainer(model=model, peft_config=lora, args=SFTConfig(optim="paged_adamw_8bit", ...))
LineWhy
load_in_4bit=TrueStore the frozen base weights at 4 bits (half a byte) instead of 2 bytes.
nf4NormalFloat-4: its 16 levels sit at the quantiles of a normal distribution, where pretrained weights actually are. The default fp4 loses more accuracy.
use_double_quantEach block of 64 weights has a scale constant. Double quantization stores those constants in 8-bit, saving about 0.37 bits per parameter — roughly 0.3 GB on an 8B model.
compute_dtype=bfloat16Weights are stored in 4-bit but computed in bf16. This keeps quality up, and it is why QLoRA is slower per step than bf16 LoRA. Use float16 only on GPUs without bf16 (T4, V100).
prepare_model_for_kbit_trainingKeeps norm layers in fp32, enables gradient checkpointing and makes gradients flow back to the adapters.
paged_adamw_8bitAdapter optimizer states in 8-bit, and memory spikes are paged to CPU instead of crashing.

In Transformers 5 the loading argument is dtype; the older torch_dtype still works but is deprecated.

Memory budget for Llama-3.1-8B

Partbf16 LoRAQLoRA
Base weightsabout 16 GBabout 5.5 GB (linear layers 4-bit; embeddings and output layer stay bf16, about 2 GB of that)
LoRA weights, gradients, Adam states (42M params)under 1 GBunder 1 GB
Activations (seq 2,048, micro-batch 2, checkpointing on)a few GBa few GB
Fits ona 40–80 GB GPU comfortablya 16–24 GB GPU

Two things to change in special cases

  • FSDP across several GPUs: add bnb_4bit_quant_storage=torch.bfloat16 so all shards share one storage dtype.
  • Quality first, memory second: load_in_8bit=True is closer to bf16 but saves less.

A real-life example

An early-stage Bengaluru startup builds a Hindi customer-support model on 12,000 real conversations (about 700 tokens each). They have one rented L4 GPU (24 GB). A bf16 LoRA run on the 8B base runs out of memory at sequence length 2,048, because the weights alone take 16 GB.

With the QLoRA config above, peak memory is about 14 GB. Two epochs are 12,000 × 700 × 2 ≈ 17M tokens. At the roughly 1,000 tokens per second they measured on the L4, that is about 4.7 hours. At an on-demand L4 price of around one dollar an hour, the run costs about $5.

After training, they do not ship the 4-bit model. They reload the base in bf16, attach the adapter, merge, and then quantize for serving with a serving format (AWQ or FP8) that vLLM runs fast.

Follow-up questions to expect

  • "Is QLoRA as good as 16-bit LoRA?" — The QLoRA paper (2023) reported NF4 with double quantization matching 16-bit fine-tuning on its benchmarks. In practice the gap is usually small, but check on your own eval.
  • "Why not quantize the LoRA weights too?" — They are the part being trained and need precise gradients, and they are under 1% of the parameters anyway.
  • "Would you serve the model in NF4?" — Usually not. bitsandbytes 4-bit is built for memory saving, and its inference is slower than AWQ, GPTQ or FP8 kernels in serving engines.