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-rematThe 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:
| Step | Tokens/s | Peak a GPU | Held-out loss | |
|---|---|---|---|---|
| Linnet, PyTorch | 0.55 s | 29.4K | 48 GiB | 1.91 to 1.342 |
Linnet, JAX, --no-remat | 0.53 s | 30.9K | 62 GiB | 1.91 to 1.345 |
| Linnet, JAX | 0.66 s | 24.6K | 39 GiB | 1.91 to 1.344 |
| TRL, FSDP2 | 0.57 s | 28.6K | 52 GiB |
What it shows
- Sharding needs nothing in the source.
linnet.torch.fully_shardsplits 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 ownfully_shardcannot 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=Truewraps each layer injax.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_weightsgathers 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()