指南

自定义损失函数

对于内置损失函数之外的场景,用户还可以选择更灵活的自定义损失函数:将手动实现的损失函数传入 forward_backward_custom 方法来计算损失和其他指标。

自定义损失函数始终在本地 Python 进程中执行。PyTRIO 从服务端取得当前模型的逐 Token logprob,在本地调用用户定义的函数,再把损失对 logprob 的梯度交给服务端完成模型参数的反向传播。

自定义损失函数通常执行得更慢,建议尽量使用异步方法 forward_backward_custom_async,避免阻塞训练流程。异步方法的介绍见异步

用法

定义损失函数

首先需要定义一个损失函数,函数签名为:

def logprob_squared_loss(
    data: list[trio.Datum],
    logprobs: list[torch.Tensor],
) -> tuple[torch.Tensor, dict[str, float]]:
    ...

其中:

  • data 是输入数据列表,每个元素都是一个 Datum 对象,顺序与调用 forward_backward_custom 时传入的数据一致;
  • logprobs 是当前模型前向传播输出的逐 Token 对数概率列表,与 data 一一对应,每个张量的长度与对应 Datum.model_input 相同;

准备 Datum 与额外数据

自定义损失可能同时使用三类数据:

数据来源或存放位置示例
服务端前向所需的数据Datum.model_inputDatum.loss_fn_inputs输入 Token、target_tokens
当前模型的输出PyTRIO 传给损失函数的 logprobs当前策略的逐 Token logprob
仅供自定义算法使用的本地数据独立的 Python 数据结构,通过闭包传入sampling/reference logprob、sequence advantage、completion 长度、分组关系、归一化参数

loss_fn_inputs 是服务端损失函数的强类型张量输入,只用于保存对应损失 schema 规定的字段。在 forward_backward_custom 中,用户只需传入与 model_input 等长的 target_tokens。自定义键、标量、dataclass 或其他 Python 对象不能放在这里。

当损失函数还需要本地辅助数据时,可以先定义一个工厂函数,再用闭包把这些数据绑定到最终的二参数损失函数中。下面以简化的 sequence-level objective 为例:

from collections.abc import Callable
from dataclasses import dataclass


@dataclass(frozen=True)
class SequenceLossMeta:
    sampling_logprobs: list[float]
    advantage: float
    completion_tokens: int


def make_sequence_loss_fn(
    metas: list[SequenceLossMeta],
) -> Callable[
    [list[trio.Datum], list[torch.Tensor]],
    tuple[torch.Tensor, dict[str, float]],
]:
    if not metas:
        raise ValueError("metas must not be empty")

    # 固定当前 batch 的顺序快照。
    batch_metas = tuple(metas)

    def sequence_loss_fn(
        data: list[trio.Datum],
        logprobs: list[torch.Tensor],
    ) -> tuple[torch.Tensor, dict[str, float]]:
        if not (len(data) == len(logprobs) == len(batch_metas)):
            raise ValueError("data, logprobs and metas must have the same length")

        objectives = []
        for meta, current in zip(batch_metas, logprobs, strict=True):
            if meta.completion_tokens <= 0:
                raise ValueError("completion_tokens must be positive")
            if current.numel() < meta.completion_tokens:
                raise ValueError("logprob sequence is shorter than the completion")

            current_completion = current[-meta.completion_tokens :].float()
            sampling = torch.as_tensor(
                meta.sampling_logprobs,
                dtype=current_completion.dtype,
                device=current_completion.device,
            )
            if current_completion.numel() != sampling.numel():
                raise ValueError("sampling logprobs must match completion tokens")

            sequence_ratio = torch.exp((current_completion - sampling).mean())
            objectives.append(sequence_ratio * meta.advantage)

        loss = -torch.stack(objectives).mean()
        return loss, {"sequence_loss": float(loss.detach().item())}

    return sequence_loss_fn

这里返回的 sequence_loss_fn 仍然符合 PyTRIO 要求的固定签名,batch_metas 则保留在本地闭包中:

loss_fn = make_sequence_loss_fn(batch_metas)
future = training_client.forward_backward_custom(
    data=batch_data,
    loss_fn=loss_fn,
)
result = future.result()

batch_data[i]batch_metas[i] 和 PyTRIO 返回的 logprobs[i] 必须描述同一条样本。每个 batch 应创建自己的闭包,并在调用前完成 reference logprob 等数据的计算,不要在损失函数内部发起网络请求。

完整的 GSPO sequence-level clipping 实现可以参考 07-gspo/loss.py,其中的 GSPOMetamake_gspo_loss_fn 使用了相同的数据拆分方式。

调用损失函数

定义好损失函数后,在 forward_backward_custom 方法中传入损失函数即可:

future = training_client.forward_backward_custom(
    data=data,
    loss_fn=your_custom_loss_fn,
)
result = future.result()

推荐使用对应的异步方法:

future = await training_client.forward_backward_custom_async(
    data=data,
    loss_fn=your_custom_loss_fn,
)
result = await future

loss = result.get("metrics", {}).get("loss:sum", 0.0)
print(loss)

示例

下面是一个简单的示例:定义一个 logprob_squared_loss 函数,计算逐 Token 对数概率的平方和,并将其作为损失进行优化。

import asyncio

import swanlab
import torch
from datasets import load_dataset

import pytrio as trio

EPOCHS = 10
BATCH_SIZE = 2


# 自定义损失函数:逐 Token 对数概率的平方和
def logprob_squared_loss(
    _: list[trio.Datum], logprobs: list[torch.Tensor]
) -> tuple[torch.Tensor, dict[str, float]]:  
    flat = torch.cat([x.reshape(-1) for x in logprobs])
    loss = (flat**2).sum()
    return loss, {"logprob_squared_loss": float(loss.detach().item())}


async def main():
    # 1. 与 TRIO 建立连接
    service_client = trio.ServiceClient()

    # 2. 创建一个 LoRA 训练客户端
    training_client = await service_client.create_lora_training_client_async(
        base_model="Qwen/Qwen3-4B-Instruct-2507",
        seed=42,
        train_mlp=True,
        train_attn=True,
        train_unembed=False,
    )
    tokenizer = training_client.get_tokenizer()

    # 3. 加载数据集,并转换为训练需要的格式
    dataset = load_dataset(
        "HuggingFaceTB/smoltalk",
        "everyday-conversations",
        split="train[:10]",
    )

    all_samples: list[trio.Datum] = []

    for example in dataset:
        text = tokenizer.apply_chat_template(example["messages"], tokenize=False)
        tokens = tokenizer.encode(text, add_special_tokens=False)

        input_ids = tokens[:-1]
        target_tokens = tokens[1:]

        all_samples.append(
            trio.Datum(
                model_input=trio.ModelInput.from_ints(input_ids),
                loss_fn_inputs={
                    "target_tokens": target_tokens,
                },
            )
        )

    batches = [
        all_samples[i : i + BATCH_SIZE] for i in range(0, len(all_samples), BATCH_SIZE)
    ]

    # 4. 初始化 SwanLab,记录训练指标
    swanlab.init(
        project="trio-custom-loss",
        experiment_name="logprob-squared-loss",
    )

    # 5. 训练
    for epoch in range(EPOCHS):
        futures = []

        for batch in batches:
            future = await training_client.forward_backward_custom_async(  
                data=batch,
                loss_fn=logprob_squared_loss,  
            )
            tokens_in_batch = sum(len(d.loss_fn_inputs["target_tokens"]) for d in batch)
            futures.append((future, tokens_in_batch))

            await training_client.optim_step_async(
                trio.AdamParams(
                    learning_rate=1e-4,
                    beta1=0.9,
                    beta2=0.999,
                    eps=1e-8,
                    weight_decay=0,
                )
            )

        for future, tokens_in_batch in futures:
            result = await future
            loss_sum = float(result.get("metrics", {}).get("loss:sum", 0.0))
            loss_mean = loss_sum / tokens_in_batch
            swanlab.log({"train/loss": loss_mean})

        print(f"Epoch {epoch} completed")


if __name__ == "__main__":
    asyncio.run(main())

注意:logprob_squared_loss 只是一个用于示例的损失函数,实际训练效果并不好,请勿使用到自己的训练中。

原理

模型的前向计算图位于服务端,用户自定义的损失计算图则在本地 Python 进程中创建,两者之间没有一张跨越网络的连续计算图。PyTRIO 因此需要分两段应用链式法则,最终得到损失对模型参数的导数:

lossθ\frac{\partial \text{loss}}{\partial \theta}

训练器默认的 Cross Entropy Loss 的计算过程如下:

loss_elementwise = -logprobs * weights
loss = loss_elementwise.sum()

在一般的学习框架中,通常取 wi[0,1]w_i \in [0, 1] 来表示每个位置的损失权重,即:

  • wi=0w_i = 0 时,该位置的损失不计算;
  • wi=1w_i = 1 时,该位置的损失正常计算;
  • wi(0,1)w_i \in (0, 1) 时,该位置的损失会被缩放。

损失对参数 θ\theta 的梯度为:

lossθ=i(logpiwi)θ\frac{\partial \text{loss}}{\partial \theta} = \sum_i \frac{\partial (-\log p_i \cdot w_i)}{\partial \theta}

如果将 wiw_i 视为常数(即 wiw_i 不依赖于 θ\theta),那么:

lossθ=iwilogpiθ\frac{\partial \text{loss}}{\partial \theta} = \sum_i -w_i \frac{\partial \log p_i}{\partial \theta}

在自定义损失流程中,SDK 会根据本地计算的 logprob 梯度生成代理权重,这里记为 w~i\tilde{w}_i

w~i=LlogpiR\tilde{w}_i = -\frac{\partial L}{\partial \log p_i} \in \mathbb{R}

将它代入 Cross Entropy 形式后,有:

L~θ=iLlogpilogpiθ\frac{\partial \tilde{L}}{\partial \theta} = \sum_i \frac{\partial L}{\partial \log p_i} \frac{\partial \log p_i}{\partial \theta}

这正是链式法则:

Lθ=iLlogpilogpiθ\frac{\partial L}{\partial \theta} = \sum_i \frac{\partial L}{\partial \log p_i} \frac{\partial \log p_i}{\partial \theta}

其中 logpiθ\frac{\partial \log p_i}{\partial \theta} 在服务端计算,Llogpi\frac{\partial L}{\partial \log p_i} 则由本地 PyTorch autograd 计算并传回服务端。

换句话说,这相当于构造了一个关于 logprobs 的线性替代目标(surrogate objective):

L~(θ)=ilogpiLlogpi\tilde{L}(\theta) = \sum_i \log p_i \cdot \frac{\partial L}{\partial \log p_i}

服务端会把 Llogpi\frac{\partial L}{\partial \log p_i} 视为常数。这个 surrogate objective 的形式与原始损失不同,但本次前向反向过程中产生的模型参数梯度严格等价。

SDK 内部生成的代理权重 w~i\tilde{w}_i 不受 [0,1][0, 1] 限制,也可能为负数。它只用于把本地算出的 logprob 梯度交回服务端,不能用来承载用户的额外元数据。

执行流程

forward_backward_custom 在客户端与服务器之间分两阶段完成梯度计算:

  1. 准备数据:客户端构造 Datum 对象列表,并准备目标 Token;算法需要的其他本地数据可以绑定在损失函数闭包中。
  2. 前向计算:服务器执行一次 forward,计算目标 token 的 logprobs。
  3. 客户端计算自定义损失:客户端把返回值重建为可求导的 PyTorch 张量,再调用用户定义的 loss_fn(data, logprobs);闭包中的额外数据也在这一步参与计算。
  4. 客户端反向传播到 logprobs:客户端对该损失执行反向传播,得到 Llogprobs\frac{\partial L}{\partial \text{logprobs}},即每个 logprob 对最终损失的梯度。
  5. 服务器执行 surrogate forward-backward:服务器使用这些梯度作为权重,构造 surrogate loss 并对其执行 forward-backward,从而得到与原始自定义损失完全一致的参数梯度。

为什么不需要上传自定义函数

在这一设计中,服务器只需要:

  • 计算目标 token 的 logprobs;
  • 接收客户端返回的 Llogprobs\frac{\partial L}{\partial \text{logprobs}}
  • 对 surrogate objective 执行标准的梯度计算。

因此,用户定义的 Python 函数始终保留在客户端执行,PyTRIO 不会对其进行 pickle,也不会将其发送到服务器。

这篇文档对你有帮助吗?

本页目录