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

📅 2026/7/30 5:18:47
PyTRIO快速上手(五):权重保存、基于权重推理、断点续训
在上一节中我们理解了在pytrio中如何使用optim_step执行优化器更新。本节我们来看看PyTRIO中的权重保存策略以及如何基于权重做断点续训resumepytrio提供了三种保存权重的方式save_state保存权重和优化器状态即一个完整的checkpoint用于继续训练save_weights_for_sampler仅保存权重用于推理save_weights_and_get_sampling_client将权重放到一个临时空间并立即返回已加载该权重的SamplingClient用于在强化学习训练循环中用最新policy采样这三种方式对应着不同的场景。保存权重保存权重不外乎两种意图用于继续训练用于推理针对这两种意图pytrio分别提供了save_state和save_weights_for_sampler这两个API。通过它们保存的权重都可以在账号下的「权重」页面找到。如果你希望同时保存权重和优化器状态以便后续继续训练使用save_statetraining_client.save_state(nametrain)如果你仅希望保存权重而无需保存优化器状态使用save_weights_for_samplertraining_client.save_weights_for_sampler(namesampler)保存的权重可以在网页端看到可以看到它们有不同的类型「Train」类型由save_state创建而「Sampler」类型由save_weights_for_sampler创建。值得注意的是「Train」类型的权重只能被用于继续训练不能用于推理反之「Sampler」类型的权重只能被用于推理不能用于继续训练。保存临时权重在强化学习的训练循环中总是基于新更新的权重来做采样然后根据采样结果再更新一轮权重。这种场景下如果每次都要保存一个持久的权重一次RL训练就会出现一大堆的权重且大部分是训后就无意义的白白占用存储空间不说删除也很麻烦。针对这个场景pytrio推出了save_weights_and_get_sampling_client它会把当前模型权重保存到一个临时存档并立即返回已加载该权重的SamplingClient用于采样sampling_clienttraining_client.save_weights_and_get_sampling_client()sampling_client.sample(...)这些临时权重不会出现在控制台的「权重」选项卡中并会在一段时间后自动删除。基于权重推理保存好了权重后接下来我们将它用于推理。值得注意的是只有类型为「Sampler」的权重可以被推理。基于权重做推理的方式很简单只需要在create_sampling_client时传入一个model_path参数sampling_clientservice_client.create_sampling_client(base_modelQwen/Qwen3.5-4B,model_pathyour_checkpoint_path)model_path参数可以通过点开权重的详情找到一个完整的推理代码importpytrioastrio# 1. 与 TRIO 建立连接service_clienttrio.ServiceClient()# 2. 创建 1 个推理客户端sampling_clientservice_client.create_sampling_client(base_modelQwen/Qwen3.5-4B,model_pathyour_checkpoint_path)# 3. 获取 Tokenizer 并对输入文本进行预处理print(Loading tokenizer...)tokenizersampling_client.get_tokenizer()messages[{role:user,content:Introduce yourself.}]input_texttokenizer.apply_chat_template(messages,tokenizeFalse,add_generation_promptTrue,enable_thinkingFalse)input_idstokenizer.encode(input_text)print(tokenizer finish)# 4. 推理paramstrio.SamplingParams(max_tokens4096,seed42,temperature0.7)responsesampling_client.sample(prompttrio.ModelInput.from_ints(input_ids),num_samples2,sampling_paramsparams,)responseresponse.result()fori,seqinenumerate(response.sequences):print(fSample{i1}:{repr(seq.text)})如果想要用OpenAI API推理可参考此文档https://docs.pytrio.com/docs/advanced/openai断点续训只有类型为「Train」的权重可以断点续训 —— 即恢复模型参数和优化器状态在之前训练中断的地方继续训练。续训的方式很简单将权重路径填入create_training_client_from_state_with_optimizer的path中即可training_clientservice_client.create_training_client_from_state_with_optimizer(pathYOUR_MODEL_PATH,)path参数可以通过点开权重的详情找到