自定义损失函数
对于内置损失函数之外的场景,用户还可以选择更灵活的自定义损失函数:将手动实现的损失函数传入 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_input 和 Datum.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,其中的 GSPOMeta 与 make_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 因此需要分两段应用链式法则,最终得到损失对模型参数的导数:
训练器默认的 Cross Entropy Loss 的计算过程如下:
loss_elementwise = -logprobs * weights
loss = loss_elementwise.sum()在一般的学习框架中,通常取 来表示每个位置的损失权重,即:
- 当 时,该位置的损失不计算;
- 当 时,该位置的损失正常计算;
- 当 时,该位置的损失会被缩放。
损失对参数 的梯度为:
如果将 视为常数(即 不依赖于 ),那么:
在自定义损失流程中,SDK 会根据本地计算的 logprob 梯度生成代理权重,这里记为 :
将它代入 Cross Entropy 形式后,有:
这正是链式法则:
其中 在服务端计算, 则由本地 PyTorch autograd 计算并传回服务端。
换句话说,这相当于构造了一个关于 logprobs 的线性替代目标(surrogate objective):
服务端会把 视为常数。这个 surrogate objective 的形式与原始损失不同,但本次前向反向过程中产生的模型参数梯度严格等价。
SDK 内部生成的代理权重 不受 限制,也可能为负数。它只用于把本地算出的 logprob 梯度交回服务端,不能用来承载用户的额外元数据。
执行流程
forward_backward_custom 在客户端与服务器之间分两阶段完成梯度计算:
- 准备数据:客户端构造
Datum对象列表,并准备目标 Token;算法需要的其他本地数据可以绑定在损失函数闭包中。 - 前向计算:服务器执行一次 forward,计算目标 token 的 logprobs。
- 客户端计算自定义损失:客户端把返回值重建为可求导的 PyTorch 张量,再调用用户定义的
loss_fn(data, logprobs);闭包中的额外数据也在这一步参与计算。 - 客户端反向传播到 logprobs:客户端对该损失执行反向传播,得到 ,即每个 logprob 对最终损失的梯度。
- 服务器执行 surrogate forward-backward:服务器使用这些梯度作为权重,构造 surrogate loss 并对其执行 forward-backward,从而得到与原始自定义损失完全一致的参数梯度。
为什么不需要上传自定义函数
在这一设计中,服务器只需要:
- 计算目标 token 的 logprobs;
- 接收客户端返回的 ;
- 对 surrogate objective 执行标准的梯度计算。
因此,用户定义的 Python 函数始终保留在客户端执行,PyTRIO 不会对其进行 pickle,也不会将其发送到服务器。