在上一节中,我们理解了Datum数据类型,及其在forward_backward函数中的作用原理。
本节我们来看PyTRIO中的三个内置损失函数,以及如何定制自己的损失函数。
内置损失函数
PyTRIO 为 sft 和 rl 提供了内置的损失函数。
内置损失函数全程都在后台GPU上计算,相比自定义损失函数,速度上要更快。
可以通过将字符串传递给forward_backward的loss_fn参数来选择损失函数:
future=training_client.forward_backward(data,loss_fn="cross_entropy",)在上面的代码中,就是使用了cross_entropy(交叉熵损失)。
目前PyTRIO提供了三种内置损失函数:
| 损失函数 | 适用场景 | 说明 |
|---|---|---|
cross_entropy | 监督学习 | 标准交叉熵损失,适用于分类任务。以模型输出的 logits 和目标标签计算负对数似然。 |
importance_sampling | 离线强化学习 | 使用重要性采样对 off-policy 数据进行修正,通过行为策略与目标策略的概率比值对梯度加权。 |
ppo | 在线强化学习 | Proximal Policy Optimization 损失,通过裁剪概率比值限制策略更新幅度,提升训练稳定性。 |
它们的详细分析可以看官方文档:https://docs.pytrio.com/docs/guide/loss_fn
损失函数的返回值
一次forward_backward计算后,会得到两类返回值:
loss_fn_outputs:对batch中每个样本的损失函数计算中间值,比如logprobs、elementwise_loss等,可以用于合成各类高阶指标。metrics:反应训练情况的指标,比如loss_sum、loss_mean等,常用于打印到终端或记录到swanlab、tensorboard、wandb等训练可观测平台。
future=training_client.forward_backward(data,loss_fn="cross_entropy",)print(fwdbwd_result.loss_fn_outputs)print(fwdbwd_result.metrics)这里举个例子,比如你希望打印这个batch的平均loss,可以:
print(fwdbwd_result.metrics["loss_mean"])自定义损失函数
对于内置损失函数满足不了的场景,PyTRIO 提供了更灵活的forward_backward_custom实现定制化损失函数。
forward_backward_custom的输入参数是data和loss_fn,在参数名上和forward_backward一样,区别在于forward_backward_custom的loss_fn传入的是一个自己实现的损失函数。
损失函数有自己的定义规范:
defloss_fn_custom(data:list[trio.Datum],logprobs:list[torch.Tensor])->tuple[torch.Tensor,dict[str,float]]:...returnloss,metrics我们来解读一下。首先入参是**data和logprobs****:**
data:一个由Datum组成的列表logprobs:由 PyTRIO 自动对Datum中的target_tokens做前向传播计算得到的logprobs(负对数概率)列表。
返回值是:
loss:损失值,是一个标量metrics:一个字典,用于放一些指标便于打印(可以被forward_backward计算结果的metrics字段拿到)
**下面举个例子。**比如我们希望实现这样一个损失函数,逻辑是希望每个 logprob 尽可能接近 0(也就是概率接近 1),公式为:
实现的代码为:
deflogprob_squared_loss(data:list[trio.Datum],logprobs:list[torch.Tensor])->tuple[torch.Tensor,dict[str,float]]:flat_logprobs=torch.cat(logprobs)loss=(flat_logprobs**2).sum()returnloss,{"logprob_squared_loss":loss.item()}将这个损失函数传入到forward_backward_custom中,并打印metrics:
future=training_client.forward_backward_custom(data,logprob_squared_loss)result=future.result()print(f"Loss:{result.loss}, Metrics:{result.metrics}")让我们改造一个「第一节」中的sft案例为自定义损失函数:
importpytrioastrioimporttorch# 1. 与TRIO建立连接service_client=trio.ServiceClient()# 2. 创建1个训练客户端base_model="Qwen/Qwen3.5-4B"training_client=service_client.create_lora_training_client(base_model=base_model,rank=32,)# 3. 数据集-让LLM答对什么是trioexamples=[{"input":"what is trio","output":"trio is emotionmachine's AI Infra products."},{"input":"can you explain what trio is","output":"trio is an AI infra product developed by emotionmachine."},{"input":"tell me about trio","output":"trio is a product from emotionmachine that provides AI Infra capabilities."},]# 4. 获取Tokenizerprint("Loading tokenizer...")tokenizer=training_client.get_tokenizer()print("Tokenizer finish")# 5. 处理数据集,转换为训练需要的格式defprocess_example(example:dict,tokenizer)->trio.Datum:prompt=f"Question:{example['input']}\nAnswer:"prompt_tokens=tokenizer.encode(prompt,add_special_tokens=True)prompt_weights=[0]*len(prompt_tokens)completion_tokens=tokenizer.encode(f"{example['output']}\n\n",add_special_tokens=False)completion_weights=[1]*len(completion_tokens)tokens=prompt_tokens+completion_tokens weights=prompt_weights+completion_weights input_tokens=tokens[:-1]target_tokens=tokens[1:]weights=weights[1:]# 转换为trio训练需要的格式returntrio.Datum(model_input=trio.ModelInput.from_ints(tokens=input_tokens),loss_fn_inputs=dict(weights=weights,target_tokens=target_tokens))processed_examples=[process_example(ex,tokenizer)forexinexamples]# 6. 自定义损失函数deflogprob_squared_loss(data:list[trio.Datum],logprobs:list[torch.Tensor])->tuple[torch.Tensor,dict[str,float]]:flat_logprobs=torch.cat(logprobs)loss=(flat_logprobs**2).sum()returnloss,{"logprob_squared_loss":loss.item()}# 7. 训练print("Start Training")foriterinrange(15):fwdbwd_future=training_client.forward_backward_custom(processed_examples,logprob_squared_loss)optim_future=training_client.optim_step(trio.AdamParams(learning_rate=1e-4))fwdbwd_result=fwdbwd_future.result()optim_result=optim_future.result()print(f"Iter{iter+1}Logprob_squared_loss:{fwdbwd_result.metrics['logprob_squared_loss']:.4f}")# 7. 推理与评估print("Start Sampling")sampling_base_client=service_client.create_sampling_client(base_model=base_model)training_client.save_state(name="Train")sampling_sft_client=training_client.save_weights_and_get_sampling_client(name='what-is-trio')prompt=trio.ModelInput.from_ints(tokenizer.encode("Question: what is trio\nAnswer:"))params=trio.SamplingParams(max_tokens=20,temperature=0.0,stop=["\n"])future_base=sampling_base_client.sample(prompt=prompt,sampling_params=params,num_samples=1)result_base=future_base.result()future_sft=sampling_sft_client.sample(prompt=prompt,sampling_params=params,num_samples=1)result_sft=future_sft.result()print("Base Responses:")print(f"{repr(result_base.sequences[0].text)}")print("SFT Responses:")print(f"{repr(result_sft.sequences[0].text)}")运行后的输出结果如下:
Iter1 Logprob_squared_loss: 2173.7051 ... Iter15 Logprob_squared_loss: 48.7835 Start Sampling Base Responses: ' A trio is a musical ensemble consisting of three performers. The term can also refer to a group of' SFT Responses: ' trio is emotionmachine emotionmachine emotionmachine emotionmachine emotionmachine emotionmachine emotionmachine emotionmachine emotionmachine'可以看到,本次训练中自定义损失函数已经产生了作用。
ps:
logprob_squared_loss只是个用于示例的损失函数,实际效果并不好,请勿使用到自己的训练中。
这时聪明的读者可能发现了一个小问题,在Datum类型中有个loss_fn_inputs参数,在第二节中我们做 sft 时,会传入包含weights和target_tokens的字典,参与到损失的计算。
那么在自定义 loss_fn 中,要如何调用这些参数呢?
方法其实也很简单。下面是用自定义 loss_fn 实现的交叉熵损失:
defcustom_cross_entropy_loss(data,logprobs):total_loss=0.0total_weight=0.0fordatum,token_logprobsinzip(data,logprobs):weights=torch.as_tensor(datum.loss_fn_inputs["weights"].data,dtype=token_logprobs.dtype,device=token_logprobs.device,)# token_logprobs: 模型对 target_tokens 的逐 token log p# cross entropy = -log p(target)total_loss=total_loss+-(token_logprobs*weights).sum()total_weight=total_weight+weights.sum()loss=total_loss/total_weight.clamp_min(1.0)returnloss,{"custom_cross_entropy/loss":float(loss.detach().item()),"custom_cross_entropy/tokens":float(total_weight.detach().item()),}可以看到,我们可以通过datum.loss_fn_inputs函数来拿到这些参数,执行计算。