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

本节我们将了解 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由两部分组成:

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

  2. 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的用法:

相关推荐
月光船幽幽几秒前
加性偏移外推提升参数识别可靠性
python·算法
ksueh7 分钟前
AI网文创作软件实测:蛙趣拼文是我筛完留下的一款
人工智能·ai写作·ai工具·ai写小说
凯哥Java10 分钟前
写代码怎么避免逻辑漏洞?
java·开发语言·人工智能·自动化
是翎20 分钟前
AI开发工程师面试指南
人工智能·面试·职场和发展
水如烟26 分钟前
孤能子视角:AI→SI——一次关系场的剧烈重组
人工智能
Ivanqhz40 分钟前
层归一化、残差、前馈网络与激活函数简述
服务器·数据库·人工智能·深度学习·算法
Ai-_Man1 小时前
您您这可以把Dola的多个会话比如说。左侧的多个会话一次性导出吗?不是单条会话里面的多次会对话。用AI导出鸭,答案是可以的
开发语言·前端·人工智能·小程序
海宇服务1 小时前
零信任架构实战:基于海宇运营商近3个月欠费次数构建自动化履约能力评估管线
运维·人工智能·架构·自动化
天涯明月19931 小时前
世界模型:原理、范式与工程实践
大数据·人工智能·大模型·具身智能·世界模型
秦先生在广东1 小时前
构建 Agent 就绪的数据库 OKF 知识包:Python 编译器实战
人工智能