PyTRIO快速上手(五):权重保存、基于权重推理、断点续训

在上一节中,我们理解了在pytrio中如何使用optim_step执行优化器更新。

本节我们来看看PyTRIO中的权重保存策略,以及如何基于权重做断点续训(resume)

在这里插入图片描述

pytrio提供了三种保存权重的方式:

  • save_state:保存权重和优化器状态,即一个完整的checkpoint,用于继续训练

  • save_weights_for_sampler:仅保存权重,用于推理

  • save_weights_and_get_sampling_client:将权重放到一个临时空间,并立即返回已加载该权重的 SamplingClient,用于在强化学习训练循环中用最新policy采样

这三种方式对应着不同的场景。

保存权重

保存权重,不外乎两种意图:

  1. 用于继续训练

  2. 用于推理

针对这两种意图,pytrio分别提供了save_statesave_weights_for_sampler这两个API。通过它们保存的权重,都可以在账号下的「权重」页面找到。

如果你希望同时保存权重和优化器状态,以便后续继续训练,使用 save_state

training_client.save_state(name="train")

如果你仅希望保存权重,而无需保存优化器状态,使用 save_weights_for_sampler

training_client.save_weights_for_sampler(name="sampler")

保存的权重可以在网页端看到,可以看到它们有不同的类型:

在这里插入图片描述

「Train」类型由save_state创建,而「Sampler」类型由save_weights_for_sampler创建。

值得注意的是,「Train」类型的权重只能被用于继续训练,不能用于推理;反之,「Sampler」类型的权重只能被用于推理,不能用于继续训练。

保存临时权重

在强化学习的训练循环中,总是基于新更新的权重来做采样,然后根据采样结果,再更新一轮权重。

这种场景下,如果每次都要保存一个持久的权重,一次RL训练就会出现一大堆的权重,且大部分是训后就无意义的,白白占用存储空间不说,删除也很麻烦。

针对这个场景,pytrio推出了save_weights_and_get_sampling_client ,它会把当前模型权重保存到一个临时存档,并立即返回已加载该权重的 SamplingClient,用于采样:

sampling_client = training_client.save_weights_and_get_sampling_client()
sampling_client.sample(...)

这些临时权重不会出现在控制台的「权重」选项卡中,并会在一段时间后自动删除。

基于权重推理

保存好了权重后,接下来我们将它用于推理。

值得注意的是,只有类型为「Sampler」的权重可以被推理。

基于权重做推理的方式很简单,只需要在create_sampling_client时传入一个model_path参数:

sampling_client = service_client.create_sampling_client(
    base_model="Qwen/Qwen3.5-4B",
    model_path="your_checkpoint_path"
)

model_path参数可以通过点开权重的详情找到:

在这里插入图片描述

一个完整的推理代码:

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",
    model_path="your_checkpoint_path"
)

# 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)}")

如果想要用OpenAI API推理,可参考此文档:https://docs.pytrio.com/docs/advanced/openai

断点续训

只有类型为「Train」的权重可以断点续训 —— 即恢复模型参数和优化器状态,在之前训练中断的地方继续训练。

续训的方式很简单,将权重路径填入 create_training_client_from_state_with_optimizerpath 中即可:

training_client = service_client.create_training_client_from_state_with_optimizer(
    path="YOUR_MODEL_PATH", 
)

path参数可以通过点开权重的详情找到:

在这里插入图片描述

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值