本节我们将了解 PyTRIO 的数据类型Datum,以及提供的三种内置损失函数。

一、了解Datum
我们已经知道,PyTRIO执行训练依靠的是在循环中将数据一轮轮传递给forward_backward,来计算梯度。
而在数据传入forward_backward之前,需要先做一层封装,而这个封装就是Datum格式。
为了更方便地理解Datum,我们先从 sft 的数据集与损失函数的关系开始。
下面是一个经典的 sft 数据集格式:
| input | output |
|---|---|
| 刘汉宏是谁? | 刘汉宏是唐朝末期的军阀之一,主要成就是在唐末时担任义胜军节度使,为其领地的经济和军事发展做出了巨大贡献。 |
| Python是什么? | Python是一种面向对象的计算机程序设计语言,语法简洁清晰,且具有丰富和强大的类库。它能够很轻松地把用其他语言制作的各种模块轻松地联结在一起,被广泛应用于各种领域,包括Web开发、人工智能、数据处理、网络编程等。 |
可以看到,数据分为两部分:input 和 output,分别是模型的输入和我们预期的模型输出。
那 sft 中 交叉熵损失 是如何工作的呢?
**首先我们需要将数据集构建成LLM可接受的的输入和输出序列。**我们知道,LLM是一种自回归模型,在序列构建上是通过将system_prompt和数据集中的input与output组合成一个长序列,然后错开一位来实现的:

同时,在sft训练中,我们希望只训练序列的output部分 ------ 即只在output部分计算loss,而prompt部分不计算。
所以,还有一个weights参数,它一般是一个由 0 和 1 组成的向量,0 代表不需要被训练的 token,1 代表需要被训练的 token。在sft中,经常的做法是让 prompt 部分为0,output部分为 1 。
得到 **input_token、 target_token和 weights**之后,就能计算交叉熵损失:

总结来说,对一次sft的loss计算而言,我们只需集齐上述的3个组件即可。
我们再来看**Datum****,**这下就很好看懂了。
构建一个Datum的代码如下:
Python
datum = trio.Datum(
model_input=trio.ModelInput.from_ints(tokens=input_tokens),
loss_fn_inputs=dict(
weights=weights,
target_tokens=target_tokens,
)
)
Datum由两部分组成:
-
model_input:即input_token,用于给到 LLM 生成predict_token -
loss_fn_inputs:损失函数的其他输入参数,在sft中也就是需要weights和target_tokens
这样,就把一条数据的Datum构建出来了!
而对于一个数据集来说,就是把每一条数据都变成Datum格式,构建一个Datum列表:
Python
def process_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:]
# 转换为Datum格式
return trio.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) for ex in examples]
最后,把Datum列表输入到forward_backward中:
Python
fwdbwd_future = training_client.forward_backward(
processed_examples,
"cross_entropy"
)

当然,这样相当于把整个数据集作为一个batch输入给了forward_backward当中。
更推荐的做法是切片后分batch输入:
Python
batch_size=16
for i in range(12):
start_index=i*batch_size
fwdbwd_future = training_client.forward_backward(
processed_examples[start_index: start_index+batch_size],
"cross_entropy"
)
二、强化学习里的Datum
上面我们介绍了 sft 中的 Datum 应该如何构建。
可以看出,Datum 的构建逻辑是围绕损失函数 的。不同的损失函数,有不同的loss_fn_inputs。
而rl和sft的损失函数不同,也注定了在强化学习中Datum的构建方式有些区别。
RL 里的一条 Datum 通常不是原始数据集里的一条问答,而是LLM自己采样出来的一条 rollout 轨迹。
我们首先需要将prompt给到LLM,得到推理后的结果rollout_token和对应的logprobs:

然后将prompt和rollout_token拼接成一个完整序列后,错位得到input_token和target_token:

另外,我们根据rollout_token结合奖励函数,可以计算出优势值advantage,这样就把RL中loss函数(重要性采样)需要的组件凑齐了:

ok,我们来看看重要性采样(importance_sampling)的Datum构建代码,应该很好理解了:
Python
datum = trio.Datum(
model_input=trio.ModelInput.from_ints(tokens=input_tokens),
loss_fn_inputs=dict(
target_tokens=target_tokens,
logprobs=logprobs,
advantages=advantages,
),
)
在实际的RL训练循环中,我们只需要在每个step中,就LLM sample的结果组成一个Datum列表,传入forward_backward计算即可。
Python
fwdbwd_future = training_client.forward_backward(
processed_examples,
"importance_sampling"
)
三、实战案例
我们可以看几个实际的训练代码,来更深入地学会Datum的用法:
-
快速开始:快速开始 - TRIO官方文档
-
sft案例-Chat甄嬛:Chat-甄嬛 - TRIO官方文档
-
rl案例-GRPO:GRPO - TRIO官方文档