Skip to content

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.safetensors

It 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/sPeakHeld-out loss after 30 steps
Linnet, PyTorch1.04 s15.6K36 GiB1.369
Linnet, JAX1.33 s12.2K38 GiB1.374
TRL + PEFT1.83 s8.9K43 GiB1.373

What it shows ​

  • Adapters without a model class. model.add_lora(patterns, rank=16) adds A and B beside every weight whose path matches a glob pattern. The generated code computes x @ 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.pack packs examples end to end into rows of 4096 positions. The card's loss_packed entry attends within each one and learns only the answers.
  • No [tokens, vocab] logits. The loss runs through std.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}")

Browse on GitHub

Released under the MIT License.