Most production feedback is not a pair of answers with a preference between them. It is a thumbs-up or thumbs-down on one answer, a support agent marking a draft as usable, or a test that passes or fails. Paired methods such as DPO need two completions for the same prompt and a judgement between them, so teams either throw away most of their feedback or pay annotators to build pairs. KTO, Kahneman-Tversky Optimization, introduced by Ethayarajh, Xu, Muennighoff, Jurafsky and Kiela in 2024, trains directly on unpaired binary labels. The paper reports that it matches or exceeds preference-based methods at scales from 1B to 30B parameters.
This article explains the loss from first principles, shows exactly how the reference point is estimated and why that estimate dictates your batch size and data loader, works out what the method costs on GPUs, and gives a training setup with Hugging Face TRL, whose documentation and source are the reference for every configuration name used here. It assumes you know supervised fine-tuning; DPO training is useful background for the implicit-reward idea.
The idea: utility around a reference point
Prospect theory, the work of Daniel Kahneman and Amos Tversky, describes how people value outcomes: relative to a reference point, with losses weighted more heavily than equal gains, and with diminishing sensitivity far from the reference. The KTO authors argue that successful alignment losses such as DPO already behave like such "human-aware losses", and propose one that maximises the utility of individual generations directly instead of the likelihood of pairwise preferences.
In practical terms: for each example the model has an implicit reward, beta times the log ratio of the policy's probability for the completion to the reference model's probability. A desirable completion should have a reward above a reference point, an undesirable one below it. Each is pushed through a sigmoid, which saturates, so examples that are already far on the correct side contribute almost no gradient. That saturation is the diminishing sensitivity, and the separate weights for desirable and undesirable examples are where loss aversion and data imbalance are expressed.
The loss, exactly
For a prompt x and completion y, let r be log pi_theta(y|x) minus log pi_ref(y|x), summed over completion tokens. With reference point z0, the value is sigma(beta (r - z0)) for desirable completions and sigma(beta (z0 - r)) for undesirable ones, and the loss is the weighted mean of 1 minus the value. TRL implements Eq. 7 of the paper like this:
# The KTO loss as TRL computes it (loss_type="kto"), per batch.
# logratio = log pi_theta(y|x) - log pi_ref(y|x), summed over completion tokens.
kl = (kl_policy_logps - kl_ref_logps).mean().detach() # mismatched pairs
kl = gather_across_ranks(kl).mean().clamp(min=0) # one reference point for the step
chosen_losses = 1 - sigmoid(beta * (chosen_logratios - kl)) # label True
rejected_losses = 1 - sigmoid(beta * (kl - rejected_logratios)) # label False
loss = cat(desirable_weight * chosen_losses,
undesirable_weight * rejected_losses).mean()Beta plays the same role as in DPO: larger beta keeps the policy closer to the reference model. The reference point z0 is an estimate of the KL divergence between the policy and the reference, so a completion only counts as "rewarded" if the policy raised its probability by more than it raised probabilities on average. Without z0, the model could satisfy desirable examples simply by drifting everywhere.
How the reference point is estimated
The true KL divergence would require sampling from the policy, which is expensive. KTO instead uses completions it already has: each prompt is paired with a completion that belongs to a different prompt, and the average log ratio over those mismatched pairs estimates how much the policy has shifted probability mass in general. In TRL the mismatched completion for example i is the completion of example i - 1 within the same batch, a rotation by one position.
TRL then detaches the estimate so no gradient flows through it, averages it across all data-parallel ranks, and clamps it at zero. Three practical rules follow, and TRL enforces or documents each. The sampling strategy must be sequential, because the pairing is precomputed against neighbours in a fixed order; KTOConfig defaults train_sampling_strategy to sequential for this reason. The per-device batch size must be greater than 1, otherwise there is no neighbour. And TRL's guidance is to use a per-step batch of at least 4 and an effective batch between 16 and 128, because a tiny per-step batch gives a noisy reference point regardless of how much you accumulate.
What it costs on the GPU
Count forward passes. For every example, the policy processes the matched sequence and the mismatched KL sequence, and the reference model processes both as well. That is four sequence forwards per example, compared with roughly two policy and two reference forwards per pair in DPO, which however covers two completions. Only the matched policy forward is differentiated, but whether the KL forward keeps activations depends on the implementation, so measure peak memory rather than assuming.
Memory is dominated, as in any full fine-tune, by the policy's weights, gradients and optimizer states. A rough estimate for full fine-tuning with AdamW in mixed precision is about 16 bytes per parameter before activations, around 130 GB for an 8B model, which is why multi-GPU sharding with FSDP or DeepSpeed ZeRO is normal at that size. A frozen reference model adds about 2 bytes per parameter in bf16, roughly 16 GB for 8B. Two settings remove most of that overhead. With precompute_ref_log_probs=True TRL computes reference log probabilities for the whole dataset before training, so the reference model does not stay in memory; it cannot be combined with sync_ref_model or streaming datasets. With a PEFT adapter, training only the adapter shrinks gradients and optimizer state to the adapter's size, and TRL can use the base model with adapters disabled as the reference.
Sequence length matters twice because of the KL pairs. KTOConfig's max_length (default 1024) truncates prompt plus completion from the right; older TRL versions exposed separate prompt and completion length settings, so check the configuration of the version you install. Gradient checkpointing is on by default in KTOConfig, which trades extra compute for lower activation memory.
Preparing the data
KTO wants one row per judgement: a prompt, a completion and a boolean label, in plain text or conversational message format. TRL also accepts paired preference data and converts it by splitting each pair into a desirable and an undesirable row, which is useful for comparing KTO with DPO on the same data.
from datasets import Dataset, load_dataset
# 1. Raw thumbs feedback -> KTO's unpaired format.
rows = [
{"prompt": [{"role": "user", "content": "Summarise ticket 8812"}],
"completion": [{"role": "assistant", "content": "Customer cannot reset password ..."}],
"label": True},
{"prompt": [{"role": "user", "content": "Refund policy for annual plans?"}],
"completion": [{"role": "assistant", "content": "All refunds are instant."}], # wrong
"label": False},
]
ds = Dataset.from_list(rows)
# 2. Balance the loss, not the data. TRL's guidance: keep
# (desirable_weight * n_desirable) / (undesirable_weight * n_undesirable) between 1 and 4/3.
n_pos = sum(ds["label"]); n_neg = len(ds) - n_pos
def check_ratio(w_d, w_u):
r = (w_d * n_pos) / (w_u * n_neg)
assert 1.0 <= r <= 4 / 3, f"weighted ratio {r:.2f} outside [1, 1.33]"
return rDecide what a label means before collecting it. "The user clicked thumbs-up" is noisy; "a reviewer judged the answer correct and complete" is a label. Remove prompts that appear with contradictory labels for the same completion, deduplicate near-identical completions, and keep a held-out set of prompts for evaluation. The real-world appeal of KTO is that undesirable examples are cheap to collect from production failures; make sure they reflect the failures you actually want to remove.
Worked example: imbalanced support data
Suppose a support assistant has 9,000 answers marked desirable and 3,000 marked undesirable. With both weights at 1, desirable examples dominate the loss three to one. TRL's guidance is to upweight the rarer class so that desirable_weight x positives divided by undesirable_weight x negatives lands between 1 and 4/3. Keeping desirable_weight at 1.0, the undesirable weight must be between 2.25 (ratio 9000 / 6750 = 1.33) and 3.0 (ratio exactly 1). Choose 2.5, which gives 9000 / 7500 = 1.2.
With 8 GPUs, a per-device batch of 8 and no accumulation, the effective batch is 64, inside the recommended 16 to 128 range. On 2 GPUs you would accumulate 4 steps to reach the same 64. Recompute the effective batch whenever the GPU count changes.
from datasets import load_dataset
from peft import LoraConfig
from trl import KTOConfig, KTOTrainer
train = load_dataset("json", data_files="kto_train.jsonl", split="train") # 9,000 True / 3,000 False
args = KTOConfig(
output_dir="support-8b-kto",
beta=0.1,
desirable_weight=1.0,
undesirable_weight=2.5, # (1.0 * 9000) / (2.5 * 3000) = 1.2
learning_rate=5e-7, # TRL: at most about 1e-6 at beta 0.1
per_device_train_batch_size=8, # must be above 1; at least 4 for a usable KL estimate
gradient_accumulation_steps=1, # 8 GPUs x 8 = 64 effective, inside 16-128
num_train_epochs=1,
max_length=2048,
precompute_ref_log_probs=True, # reference model not kept in GPU memory
bf16=True,
logging_steps=10,
report_to="none",
)
trainer = KTOTrainer(
model="my-org/support-8b-sft",
args=args,
train_dataset=train,
peft_config=LoraConfig(r=16, lora_alpha=32, target_modules="all-linear"),
)
trainer.train()The learning rate follows TRL's documented guidance: with the default beta of 0.1 the rate should typically not exceed 1e-6, and the recommended range is 5e-7 to 5e-6 even for small datasets, with more epochs rather than a higher rate when data is scarce. Lower beta needs a lower rate. Adapter training typically tolerates a higher rate than full fine-tuning, so treat 5e-7 here as a conservative start.
Reading the metrics
TRL logs kl, rewards/chosen, rewards/rejected, rewards/margins, logps/chosen, logps/rejected, grad_norm and the loss. The rewards are the implicit beta-scaled log ratios. A healthy run shows the chosen reward rising, the rejected reward falling and the margin widening gradually, while kl stays small and grows slowly.
Warning signs: kl climbing fast means the policy is drifting, often from a learning rate too high for the beta. Both chosen and rejected log probabilities falling together means the model is lowering likelihood on everything, a pattern also seen with DPO, and it often precedes degraded generations; lower the learning rate or raise beta. A margin that never moves suggests labels the model cannot distinguish, or a weight imbalance overwhelming one class. Metrics are not the goal: evaluate generations on held-out prompts with a reward model or human review, and compare against the SFT starting point.
Variants and choosing a method
TRL's KTOTrainer also offers loss_type="apo_zero_unpaired", the unpaired variant of APO-zero, which raises the likelihood of desirable completions and lowers undesirable ones without estimating the KL term. It drops the mismatched forward passes and the batch-size constraint, and TRL suggests it when you believe the desirable completions are better than the model's own default outputs.
| Method | Data | Extra models | Pick it when |
|---|---|---|---|
| KTO | Unpaired binary labels | Reference (or precomputed log probs) | Feedback is thumbs-up/down or pass/fail |
| DPO | Pairs with a preference | Reference | You have or can build clean pairs |
| ORPO | Pairs | None | You want SFT and preference in one stage |
| PPO / GRPO | Prompts plus a reward signal | Reward model or verifier, sampling | You can score fresh samples online |
Comparisons: ORPO training for reference-free pairs, GRPO for online reinforcement learning with verifiable rewards, and the RLHF pipeline for where an offline method such as KTO sits in a full post-training stack.
Failure modes
| Symptom | Cause | Fix |
|---|---|---|
| Trainer refuses to start | Batch size 1 or random sampling with loss_type kto | Per-device batch at least 4; keep sequential sampling |
| kl noisy, training unstable | Tiny per-step batch | Raise per-device batch; accumulation does not fix it |
| Generations degrade, all log probs fall | Learning rate too high for beta | Stay within 5e-7 to 5e-6; lower rate or raise beta |
| Model ignores undesirable examples | Imbalance not weighted | Set weights so the ratio is 1 to 4/3 |
| Out of memory | Policy and reference both resident | precompute_ref_log_probs, PEFT, FSDP or ZeRO |
| Answers truncated mid-sentence after training | max_length cut completions | Measure lengths; raise max_length or filter |
Trade-offs
- Data cost versus signal: binary labels are cheap and plentiful but carry less information per example than a direct comparison.
- Offline versus online: KTO learns from fixed data, so it cannot discover behaviours absent from the dataset the way online methods can.
- Precomputing reference log probs saves memory but fixes the reference for the whole run and adds a full pass before training starts.
What to do next
- Export your production feedback into prompt, completion and label rows, and define in writing what a desirable label means.
- Count desirable and undesirable examples and compute weights that put the weighted ratio between 1 and 4/3.
- Run a small KTO job with a 0.5B to 1.5B model and the defaults, per-device batch of at least 4, to learn the metrics.
- Measure peak GPU memory with and without precompute_ref_log_probs before scaling up.
- Watch kl, the reward margin and both log probabilities; stop and lower the learning rate if everything falls.
- Evaluate on held-out prompts against the SFT baseline, and against DPO if you can build pairs.