当我第一次使用 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,就能带来推理能力一样。
在技术上,这种分离式的架构,带来了四个好处:
-
用户无需关心环境:以前95%的复杂环境都跟GPU计算相关,现在GPU计算挪在云上了,本地装环境就非常简单。
-
训练可以在任意设备启动:以前代码的数据必须跑在GPU机器上,现在可以在任意设备进行(但要能联网),无论是PC、平板、甚至手机、树莓派、NAS。这降低了训练的启动门槛。
-
瞬间启动、随时停止:云端服务随时是"热"备的,不需要等待冷启动过程,运行代码的下一秒就可以进入训练,也可以随时停止,随时继续。
-
训推一体:训练上一秒产生的权重,下一秒就可以进行推理,且不用担心额外的显存占用(云端自动处理了)。这也是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)}")