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.

Advertisement

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.

Advertisement

How the reference point is estimated

Unpaired batch (sequential)(prompt, completion, label) x Brotate by oneKL batchprompt i + completion i-1Policy forwardlog p_theta(y|x), matched pairsPolicy forwardmismatched KL pairsReference forwardboth sets, or precomputedLog ratios rlog p_theta - log p_refKL estimate z0mean ratio, detached, gathered, clamp at 0Lossdesirable: w_d (1 - sigma(beta (r - z0))) undesirable: w_u (1 - sigma(beta (z0 - r)))Gradients flow only through r for the matched pairs; z0 shifts the reference point but is not differentiated.
One KTO step. Each prompt is also paired with its neighbour's completion; the mismatched pairs give an estimate of the policy's average drift from the reference, which becomes the reference point z0 for every example in the step.

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 r

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

MethodDataExtra modelsPick it when
KTOUnpaired binary labelsReference (or precomputed log probs)Feedback is thumbs-up/down or pass/fail
DPOPairs with a preferenceReferenceYou have or can build clean pairs
ORPOPairsNoneYou want SFT and preference in one stage
PPO / GRPOPrompts plus a reward signalReward model or verifier, samplingYou 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

SymptomCauseFix
Trainer refuses to startBatch size 1 or random sampling with loss_type ktoPer-device batch at least 4; keep sequential sampling
kl noisy, training unstableTiny per-step batchRaise per-device batch; accumulation does not fix it
Generations degrade, all log probs fallLearning rate too high for betaStay within 5e-7 to 5e-6; lower rate or raise beta
Model ignores undesirable examplesImbalance not weightedSet weights so the ratio is 1 to 4/3
Out of memoryPolicy and reference both residentprecompute_ref_log_probs, PEFT, FSDP or ZeRO
Answers truncated mid-sentence after trainingmax_length cut completionsMeasure 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

  1. Export your production feedback into prompt, completion and label rows, and define in writing what a desirable label means.
  2. Count desirable and undesirable examples and compute weights that put the weighted ratio between 1 and 4/3.
  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.
  4. Measure peak GPU memory with and without precompute_ref_log_probs before scaling up.
  5. Watch kl, the reward margin and both log probabilities; stop and lower the learning rate if everything falls.
  6. Evaluate on held-out prompts against the SFT baseline, and against DPO if you can build pairs.
Key takeaway: KTO aligns a model from unpaired thumbs-up and thumbs-down data by pushing the implicit reward of desirable completions above, and undesirable ones below, a reference point estimated from mismatched completions in the same batch. That estimate is why sampling must stay sequential and per-device batches must be at least 4, and it adds forward passes that precomputed reference log probabilities and adapters can offset. Weight the rarer label so the weighted ratio is between 1 and 4/3, keep the learning rate between 5e-7 and 5e-6, watch kl and both log probabilities, and judge success by held-out generations rather than training metrics.