PyTRIO快速入门(三):损失函数
2026/7/30 12:14:26 网站建设 项目流程

在上一节中,我们理解了Datum数据类型,及其在forward_backward函数中的作用原理。

本节我们来看PyTRIO中的三个内置损失函数,以及如何定制自己的损失函数。

内置损失函数

PyTRIO 为 sft 和 rl 提供了内置的损失函数。

内置损失函数全程都在后台GPU上计算,相比自定义损失函数,速度上要更快。

可以通过将字符串传递给forward_backwardloss_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计算后,会得到两类返回值:

  1. loss_fn_outputs:对batch中每个样本的损失函数计算中间值,比如logprobs、elementwise_loss等,可以用于合成各类高阶指标。

  2. 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的输入参数是dataloss_fn,在参数名上和forward_backward一样,区别在于forward_backward_customloss_fn传入的是一个自己实现的损失函数。

损失函数有自己的定义规范:

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

我们来解读一下。首先入参是**datalogprobs****:**

  • 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 时,会传入包含weightstarget_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函数来拿到这些参数,执行计算。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询