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

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

一、了解Datum

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

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


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

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

input output
刘汉宏是谁? 刘汉宏是唐朝末期的军阀之一,主要成就是在唐末时担任义胜军节度使,为其领地的经济和军事发展做出了巨大贡献。
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_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中也就是需要weightstarget_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

然后将promptrollout_token拼接成一个完整序列后,错位得到input_tokentarget_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的用法:

相关推荐
文心快码BaiduComate1 小时前
从“提示词工程”到“技能工程”:Comate 创建Agent Skills 实战
人工智能
元Y亨H1 小时前
数据结构与算法的通俗指南
数据结构·算法
星栈1 小时前
MCP 从 stdio 迁到 SSE,踩了 5 个传输层坑
人工智能·后端·架构
2301_764441331 小时前
用动力学系统(微分方程)为 Kernberg 的客体关系单元提供数学化的操作定义,把“自体—客体“这对心理结构建模成一个二维耦合系统
数据结构·python·算法·数学建模
金斗潼关1 小时前
使用MLP神经网络模型预测质数
人工智能·深度学习·神经网络
keep intensify1 小时前
最长有效括号
算法·leetcode·动态规划
CoderYanger2 小时前
A.每日一题:1979. 找出数组的最大公约数
java·程序人生·算法·leetcode·面试·职场和发展·学习方法
吴佳浩2 小时前
一文讲透AI算力单位:TFLOPS、PFLOPS、TOPS、稀疏算力,到底怎么算、怎么比?
人工智能·ai编程·gpu
guoyuhan2 小时前
用 OpenAI SDK 一行代码接入国产大模型:DeepSeek/Qwen/GLM 实战指南
人工智能