Skip to content

Reinforcement learning and preferences ​

Llama 3.1 8B trained on one GPU with LoRA adapters, two ways: GRPO against a reward, with answers sampled by linnet.serve, and DPO on preference pairs.

Run ​

bash
python examples/08-rl/grpo.py     # learns to get two-number multiplications right
python examples/08-rl/dpo.py      # UltraFeedback's chosen and rejected answers

On one H100, against TRL on the same settings:

Linnet, PyTorchLinnet, JAXTRL
GRPO step: 16 prompts × 8 answers of up to 128 tokens1.18 s1.44 s3.14 s (vLLM sampling)
DPO step: 32 pairs1.91 s2.41 s4.63 s

GRPO ​

The reward is 1 when an answer's last number is the product and 0 otherwise. Over 40 steps, the share of right answers goes from 0.84 over the first ten steps to 0.99 over the last ten. The model learns to work each product out: its answers grow from about 11 tokens to 55.

Each step, linnet.train.grpo:

  1. copies the policy's weights, adapters merged, into the engine's model (Engine.load_weights), with no new compilation;
  2. samples group answers to each prompt with the engine;
  3. scores them with reward, and gives each answer its advantage over its group's mean;
  4. trains the policy on the packed answers with the clipped policy-gradient loss.

correction_cap=2.0 weighs each token by how much likelier the policy finds it than the engine did, capped at 2 (truncated importance sampling). That corrects for the two computing the same model with different kernels. The engine is the serving engine from 04-serve: continuous batching, packed prompts, one CUDA graph a step.

DPO ​

linnet.train.dpo packs each pair's two answers into one batch, sums their log-probabilities from the card's log_probs_packed entry, and computes -log sigmoid(beta * (margin - reference margin)). Without a separate reference model, the reference log-probabilities are the model's own, computed for every pair before training.

Both run fully sharded across GPUs in PyTorch (fully_shard) and over a mesh in JAX (linnet.jax.grpo and linnet.jax.dpo with mesh=). Training has every option.

Source ​

dpo.py ​

python
"""Preference optimization (DPO) for Llama 3.1 8B with LoRA, on one GPU:
UltraFeedback's chosen and rejected answers. The reference model's
log-probabilities are computed first, by the same model before it trains."""

import argparse

import torch
from datasets import load_dataset
from transformers import AutoTokenizer

from linnet import nest
from linnet.packing import Pair
from linnet.train.dpo import dpo

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("--pairs", type=int, default=32, help="pairs a step")
parser.add_argument("--lr", type=float, default=2e-5)
parser.add_argument("--beta", type=float, default=0.1)
args = parser.parse_args()

tokenizer = AutoTokenizer.from_pretrained(nest.resolve(args.card).weights.repo)


def answer(text: str) -> list[int]:
    return tokenizer(text + "<|eot_id|>", add_special_tokens=False)["input_ids"][:512]


pairs = []
for row in load_dataset("trl-lib/ultrafeedback_binarized", split="train").shuffle(seed=0):
    conversation = row["chosen"][:-1]
    if not conversation or conversation[-1]["role"] != "user":
        continue
    text = tokenizer.apply_chat_template(conversation, add_generation_prompt=True, tokenize=False)
    prompt = tokenizer(text, add_special_tokens=False)["input_ids"][-512:]
    pairs.append(
        Pair(prompt, answer(row["chosen"][-1]["content"]), answer(row["rejected"][-1]["content"]))
    )
    if len(pairs) == args.steps * args.pairs:
        break

model = nest.load(
    args.card,
    backend="torch",
    device="cuda",
    numerics="fast",
    compile="inductor",
    generics={"Batch": 1, "MaxSeq": 2048, "T": "bf16"},
    cast_dtype=True,
)
model.add_lora(["layers.*.attention.*_proj.weight", "layers.*.mlp.*.weight"], rank=16, alpha=32)


def report(step):
    print(
        f"step {step.step}: loss {step.loss:.3f}, accuracy {step.accuracy:.2f}, "
        f"margin {step.margin:.3f}, {step.seconds:.2f} s",
        flush=True,
    )


trained = [p for p in model.parameters() if p.requires_grad]
dpo(
    model,
    pairs,
    optimizer=torch.optim.AdamW(trained, lr=args.lr, weight_decay=0.0),
    steps=args.steps,
    pairs_per_step=args.pairs,
    beta=args.beta,
    tokens=4096,
    on_step=report,
)
print(f"peak {torch.cuda.max_memory_allocated() / 2**30:.1f} GiB")

grpo.py ​

python
"""Reinforcement learning (GRPO) for Llama 3.1 8B with LoRA, on one GPU: the
model learns to get two-number multiplications right. A `linnet.serve`
Engine samples a group of answers to each prompt from a second copy of the
model, which takes the policy's weights before every step; the policy
learns from each answer's reward against its group's."""

import argparse
import itertools
import random
import re

import torch
from transformers import AutoTokenizer

from linnet import nest
from linnet.packing import Prompt
from linnet.serve import Engine
from linnet.train.grpo import grpo

parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--card", default="llama-3.1-8b-instruct")
parser.add_argument("--steps", type=int, default=40)
parser.add_argument("--prompts", type=int, default=16, help="prompts a step")
parser.add_argument("--group", type=int, default=8, help="answers sampled per prompt")
parser.add_argument("--lr", type=float, default=5e-5)
parser.add_argument("--rank", type=int, default=16)
args = parser.parse_args()

tokenizer = AutoTokenizer.from_pretrained(nest.resolve(args.card).weights.repo)
eos = tokenizer.convert_tokens_to_ids(["<|eot_id|>", "<|end_of_text|>", "<|eom_id|>"])


def prompts():
    """Endless multiplications as chat prompts, each carrying its answer."""
    draw = random.Random(0)
    for _ in itertools.count():
        a, b = draw.randrange(10, 100), draw.randrange(10, 100)
        text = tokenizer.apply_chat_template(
            [{"role": "user", "content": f"What is {a} * {b}?"}],
            add_generation_prompt=True,
            tokenize=False,
        )
        yield Prompt(tokenizer(text, add_special_tokens=False)["input_ids"], a * b)


def reward(prompt: Prompt, completion: list[int]) -> float:
    """1 when the answer's last number is the product, 0 otherwise."""
    text = tokenizer.decode(completion, skip_special_tokens=True)
    numbers = re.findall(r"\d[\d,]*", text)
    return 1.0 if numbers and numbers[-1].replace(",", "") == str(prompt.data) else 0.0


# The policy trains its adapters; the engine samples from a copy of the model
# compiled for serving: 64 cache rows, each step one CUDA graph.
policy = nest.load(
    args.card,
    backend="torch",
    device="cuda",
    numerics="fast",
    compile="inductor",
    generics={"Batch": 1, "MaxSeq": 1024, "T": "bf16"},
    cast_dtype=True,
)
policy.add_lora(
    ["layers.*.attention.*_proj.weight", "layers.*.mlp.*.weight"],
    rank=args.rank,
    alpha=2 * args.rank,
)
engine = Engine(
    nest.load(
        args.card,
        backend="torch",
        device="cuda",
        numerics="fast",
        compile=True,
        generics={"Batch": 64, "MaxSeq": 1024, "T": "bf16"},
        cast_dtype=True,
    )
)


def report(step):
    print(
        f"step {step.step}: reward {step.reward:.3f}, answers {step.length:.1f} tokens, "
        f"sample {step.sample_seconds:.2f} s, train {step.train_seconds:.2f} s",
        flush=True,
    )


trained = [p for p in policy.parameters() if p.requires_grad]
grpo(
    policy,
    engine,
    prompts(),
    reward,
    optimizer=torch.optim.AdamW(trained, lr=args.lr),
    steps=args.steps,
    group=args.group,
    prompts_per_step=args.prompts,
    max_new_tokens=128,
    temperature=1.0,
    eos=eos,
    correction_cap=2.0,
    on_step=report,
)
print(f"peak {torch.cuda.max_memory_allocated() / 2**30:.1f} GiB")

Browse on GitHub

Released under the MIT License.