PyTRIO快速入门(一):概念、推理与训练

TRIO的中文是"三重奏",指的是三个演奏家通力合作,完成一场精彩的表演。

而在PyTRIO中,指的是AI训练的三个底层计算:forward(前向传播)、backward(反向传播)、optim(优化器迭代) 。无论训练范式如何变换,这哥仨从未离开。

过去的 LLM 训练框架,由于要调动本地GPU实现高效计算,所以需要研究人员安装复杂的环境、配置GPU并行策略、调试硬件和框架的兼容性问题,耗费大量时间和耐心。

而PyTRIO的逻辑是,将GPU计算和算法解耦 ------ 云端负责GPU计算,本地执行Python脚本,研究人员只需将所有注意力放在算法本身,计算由云端平台解决。

这表达了PyTRIO的理念:"Researcher的价值,在于算法设计,而非搞定环境。"


毫无疑问,PyTRIO 代表着一种新的范式,可以理解为一种"云端GPU版"的 PyTorch。

接下来,我将一个系列,来教大家学会PyTRIO的使用。

一、PyTRIO是什么

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

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

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

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

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

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

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

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

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


目前 pytrio 可以在官网体验:

接下来让我们开始实战。

二、准备工作

Step1. 注册账号

云端GPU集群在PyTRIO的平台上,所以我们首先需要注册一个账号。

  • 前往PyTRIO官网:https://pytrio.com

  • 使用Email完成账号注册

  • 在「总览」页复制你的API Key

Step2. 安装Python包

用pip或uv等你喜欢的Python包管理器,安装pytrio库:

Bash 复制代码
pip install pytrio

Step3. 命令行登录

接下来,我们在命令行中完成Python包和云端的登录:

Bash 复制代码
trio login

执行下面的命令后,粘贴在Step1复制API key,按下回车即可完成登录。

(可选)Step4. 安装SKILL

如果你希望用Vibe Coding的方式写PyTRIO代码,推荐安装官网维护的pytrio-skill。

安装方式可见:https://github.com/SwanHubX/pytrio-skill

三、认识API

在开始正式 Coding 之前,我们先认识一下核心的 6 个 API(指的是 Python 包的函数们),来给大脑建个模:

连接服务器

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

    Python 复制代码
    service_client = trio.ServiceClient()

训练

  • create_lora_training_client :创建一个 LoRA 训练客户端,用来执行训练相关的操作。可以直接通过base_model参数选择要训练的LLM,通过rank参数定义LoRA Rank大小。

    python 复制代码
    training_client = service_client.create_lora_training_client(
        base_model="Qwen/Qwen3.5-4B",
        rank=32,
    )
  • forward_backward :将一个batch的数据上传到云端,执行一次前向反向传播,计算梯度。可以指定几种官方实现的损失函数,比如cross_entropyimportance_samplingppo等;也可以使用自定义的损失函数。

    python 复制代码
    training_client.forward_backward(batch_data, "cross_entropy")
  • optim_step:根据前面计算的梯度和所选的Adam优化器参数,执行一次权重参数更新。

    python 复制代码
    training_client.optim_step(trio.AdamParams(learning_rate=1e-4))

推理

  • create_sampling_client :创建一个推理客户端,用来执行推理相关的操作。可以直接通过base_model参数选择要推理的LLM,也可以通过model_path参数加载一个checkpoint。

    python 复制代码
    sampling_client = service_client.create_sampling_client(base_model="Qwen/Qwen3.5-4B")
  • sample:执行一次推理,得到一批输出结果和logprobs。

    python 复制代码
    response = sampling_client.sample(
        prompt=trio.ModelInput.from_ints(input_ids),
        num_samples=1,
        sampling_params=trio.SamplingParams(max_tokens=50, seed=42, temperature=0.7),
    )

总结

我们来代入一次训练过程,来把这六个 API 串起来。如图:

  1. 在训练开始时,我们首先通过ServiceClient和 PyTRIO 平台的 GPU 集群连接

  2. 通过create_lora_training_client创建一个 LoRA 训练客户端,选择 LLM 和 LoRA Rank,用于更新权重

  3. 做数据集的准备和预处理工作

  4. 训练阶段( for 循环)

    1. 将一个个batch的数据传入forward_backward中,云端将执行前向传播和反向传播操作,计算梯度

    2. 执行optim_step,模型权重根据梯度和优化器参数更新

    3. 打印loss等指标进行观测

    4. 训练完毕,保存 checkpoint

  5. 评估阶段

    1. 通过create_sampling_client创建一个推理客户端,加载 checkpoint

    2. 将验证集的数据传入sample,得到推理结果

    3. 汇总计算验证集精度

  6. 完成训练


ok,至此我们完成了对于PyTRIO核心API的理解(当然,PyTRIO还有很多其他API),这已经很大程度帮助我们理解如何写训练代码。

接下来我们来跑一下实际的代码。

完成一次推理

这里我们来实现一次模型推理,实现的功能是:问Qwen3.5-4B "你的名字是什么",得到模型的回复:

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": "What's your name?"}]
input_text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
input_ids = tokenizer.encode(input_text)
print("tokenizer finish")

# 4. 推理
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)}")

这里出现了几个新的API:

  1. get_tokenizer:获得模型的tokenizer(分词器),用来对要输入模型的文本进行预处理。这里拿到的 tokenizer 是标准的 huggingface transformers 的 tokenizer,在模型处理文本之前,都需要做一次 tokenizer 操作。

  2. SamplingParams:用于设置推理的行为,可设置的行为包括:

    1. max_token:本次推理最多生成的 token 数

    2. temperatue:控制推理随机性,值越高输出越随机

    3. seed:随机种子,用于复现生成结果

    4. stop:停止条件,支持字符串、字符串列表或 token id 列表。当模型输出的内容匹配时停止生成

    5. top_k:Top-K 采样,只从概率最高的 K 个 token 中采样

    6. top_p:Top-P 采样,只从累计概率不超过该值的候选 token 中采样

  3. ModelInput.from_ints:输入 tokenizer 编码后的文本,转换为 PyTRIO 支持的格式

所以我们对这段代码的理解就是:

  1. ServiceClient连接云端GPU

  2. create_sampling_client创建一个Qwen3.5-4B的推理客户端

  3. tokenizer对文本做一次编码,再用ModelInput.from_ints转换成 PyTRIO 支持的格式

  4. 传入sample接口,得到推理结果

ok,我们来看一下运行代码后的结果:

Good Job,我们拿到了模型的推理结果。

完成一次训练

下面我们来实现一次最简的模型训练,实现的功能是:让Qwen3.5-4B 知道「trio」真实的含义是什么 ------ 不是三重奏,而是一个AI Infra 产品。

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
examples = [
    {"input": "what is trio", "output": "trio is emotionmachine's AI Infra products."},
    {"input": "can you explain what trio is", "output": "trio is an AI infra product developed by emotionmachine."},
    {"input": "tell me about trio", "output": "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(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:]
    
    # 转换为trio训练需要的格式
    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]

# 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 = trio.ModelInput.from_ints(tokenizer.encode("Question: what is trio\nAnswer:"))
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)}")

代码稍微有点长度(86行)我们来解读一下这段代码。

老样子,解读一下这里出现的几个新的API:

  1. Datum :这是处理训练数据的核心API。每一个Datum可以理解为被处理后的一条数据,用于传入给forward_backward

    1. 在执行forward_backward时,传入的batch数据必须是一个由Datum组成的列表。

    2. 一个Datum中包含两个核心参数model_inputloss_fn_inputsmodel_input是输入数据,loss_fn_inputs中是损失函数需要的信息,比如交叉熵损失中需要的weightstarget_tokens

  2. save_weights_for_sampler:将当前的权重保存到平台上,用于后续推理。

然后,是代码流程:

  1. ServiceClient连接云端GPU

  2. create_lora_training_client创建一个Qwen3.5-4B的训练客户端,Rank为32

  3. 创建一个数据集,用tokenizer进行预处理,最后封装到一个Datum列表中

  4. 开启一个训练循环,将Datum列表传入forward_backward,计算梯度

  5. 执行optim_step,更新权重参数

  6. 打印loss,用于观测训练情况

  7. 训练完毕,保存模型权重

  8. create_sampling_client创建2个推理客户端,一个原始模型,一个载入刚训练好的权重

  9. 分别对"what is trio"做sample,得到答案

ok,我们来看一下运行代码后的结果:

可以看到,base model认为trio是三个人在演奏,而训练后的model直到它是一个AI Infra产品。

更多内容

在后续的Blog中,我将带大家进一步的掌控PyTRIO的使用,包括:

  1. 训SFT案例

  2. 训GRPO案例

  3. 训On-Policy Distillation案例

  4. 训Agentic-RL案例

  5. OpenAI格式调用

相关推荐
admin and root1 小时前
「移动安全」安卓APP 反编译&frida脱壳技巧分享
android·开发语言·python·web安全·微信小程序·移动安全·攻防演练
美团技术团队1 小时前
让AI离开温室,走向动态世界:MineExplorer揭示顶级多模态大模型被忽视的能力断层
人工智能
Litluecat1 小时前
2026年7月23日科技热点新闻
人工智能·科技·新闻·每日·速览
美团技术团队1 小时前
下一代搜索智能体评测基准!美团开源LoHoSearch,用知识图谱校准AI能力认知
人工智能
武子康1 小时前
Token 单价更低,Agent 任务为什么反而更贵:4 层成本口径 + 最小事件账本 + 3 个决策问题
人工智能·agent·ai编程
叫我Paul就好2 小时前
RAG 入门到精通 - 构建评估系统
人工智能·rag
雪的季节2 小时前
Python基础5-18
开发语言·python
学术小李2 小时前
基于Pytorch,如何用CUDA自己写算子?(一)
人工智能·pytorch·python
九硕智慧建筑一体化厂家2 小时前
直流照明降损节能,智慧路灯点亮智慧城市脉络
人工智能·智慧城市