PyTRIO快速入门(二):Datum构建

0 阅读4分钟

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

在这里插入图片描述

一、了解Datum

我们已经知道,PyTRIO执行训练依靠的是在循环中将数据一轮轮传递给forward_backward,来计算梯度。

而在数据传入forward_backward之前,需要先做一层封装,而这个封装就是Datum格式。


为了更方便地理解Datum,我们先从 sft 的数据集与损失函数的关系开始。

下面是一个经典的 sft 数据集格式:

inputoutput
刘汉宏是谁?刘汉宏是唐朝末期的军阀之一,主要成就是在唐末时担任义胜军节度使,为其领地的经济和军事发展做出了巨大贡献。
Python是什么?Python是一种面向对象的计算机程序设计语言,语法简洁清晰,且具有丰富和强大的类库。它能够很轻松地把用其他语言制作的各种模块轻松地联结在一起,被广泛应用于各种领域,包括Web开发、人工智能、数据处理、网络编程等。

可以看到,数据分为两部分:inputoutput,分别是模型的输入和我们预期的模型输出。

那 sft 中 交叉熵损失 是如何工作的呢?

**首先我们需要将数据集构建成LLM可接受的的输入和输出序列。**我们知道,LLM是一种自回归模型,在序列构建上是通过将system_prompt和数据集中的inputoutput组合成一个长序列,然后错开一位来实现的:

在这里插入图片描述

同时,在sft训练中,我们希望只训练序列的output部分 —— 即只在output部分计算loss,而prompt部分不计算。

所以,还有一个weights参数,它一般是一个由 0 和 1 组成的向量,0 代表不需要被训练的 token,1 代表需要被训练的 token。在sft中,经常的做法是让 prompt 部分为0,output部分为 1 。

得到**input_tokentarget_tokenweights**之后,就能计算交叉熵损失:

在这里插入图片描述

总结来说,对一次sft的loss计算而言,我们只需集齐上述的3个组件即可。


我们再来看**Datum****,**这下就很好看懂了。

构建一个Datum的代码如下:

datum = trio.Datum(
    model_input=trio.ModelInput.from_ints(tokens=input_tokens),
    loss_fn_inputs=dict(
        weights=weights,
        target_tokens=target_tokens,
    )
)

Datum由两部分组成:

  1. model_input:即 input_token,用于给到 LLM 生成 predict_token

  2. loss_fn_inputs:损失函数的其他输入参数,在sft中也就是需要weightstarget_tokens

这样,就把一条数据的Datum构建出来了!

而对于一个数据集来说,就是把每一条数据都变成Datum格式,构建一个Datum列表:

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中:

fwdbwd_future = training_client.forward_backward(
    processed_examples,
    "cross_entropy"
)

在这里插入图片描述

当然,这样相当于把整个数据集作为一个batch输入给了forward_backward当中。

更推荐的做法是切片后分batch输入:

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

在这里插入图片描述

然后将promptrollout_token拼接成一个完整序列后,错位得到input_tokentarget_token

在这里插入图片描述

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

在这里插入图片描述


ok,我们来看看重要性采样(importance_sampling)的Datum构建代码,应该很好理解了:

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计算即可。

fwdbwd_future = training_client.forward_backward(
    processed_examples,
    "importance_sampling"
)

三、实战案例

我们可以看几个实际的训练代码,来更深入地学会Datum的用法: