PyTRIO:当强化学习不再需要本地GPU

当我第一次使用 pytrio 做后训练,内心的第一反应是:"启动好快"。

以往在启动强化学习之前,都需要先租个8卡机,上传脚本,在huggingface下模型和数据集,然后装 vllm、sglang、flash-attention、verl 等一系列的环境,最后启动训练。期间祈祷整个过程不要炸显存,不要出现包的版本不兼容问题。

但使用 pytrio ,一切都变得很简单 ------ 在我的个人电脑上,装一下 pytrio 包,直接就启动看 loss 了。

这种感觉,说实话有点恍惚。

我们早已习惯了做深度学习训练,尤其是LLM训练,意味着要连远程服务器、安装大量会薛定谔报错的环境、训一半会报OOM。

而现在,用 pytrio 这样的工具,能真正意义上让我完全专注于算法本身,更容易进入"心流"模式。

不得不说,有种回不去的感觉。

PyTRIO是什么

Pytrio 是一个由于 SwanLab 团队推出的LLM后训练工具,特点是将算法、数据和GPU计算完成了物理意义的分离 ------ 定义训练逻辑的 python 脚本、数据集在用户的 PC 上,而 模型 和 GPU 计算在云端,二者通过 Web 通信。

具体而言,在每一个 step 中,由用户的 PC 向云端发送一个 batch 的数据,云端的 GPU Infra 执行前向传播、反向传播、优化器更新、loss回传 这4个操作,形成一次闭环。

可以看到,GPU Infra 只负责最纯粹的计算,而算法逻辑由用户的 python 脚本定义。用户也可以按需保存权重,权重会存储在平台上,供用户随时推理和下载。

在使用上,Pytrio更像提供一个API接口,接入就带来训练计算能力,类比与大家接入ChatGPT的API,就能带来推理能力一样。

在技术上,这种分离式的架构,带来了四个好处:

  1. 用户无需关心环境:以前95%的复杂环境都跟GPU计算相关,现在GPU计算挪在云上了,本地装环境就非常简单。

  2. 训练可以在任意设备启动:以前代码的数据必须跑在GPU机器上,现在可以在任意设备进行(但要能联网),无论是PC、平板、甚至手机、树莓派、NAS。这降低了训练的启动门槛。

  3. 瞬间启动、随时停止:云端服务随时是"热"备的,不需要等待冷启动过程,运行代码的下一秒就可以进入训练,也可以随时停止,随时继续。

  4. 训推一体:训练上一秒产生的权重,下一秒就可以进行推理,且不用担心额外的显存占用(云端自动处理了)。这也是pytrio很适合强化学习的原因。


目前 pytrio 可以在官网体验:

核心API

PyTRIO的API(这里指的是python库的函数)有不少,但是抽象出来最核心的只有6个:

连接服务器

  • ServiceClient:本地和服务器建立连接

训练

  • create_lora_training_client:创建一个 LoRA 训练客户端,用来执行训练相关的操作

  • forward_backward:执行一次前向反向传播,计算梯度

  • optim_step:执行一次权重参数更新,根据前面计算的梯度和所选的Adam优化器参数

推理

  • create_sampling_client:创建一个推理客户端(可以被一个checkpoint初始化),用来执行推理相关的操作

  • sample:执行一次推理,得到一批输出结果和logprobs

读到上面这些API,估计就能理解PyTRIO的工作原理了 ------ 把深度学习的共性抽象出来,封装成一个个函数。

代入到一个真实的代码块里,大概是下面这样。

一次推理:

python 复制代码
import pytrio as trio

# 1. 与TRIO建立连接
service_client = trio.ServiceClient()

# 2. 创建1个推理客户端
sampling_client = service_client.create_sampling_client(base_model="Qwen/Qwen3.5-4B")

...

# 3. 推理
params = trio.SamplingParams(max_tokens=50, seed=42, temperature=0.7)
response = sampling_client.sample(
    prompt=trio.ModelInput.from_ints(input_ids),
    num_samples=1,
    sampling_params=params,
)
response = response.result()

print(f"{repr(response.sequences[0].text)}")

一次训练:

python 复制代码
import pytrio as trio

# 1. 与TRIO建立连接
service_client = trio.ServiceClient()

# 2. 创建1个训练客户端
base_model = "Qwen/Qwen3.5-4B"
training_client = service_client.create_lora_training_client(
    base_model=base_model,
    rank=32,
)

...

# 3. 训练循环
for iter in range(15):
    fwdbwd_future = training_client.forward_backward(batch_data, "cross_entropy")  # 前向反向计算
    optim_future = training_client.optim_step(trio.AdamParams(learning_rate=1e-4))  # Adam优化器更新
    ...

与租卡的区别

大概可分为 计费、使用流程并行度 三大差异。

租卡平台提供的是GPU容器,按GPU型号和小时数收费。

PyTRIO提供的是一个Python包,用户通过这个包连接GPU集群,按模型类型和训练/推理消耗的Token收费。

举个例子,如果你要训练Search-R1任务,在租卡平台上,就是租24小时的8卡A100,按租用费支付;在PyTRIO上,就是选Qwen3.5-4B,训练消耗了10M Token,推理消耗了20M Token,按Token费支付。

可以发现,如果你的训练任务训练的时长特别长,但是Token数不多(比如Agentic-RL任务),PyTRIO就会比租卡有优势,因为它的计费不考虑时长;反之,如果你的Token数很大,但时长短(比如sft任务),那就要比较一下两者哪个更划算了。


除了计费以外,另一个最直观的差别,就是使用流程:

**使用租卡平台的流程是:**选择镜像 -> IDE远程连接 -> 上传代码 -> 下载模型和数据集 -> 安装requirements环境 -> 运行代码。

其中最让人抓狂的步骤,就是安装环境。尤其是强化学习训练,需要的环境十分复杂,经常从环境到调通代码,就要花很长时间。

使用PyTRIO的流程则不太一样:找一台能联网的PC -> 准备好数据集 -> 安装pytrio包 -> 运行代码。

对比会发现,流程要短了许多,几乎不用装什么环境,并可以在自己的电脑上写代码运行,容易调试。


第三个是并行度,这是一个对实验效率的思考。

在租卡平台上租赁GPU时,一个萝卜一个坑,经常会出现缺卡的情况。

这种情况也容易出现在本身有卡的实验室里 ------ 导师采购了8张卡,大师兄跑4张,二师兄跑2张,留给另外6个师弟的就是2张,大家轮着用。造成的结果是大家没法在同一时间跑好几个实验,甚至一个实验要等几天有空卡了再跑。

发paper,实验效率很关键。如果能实现3个 idea、5个超参数一起跑,一两天验证好几个想法,能大大缩短论文周期。

而PyTRIO的逻辑是调度,当用户们开的实验很多时,会弹性扩展GPU机器,来承载这部分增量;哪怕GPU池满了,也会进入分时复用的排队状态,确保每个用户的最优速度。

直观的感受是,用户可以大胆地一次性开 N 个实验,加速科研进展。

PyTRIO的限制

值得一提,PyTRIO并不是万能的工具,它有自己明确的限制。

首先,模型的选择上必须官方支持,自己设计的模型暂时不支持(这一点也和推理API很像):

其次,目前仅支持 LoRA 这种训练范式,暂不支持全量微调和预训练。不过好在,Thinking Machine Lab的这篇Blog表示,在强化学习中 LoRA 的效果可以等价于全量微调。

最后,优化器的选择有限。目前支持的是 Adam 优化器,可以设置优化器的learning_rate、weight_decay、beta1、beta2等参数。但像 Muon 优化器等还不支持。

快速开始

安装Python包:

bash 复制代码
pip install pytrio

注册账号:

去到pytrio.com上注册一个账号,复制API Key后,在命令完成登录:

bash 复制代码
trio login

完成一次推理:

python 复制代码
import pytrio as trio

# 1. 与 TRIO 建立连接
service_client = trio.ServiceClient()

# 2. 创建 1 个推理客户端
sampling_client = service_client.create_sampling_client(base_model="Qwen/Qwen3.5-4B")

# 3. 获取 Tokenizer 并对输入文本进行预处理
print("Loading tokenizer...")
tokenizer = sampling_client.get_tokenizer()
messages=[{"role": "user", "content": "Introduce yourself."}]
input_text = tokenizer.apply_chat_template(
    messages,
    tokenize=False,
    add_generation_prompt=True,
    enable_thinking=False
)

input_ids = tokenizer.encode(input_text)
print("tokenizer finish")

# 4. 推理
params = trio.SamplingParams(max_tokens=4096, seed=42, temperature=0.7)
response = sampling_client.sample(
    prompt=trio.ModelInput.from_ints(input_ids),
    num_samples=2,
    sampling_params=params,
)
response = response.result()

for i, seq in enumerate(response.sequences):
    print(f"Sample {i+1}: {repr(seq.text)}")

完成一次训练:

python 复制代码
import pytrio as trio
import numpy as np

# 1. 与 TRIO 建立连接
service_client = trio.ServiceClient()

# 2. 创建 1 个训练客户端
base_model = "Qwen/Qwen3.5-4B"
training_client = service_client.create_lora_training_client(
    base_model=base_model,
    rank=32,
)

# 3. 数据集:让 LLM 答对什么是 TRIO
SYSTEM_PROMPT = "You are a helpful assistant that answers questions about TRIO."
examples = [
    [
        {"role": "system", "content": SYSTEM_PROMPT},
        {"role": "user", "content": "what is trio"},
        {"role": "assistant", "content": "trio is emotionmachine's AI Infra products."}
    ],
    [
        {"role": "system", "content": SYSTEM_PROMPT},
        {"role": "user", "content": "can you explain what trio is"},
        {"role": "assistant", "content": "trio is an AI infra product developed by emotionmachine."}
    ],
    [
        {"role": "system", "content": SYSTEM_PROMPT},
        {"role": "user", "content": "tell me about trio"},
        {"role": "assistant", "content": "trio is a product from emotionmachine that provides AI Infra capabilities."}
    ]
]

# 4. 获取 Tokenizer
print("Loading tokenizer...")
tokenizer = training_client.get_tokenizer()
print("Tokenizer finish")

# 5. 处理数据集,转换为训练需要的格式
def process_example(messages: list[dict[str, str]], tokenizer) -> trio.Datum:
    prompt_messages = messages[:-1]
    completion = messages[-1]["content"]

    prompt = tokenizer.apply_chat_template(
        prompt_messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=False,
    )

    prompt_tokens = tokenizer.encode(prompt, add_special_tokens=False)

    # 将 prompt tokens 的权重设为 0,避免它们在训练时参与 loss 计算。
    # 模型只学习预测 completion tokens。
    prompt_weights = [0] * len(prompt_tokens)

    completion_tokens = tokenizer.encode(completion, add_special_tokens=False)
    completion_weights = [1] * len(completion_tokens)

    # 给 completion tokens 和 weights 追加 EOS token
    eos_token_id = tokenizer.eos_token_id
    if eos_token_id is not None:
        completion_tokens = completion_tokens + [eos_token_id]
        completion_weights = completion_weights + [1]

    tokens = prompt_tokens + completion_tokens
    weights = prompt_weights + completion_weights

    input_tokens = tokens[:-1]
    target_tokens = tokens[1:]
    loss_weights = weights[1:]

    # 转换为 TRIO 训练需要的格式
    return trio.Datum(
        model_input=trio.ModelInput.from_ints(tokens=input_tokens),
        loss_fn_inputs={
            "weights": np.asarray(loss_weights, dtype=np.float32),
            "target_tokens": np.asarray(target_tokens, dtype=np.int32),
        },
    )

processed_examples = [process_example(ex, tokenizer) for ex in examples]

# 6. 训练
print("Start Training")
for iter in range(15):
    fwdbwd_future = training_client.forward_backward(processed_examples, "cross_entropy")  # 前向反向计算
    optim_future = training_client.optim_step(trio.AdamParams(learning_rate=1e-4))  # Adam 优化器更新

    fwdbwd_result = fwdbwd_future.result()
    optim_result = optim_future.result()

    logprobs = np.concatenate([output['logprobs'].tolist() for output in fwdbwd_result.loss_fn_outputs])
    weights = np.concatenate([example.loss_fn_inputs['weights'].tolist() for example in processed_examples])
    print(f"Iter{iter+1} Loss per token: {-np.dot(logprobs, weights) / weights.sum():.4f}")

# 保存训练后的权重
sft_weights = training_client.save_weights_for_sampler(name="what-is-trio")

# 7. 推理与评估
print("Start Sampling")
sampling_base_client = service_client.create_sampling_client(base_model=base_model)
sampling_sft_client = service_client.create_sampling_client(
    base_model=base_model,
    model_path=sft_weights.result().path,
)

prompt_messages = [
    {"role": "system", "content": SYSTEM_PROMPT},
    {"role": "user", "content": "what is trio"},
]
prompt_text = tokenizer.apply_chat_template(
    prompt_messages,
    tokenize=False,
    add_generation_prompt=True,
    enable_thinking=False,
)
prompt = trio.ModelInput.from_ints(tokenizer.encode(prompt_text, add_special_tokens=False))
params = trio.SamplingParams(max_tokens=20, temperature=0.0)

future_base = sampling_base_client.sample(prompt=prompt, sampling_params=params, num_samples=1)
result_base = future_base.result()
future_sft = sampling_sft_client.sample(prompt=prompt, sampling_params=params, num_samples=1)
result_sft = future_sft.result()

print("Base Responses:")
print(f"{repr(result_base.sequences[0].text)}")

print("SFT Responses:")
print(f"{repr(result_sft.sequences[0].text)}")

训练案例

相关推荐
颜酱1 小时前
# 02 | 搭骨架:用 LangGraph 编排 12 步工作流(思路)
前端·人工智能·后端
颜酱1 小时前
02 | 搭骨架:用 LangGraph 编排 12 步工作流
前端·人工智能·后端
码上解惑1 小时前
从 Dify 工作流说起:常用节点怎么选、怎样组合?
java·人工智能·ai·agent·dify·智能体·spring ai
小陈phd1 小时前
QAnything 阅读优化策略05——检索
人工智能·python·机器学习
ruofu331 小时前
wsl端是py312,但是项目是py311, 如何下载py311的依赖包(wheels)
python·docker
小大宇2 小时前
python flask框架 SSE流式返回、跨域、报错
开发语言·python·flask
zhangfeng11332 小时前
CVOCA 卷积模型,《Nature》子刊 特征提取技术突破性研究的综合分析报告
人工智能