"""Response-only next-token loss with exclusive sequence-length boundaries."""
import mlx.core as mx
import mlx.nn as nn


def response_loss(model, batch, lengths):
    logits = model(batch[:, :-1])
    target_positions = mx.arange(1, batch.shape[1])
    # Tokens at indices [prompt_length, sequence_length) are actual response
    # targets. sequence_length itself is the first padding token, not a target.
    response = (target_positions >= lengths[:, :1]) & (target_positions < lengths[:, 1:])
    losses = nn.losses.cross_entropy(logits, batch[:, 1:])
    count = response.sum()
    total = mx.where(response, losses, 0).astype(mx.float32).sum()
    return total / count, count
