PyTRIO快速入门(三):损失函数
在上一节中我们理解了Datum数据类型及其在forward_backward函数中的作用原理。本节我们来看PyTRIO中的三个内置损失函数以及如何定制自己的损失函数。内置损失函数PyTRIO 为 sft 和 rl 提供了内置的损失函数。内置损失函数全程都在后台GPU上计算相比自定义损失函数速度上要更快。可以通过将字符串传递给forward_backward的loss_fn参数来选择损失函数futuretraining_client.forward_backward(data,loss_fncross_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等训练可观测平台。futuretraining_client.forward_backward(data,loss_fncross_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_logprobstorch.cat(logprobs)loss(flat_logprobs**2).sum()returnloss,{logprob_squared_loss:loss.item()}将这个损失函数传入到forward_backward_custom中并打印metricsfuturetraining_client.forward_backward_custom(data,logprob_squared_loss)resultfuture.result()print(fLoss:{result.loss}, Metrics:{result.metrics})让我们改造一个「第一节」中的sft案例为自定义损失函数importpytrioastrioimporttorch# 1. 与TRIO建立连接service_clienttrio.ServiceClient()# 2. 创建1个训练客户端base_modelQwen/Qwen3.5-4Btraining_clientservice_client.create_lora_training_client(base_modelbase_model,rank32,)# 3. 数据集-让LLM答对什么是trioexamples[{input:what is trio,output:trio is emotionmachines 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...)tokenizertraining_client.get_tokenizer()print(Tokenizer finish)# 5. 处理数据集转换为训练需要的格式defprocess_example(example:dict,tokenizer)-trio.Datum:promptfQuestion:{example[input]}\nAnswer:prompt_tokenstokenizer.encode(prompt,add_special_tokensTrue)prompt_weights[0]*len(prompt_tokens)completion_tokenstokenizer.encode(f{example[output]}\n\n,add_special_tokensFalse)completion_weights[1]*len(completion_tokens)tokensprompt_tokenscompletion_tokens weightsprompt_weightscompletion_weights input_tokenstokens[:-1]target_tokenstokens[1:]weightsweights[1:]# 转换为trio训练需要的格式returntrio.Datum(model_inputtrio.ModelInput.from_ints(tokensinput_tokens),loss_fn_inputsdict(weightsweights,target_tokenstarget_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_logprobstorch.cat(logprobs)loss(flat_logprobs**2).sum()returnloss,{logprob_squared_loss:loss.item()}# 7. 训练print(Start Training)foriterinrange(15):fwdbwd_futuretraining_client.forward_backward_custom(processed_examples,logprob_squared_loss)optim_futuretraining_client.optim_step(trio.AdamParams(learning_rate1e-4))fwdbwd_resultfwdbwd_future.result()optim_resultoptim_future.result()print(fIter{iter1}Logprob_squared_loss:{fwdbwd_result.metrics[logprob_squared_loss]:.4f})# 7. 推理与评估print(Start Sampling)sampling_base_clientservice_client.create_sampling_client(base_modelbase_model)training_client.save_state(nameTrain)sampling_sft_clienttraining_client.save_weights_and_get_sampling_client(namewhat-is-trio)prompttrio.ModelInput.from_ints(tokenizer.encode(Question: what is trio\nAnswer:))paramstrio.SamplingParams(max_tokens20,temperature0.0,stop[\n])future_basesampling_base_client.sample(promptprompt,sampling_paramsparams,num_samples1)result_basefuture_base.result()future_sftsampling_sft_client.sample(promptprompt,sampling_paramsparams,num_samples1)result_sftfuture_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可以看到本次训练中自定义损失函数已经产生了作用。pslogprob_squared_loss只是个用于示例的损失函数实际效果并不好请勿使用到自己的训练中。这时聪明的读者可能发现了一个小问题在Datum类型中有个loss_fn_inputs参数在第二节中我们做 sft 时会传入包含weights和target_tokens的字典参与到损失的计算。那么在自定义 loss_fn 中要如何调用这些参数呢方法其实也很简单。下面是用自定义 loss_fn 实现的交叉熵损失defcustom_cross_entropy_loss(data,logprobs):total_loss0.0total_weight0.0fordatum,token_logprobsinzip(data,logprobs):weightstorch.as_tensor(datum.loss_fn_inputs[weights].data,dtypetoken_logprobs.dtype,devicetoken_logprobs.device,)# token_logprobs: 模型对 target_tokens 的逐 token log p# cross entropy -log p(target)total_losstotal_loss-(token_logprobs*weights).sum()total_weighttotal_weightweights.sum()losstotal_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函数来拿到这些参数执行计算。

相关新闻

最新新闻

日新闻

周新闻

月新闻