Skip to content

Full fine-tuning across GPUs ​

Every weight of Llama 3.1 8B trained across four GPUs with fully sharded data parallelism: each GPU holds a quarter of every weight, gradient and AdamW state, and gathers a layer's weights only while the layer runs. The model is the same card in PyTorch and in JAX.

Run ​

bash
torchrun --nproc_per_node=4 examples/07-fsdp/sft_torch.py
python examples/07-fsdp/sft_jax.py            # one process, every GPU as a mesh
XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 python examples/07-fsdp/sft_jax.py --no-remat

The JAX run recomputes each layer in the backward pass by default. --no-remat keeps them, which is faster but needs more than XLA's default 75% of the GPU.

On some H100 hosts NCCL fails to set up NVLink SHARP; set NCCL_NVLS_ENABLE=0 if it does.

On four H100s, 30 steps of one 4096-token row per GPU:

StepTokens/sPeak a GPUHeld-out loss
Linnet, PyTorch0.55 s29.4K48 GiB1.91 to 1.342
Linnet, JAX, --no-remat0.53 s30.9K62 GiB1.91 to 1.345
Linnet, JAX0.66 s24.6K39 GiB1.91 to 1.344
TRL, FSDP20.57 s28.6K52 GiB

What it shows ​

  • Sharding needs nothing in the source. linnet.torch.fully_shard splits each layer's weights across the processes in f32, and compiles the entries to gather a layer's weights in bf16 where it first uses them. The model is one generated function, which PyTorch's own fully_shard cannot hook into.
  • JAX from the same card. linnet.jax.train(..., mesh=mesh) runs each device on its own batch (shard_map). The generated code gathers each layer's weights where it runs (linnet jax --fully-shard), and reduces gradients into the parts in f32.
  • Memory for compute. remat=True wraps each layer in jax.checkpoint: the backward pass gathers and computes the layer again, and a step keeps each layer's input alone. Peak memory falls from 62 to 39 GiB a GPU for a quarter more time a step.
  • Checkpoints. train(..., checkpoint=dir, checkpoint_every=100) resumes from the latest one, in either framework; save_weights gathers the parts and writes one SafeTensors file.

Source ​

sft_jax.py ​

python
"""The same full fine-tuning in JAX, from the same card: one process over
every local GPU as a mesh. Each device takes its own packed row and holds
part of every weight; each layer gathers its weights where it runs, and the
backward pass gathers and computes each layer again instead of keeping it
(`--no-remat` keeps it: faster, in 62 GiB a GPU rather than 39, beyond XLA's
default memory fraction, so set XLA_PYTHON_CLIENT_MEM_FRACTION=0.95)."""

import argparse
import statistics

import jax
import numpy as np
import optax
from datasets import load_dataset
from jax.sharding import Mesh
from transformers import AutoTokenizer

from linnet import nest
from linnet.jax.train import train
from linnet.packing import Example, pack

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("--lr", type=float, default=1e-5)
parser.add_argument(
    "--no-remat", dest="remat", action="store_false", help="keep each layer for backward"
)
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:
    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))
batches = pack((example(row) for row in rows.select(range(200, len(rows)))), TOKENS)

# The card's `loss_packed` entry as generated `jax.numpy`, f32 master weights
# computed in bf16.
model = nest.load(
    args.card,
    backend="jax_source",
    entry="loss_packed",
    numerics="fast",
    generics={"Batch": 1, "MaxSeq": TOKENS, "T": "bf16"},
    cast_dtype=True,
)
mesh = Mesh(np.array(jax.devices()), ("data",))
schedule = optax.warmup_cosine_decay_schedule(0.0, args.lr, 3, args.steps, 0.0)


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


params, history = train(
    model,
    batches,
    optimizer=optax.adamw(schedule, weight_decay=0.0),
    steps=args.steps,
    mesh=mesh,
    remat=args.remat,
    on_step=report,
)
step = statistics.median(s.seconds for s in history.steps[2:])
peak = max(d.memory_stats().get("peak_bytes_in_use", 0) for d in jax.devices()) / 2**30
print(
    f"{step:.2f} s a step ({mesh.size * TOKENS / step:.0f} positions/s on {mesh.size} GPUs), "
    f"peak {peak:.1f} GiB a GPU"
)

sft_torch.py ​

python
"""Full fine-tuning of Llama 3.1 8B across GPUs with fully sharded data
parallelism (FSDP): each process holds a quarter of every weight, gradient
and AdamW state (on four GPUs), and gathers a layer's weights only while the
layer runs. Run under torchrun:

    torchrun --nproc_per_node=4 examples/07-fsdp/sft_torch.py
"""

import argparse
import os
import statistics

import torch
from datasets import load_dataset
from transformers import AutoTokenizer

from linnet import nest
from linnet.torch import fully_shard
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("--lr", type=float, default=1e-5)
parser.add_argument("--out", default=None, help="write the trained weights here (bf16)")
args = parser.parse_args()
TOKENS = 4096  # positions in a packed row

torch.distributed.init_process_group("nccl")
rank, world = torch.distributed.get_rank(), torch.distributed.get_world_size()
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
device = f"cuda:{torch.cuda.current_device()}"
tokenizer = AutoTokenizer.from_pretrained(nest.resolve(args.card).weights.repo)


def example(row: dict) -> Example:
    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))
batches = list(pack([example(row) for row in rows.select(range(200, len(rows)))], TOKENS))

model = nest.load(
    args.card,
    backend="torch",
    device=device,
    numerics="fast",
    compile="inductor",
    generics={"Batch": 1, "MaxSeq": TOKENS, "T": "bf16"},
    cast_dtype=True,
    trainable=True,
)
# Splits every layer's weights across the processes, kept in f32, and
# compiles the entries to gather each layer's in bf16 where it runs.
fully_shard(model)
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):
    if rank == 0:
        print(f"step {step.step}: loss {step.loss:.3f}, {step.seconds:.2f} s", flush=True)


# One packed row a process a step: the processes take turns over the batches.
history = train(
    model,
    iter(batches[rank::world]),
    optimizer=optimizer,
    steps=args.steps,
    schedule=schedule,
    on_step=report,
)
peak = torch.tensor([torch.cuda.max_memory_allocated() / 2**30], device=device)
torch.distributed.all_reduce(peak, op=torch.distributed.ReduceOp.MAX)
if rank == 0:
    step = statistics.median(s.seconds for s in history.steps[2:])
    print(
        f"{step:.2f} s a step ({world * TOKENS / step:.0f} positions/s on {world} GPUs), "
        f"peak {float(peak):.1f} GiB a GPU"
    )
if args.out is not None:
    # Every process calls it: the parts are gathered, and the first writes.
    model.save_weights(args.out, dtype=torch.bfloat16)
torch.distributed.destroy_process_group()

Browse on GitHub

Released under the MIT License.