LoRA fine-tuning
Llama 3.1 8B from Nest fine-tuned with LoRA adapters on one GPU: Alpaca, packed into 4096-token rows, trained by linnet.train.
Run
bash
python examples/06-lora/sft.py # 30 steps of 4 rows
python examples/06-lora/sft.py --steps 300 --out adapters.safetensorsIt prints each step, the held-out loss before and after, the speed and peak memory, and writes the adapters.
On one H100, against TRL with PEFT on the same data and settings:
| Step (4 rows) | Tokens/s | Peak | Held-out loss after 30 steps | |
|---|---|---|---|---|
| Linnet, PyTorch | 1.04 s | 15.6K | 36 GiB | 1.369 |
| Linnet, JAX | 1.33 s | 12.2K | 38 GiB | 1.374 |
| TRL + PEFT | 1.83 s | 8.9K | 43 GiB | 1.373 |
What it shows
- Adapters without a model class.
model.add_lora(patterns, rank=16)addsAandBbeside every weight whose path matches a glob pattern. The generated code computesx @ W.T + (x @ A.T) @ B.T * alpha / rank, and only the adapters train.merge_lora()folds them back into the weights for serving or export. - No padding.
linnet.train.packpacks examples end to end into rows of 4096 positions. The card'sloss_packedentry attends within each one and learns only the answers. - No
[tokens, vocab]logits. The loss runs throughstd.nn.loss::linear_cross_entropy, which PyTorch computes a block of rows at a time, with the gradient computed alongside.
The same model trains under transformers' Trainer and TRL's SFTTrainer through linnet.torch.CausalLM (other trainers). On four GPUs, fully_shard splits the frozen weights and keeps the adapters whole (07-fsdp).
Source
sft.py
python
"""Supervised fine-tuning of Llama 3.1 8B with LoRA adapters on one GPU:
Alpaca packed into 4096-token rows, trained by `linnet.train`. Prints each
step and the held-out loss before and after, and saves the adapters."""
import argparse
import statistics
import torch
from datasets import load_dataset
from transformers import AutoTokenizer
from linnet import nest
from linnet.train import Example, cosine_schedule, pack, train
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--card", default="llama-3.1-8b-instruct")
parser.add_argument("--steps", type=int, default=30)
parser.add_argument("--accumulate", type=int, default=4, help="packed rows a step")
parser.add_argument("--lr", type=float, default=2e-4)
parser.add_argument("--rank", type=int, default=16)
parser.add_argument("--out", default="adapters.safetensors")
args = parser.parse_args()
TOKENS = 4096 # positions in a packed row
tokenizer = AutoTokenizer.from_pretrained(nest.resolve(args.card).weights.repo)
def example(row: dict) -> Example:
"""An Alpaca row as a chat turn: the prompt is context, the answer is learned."""
user = row["instruction"] + (f"\n\n{row['input']}" if row["input"] else "")
text = tokenizer.apply_chat_template(
[{"role": "user", "content": user}], add_generation_prompt=True, tokenize=False
)
prompt = tokenizer(text, add_special_tokens=False)["input_ids"][-1024:]
answer = tokenizer(row["output"] + "<|eot_id|>", add_special_tokens=False)["input_ids"]
return Example.prompted(prompt, answer[: 2048 - len(prompt)])
rows = load_dataset("tatsu-lab/alpaca", split="train").shuffle(seed=0).select(range(8400))
held = list(pack([example(row) for row in rows.select(range(200))], TOKENS))
batches = pack((example(row) for row in rows.select(range(200, len(rows)))), TOKENS)
# The card's source, compiled for one packed row; Inductor fuses the step.
model = nest.load(
args.card,
backend="torch",
device="cuda",
numerics="fast",
compile="inductor",
generics={"Batch": 1, "MaxSeq": TOKENS, "T": "bf16"},
cast_dtype=True,
)
# Adapters beside every attention and MLP projection; only they train.
model.add_lora(
["layers.*.attention.*_proj.weight", "layers.*.mlp.*.weight"],
rank=args.rank,
alpha=2 * args.rank,
)
def held_out() -> float:
count = sum(batch.count for batch in held)
with torch.no_grad():
return sum(
float(model.run_entry("loss_packed", batch.inputs(count, "cuda"))) for batch in held
)
before = held_out()
trained = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trained, lr=args.lr, weight_decay=0.0)
schedule = cosine_schedule(optimizer, 3, args.steps, floor=0.0)
def report(step):
print(f"step {step.step}: loss {step.loss:.3f}, {step.seconds:.2f} s", flush=True)
history = train(
model,
batches,
optimizer=optimizer,
steps=args.steps,
accumulate=args.accumulate,
schedule=schedule,
on_step=report,
)
step = statistics.median(s.seconds for s in history.steps[2:])
print(
f"held-out loss {before:.3f} -> {held_out():.3f}; {step:.2f} s a step "
f"({args.accumulate * TOKENS / step:.0f} positions/s), "
f"peak {torch.cuda.max_memory_allocated() / 2**30:.1f} GiB"
)
model.save_weights(args.out, names="linnet", include=["*.lora_a", "*.lora_b"])
print(f"adapters written to {args.out}")