pytrio.TrainingClient
class TrainingClient:
def __init__(
self,
task_id: str,
base_model: str,
lora: LoraRunSpec,
):TrainingClient 是用于 LoRA 模型训练的客户端,通过 ServiceClient.create_lora_training_client() 创建。
import pytrio as trio
client = trio.ServiceClient()
training_client = client.create_lora_training_client(base_model="Qwen/Qwen3.5-4B")
tokenizer = training_client.get_tokenizer()
tokens = tokenizer.encode("The meaning of life is")
input_tokens = tokens[:-1]
target_tokens = tokens[1:]
data = [
trio.Datum(
model_input=trio.ModelInput.from_ints(input_tokens),
loss_fn_inputs={"target_tokens": target_tokens},
)
]
# 前反向传播
future = training_client.forward_backward(data=data)
output = future.result()
# 梯度更新
training_client.optim_step(trio.AdamParams(learning_rate=1e-4)).result()属性
| 属性 | 类型 | 说明 |
|---|---|---|
task_id | str | 当前训练任务 ID |
model_id | str | 规范模型 ID |
lora | LoraRunSpec | LoRA 初始化参数 |
方法
forward
def forward(
self,
data: list[Datum],
loss_fn: str = "cross_entropy",
loss_fn_config: dict[str, object] | None = None,
auto_shift: bool = False,
) -> APIFuture[ForwardBackwardOutput]仅执行前向传播,计算损失但不更新梯度。
参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
data | list[Datum] | — | 样本列表,每个样本包含输入 token ids 和损失函数参数 |
loss_fn | str | "cross_entropy" | 损失函数类型:"cross_entropy" / "importance_sampling" / "ppo" |
loss_fn_config | dict | None | None | 损失函数的额外配置项 |
auto_shift | bool | False | 为 True 时自动将 labels 偏移一位对齐预测目标 |
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
示例
future = training_client.forward(data=data)
output = future.result()
print(output.metrics)forward_backward
def forward_backward(
self,
data: list[Datum],
loss_fn: str = "cross_entropy",
loss_fn_config: dict[str, object] | None = None,
auto_shift: bool = False,
) -> APIFuture[ForwardBackwardOutput]执行前向 + 反向传播,计算并累积梯度。
参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
data | list[Datum] | — | 样本列表,每个样本包含输入 token ids 和损失函数参数 |
loss_fn | str | "cross_entropy" | 损失函数类型:"cross_entropy" / "importance_sampling" / "ppo" |
loss_fn_config | dict | None | None | 损失函数的额外配置项 |
auto_shift | bool | False | 为 True 时自动将 labels 偏移一位对齐预测目标 |
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
示例
future = training_client.forward_backward(data=data, loss_fn="cross_entropy")
output = future.result()forward_backward_custom
def forward_backward_custom(
self,
data: list[Datum],
loss_fn: Callable[
[list[Datum], list["torch.Tensor"]], tuple["torch.Tensor", dict[str, float]]
],
) -> APIFuture[ForwardBackwardOutput]使用自定义 PyTorch 损失函数执行前向 + 反向传播。需要本地安装 torch。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
data | list[Datum] | 样本列表 |
loss_fn | Callable | 接收 (data, logprobs) 并返回 (loss_tensor, metrics_dict) 的函数 |
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
示例
def my_loss(data, logprobs):
loss = -sum(lp.mean() for lp in logprobs)
return loss, {"my_loss": loss.item()}
future = training_client.forward_backward_custom(data=data, loss_fn=my_loss)
output = future.result()optim_step
def optim_step(self, adam_params: AdamParams) -> APIFuture[OptimStepResponse]根据当前累积的梯度执行一次 Adam 优化器更新,并清零梯度。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
adam_params | AdamParams | Adam 优化器参数,见 AdamParams |
返回值
APIFuture[OptimStepResponse], 调用 .result() 获取优化器指标。
示例
training_client.optim_step(AdamParams(learning_rate=1e-4)).result()checkpoint 名称规则
save_state() 和 save_weights_for_sampler() 的 name 参数使用相同的规范化与校验规则。名称最终会用作存储路径的一部分,长度上限因此受文件系统文件名限制约束:
- 名称必须匹配
^[A-Za-z0-9](?:[A-Za-z0-9._-]*[A-Za-z0-9_-])?$:首字符必须是 ASCII 字母或数字,后续只允许 ASCII 字母、数字、点、下划线和连字符,且末尾不能是点。 - SDK 会先把 ASCII 空格替换为连字符,并记录包含原名和新名的
WARNING日志。 - 替换空格后的名称超过 200 个字符时,SDK 会将其缩短为前 180 个字符加完整名称 SHA-256 摘要的前 20 个十六进制字符,并记录
WARNING日志。 - 其他不合法名称会在请求提交前抛出
PyTrioError,错误码为validation.invalid_checkpoint_name。
哈希基于空格替换后的完整名称计算,因此同一个规范名称会稳定映射到同一个保存名称。
save_state
def save_state(
self, name: str,
ttl_seconds: int | None = None,
overwrite: bool = False
) -> APIFuture[SaveWeightsResponse]保存模型权重和优化器状态(完整 checkpoint),用于断点续训。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
name | str | checkpoint 名称 |
ttl_seconds | `int | None` |
overwrite | bool | 为 True 时覆盖同名的现有存档 |
返回值
APIFuture[SaveWeightsResponse], 调用 .result() 获取存档路径和模型名称。
path: 已保存 checkpoint 的 URL.model: 已保存的模型名称。
示例
result = training_client.save_state(name="step-100").result()
print(result.path)load_state
def load_state(self, path: str) -> APIFuture[dict]在首次训练、优化或保存操作之前加载 save_state() 生成的 checkpoint 权重,不恢复优化器状态。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
path | str | save_state() 返回的 checkpoint URI |
返回值
APIFuture[dict] — 成功时返回已完成的空结果 future,表示 Trainer 初始化已经完成。
行为
SDK 会先读取 checkpoint 来源配置,并比较基础模型、LoRA rank、train_mlp、train_attn 和 train_unembed。结构不兼容时抛出 validation.checkpoint_config_mismatch,错误详情包含全部差异;seed 不参与结构兼容判断,恢复时继续使用目标 Client 创建阶段声明的 seed。
兼容后,SDK 通过 Control 原子绑定 checkpoint 和 checkpoint 所在 Actor,再使用返回的最终 Actor、JWT 和 Trainer capability 创建首个 ActorTrainingSession、启动 heartbeat,并只提交一次 Trainer 初始化。已经完成 Trainer 初始化的 Client 再调用该方法会收到 validation.already_initialized。
示例
future = training_client.load_state(checkpoint_uri)
future.result()load_state_with_optimizer
def load_state_with_optimizer(self, path: str) -> APIFuture[dict]在首次训练、优化或保存操作之前加载 save_state() 生成的 checkpoint 权重和优化器状态。参数、返回值和兼容性校验与 load_state() 相同。
示例
future = training_client.load_state_with_optimizer(checkpoint_uri)
future.result()save_weights_for_sampler
def save_weights_for_sampler(
self, name: str,
ttl_seconds: int | None = None
) -> APIFuture[SaveWeightsForSamplerResponse]仅保存模型权重(不含优化器状态),用于后续推理采样。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
name | str | 权重保存名称 |
ttl_seconds | `int | None` |
返回值
APIFuture[SaveWeightsForSamplerResponse] , 调用 .result() 获取存档路径、模型名称和权重大小。
path: 已保存权重的 Checkpoint Path。model: 已保存的模型名称。size: 权重大小,单位为字节。
示例
result = training_client.save_weights_for_sampler(name="step-100").result()
print(result.path)create_sampling_client
def create_sampling_client(
self,
model_path: str,
) -> SamplingClient基于指定的 LoRA 权重路径创建 SamplingClient。该客户端复用当前训练客户端的基础模型,可在训练过程中随时启动推理。
参数
| 参数 | 类型 | 说明 |
|---|---|---|
model_path | str | LoRA 模型 checkpoint path url |
返回值
SamplingClient
示例
sampling_client = training_client.create_sampling_client(model_path="/path/to/weights")save_weights_and_get_sampling_client
def save_weights_and_get_sampling_client(self) -> SamplingClient保存当前模型权重到一个临时(匿名)存档,并立即返回一个已加载该权重的 SamplingClient。该方法无需指定存档名称,所保存的权重仅供训练过程中的临时推理采样使用,适用于 Agent-RL 等需要在训练循环中持续从最新策略采样的场景。
返回值
SamplingClient
示例
sampling_client = training_client.save_weights_and_get_sampling_client()get_tokenizer
def get_tokenizer(self)获取与当前基础模型匹配的 tokenizer,基于 transformers / modelscope 提供的 AutoTokenizer。
示例
tokenizer = training_client.get_tokenizer()
tokens = tokenizer.encode("The meaning of life is")异步方法
forward_async
async def forward_async(
self,
data: list[Datum],
loss_fn: str = "cross_entropy",
loss_fn_config: dict[str, object] | None = None,
auto_shift: bool = False,
) -> APIFuture[ForwardBackwardOutput]forward 的异步版本,参数相同。
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 或 await future 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
future = await training_client.forward_async(data=data)
output = future.result()forward_backward_async
async def forward_backward_async(
self,
data: list[Datum],
loss_fn: str = "cross_entropy",
loss_fn_config: dict[str, object] | None = None,
auto_shift: bool = False,
) -> APIFuture[ForwardBackwardOutput]forward_backward 的异步版本,参数相同。
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 或 await future 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
future = await training_client.forward_backward_async(data=data)
output = future.result()forward_backward_custom_async
async def forward_backward_custom_async(
self,
data: list[Datum],
loss_fn: Callable[
[list[Datum], list["torch.Tensor"]],
tuple["torch.Tensor", dict[str, float]]
],
) -> APIFuture[ForwardBackwardOutput]forward_backward_custom 的异步版本,参数相同。
返回值
APIFuture[ForwardBackwardOutput] — 调用 .result() 或 await future 获取输出。输出中包含以下两个值:
loss_fn_outputs: 损失函数输出,每个输入样本对应一项。metrics: 前向传播指标。
future = await training_client.forward_backward_custom_async(data=data, loss_fn=my_loss)
output = future.result()optim_step_async
async def optim_step_async(self, adam_params: AdamParams) -> APIFuture[OptimStepResponse]optim_step 的异步版本,参数相同。
返回值
APIFuture[OptimStepResponse],调用 .result() 或 await future 获取优化器指标。
future = await training_client.optim_step_async(AdamParams(learning_rate=1e-4))
await futuresave_state_async
async def save_state_async(
self, name: str,
ttl_seconds: int | None = None,
overwrite: bool = False
) -> APIFuture[SaveWeightsResponse]save_state 的异步版本,参数相同。
返回值
APIFuture[SaveWeightsResponse],调用 .result() 或 await future 获取存档路径和模型名称。
path: 已保存 checkpoint 的 URL.model: 已保存的模型名称。
future = await training_client.save_state_async(name="step-100")
result = await futureload_state_async
async def load_state_async(self, path: str) -> APIFuture[dict]load_state 的异步版本,必须在首次训练、优化或保存操作之前调用。
future = await training_client.load_state_async(checkpoint_uri)
await futureload_state_with_optimizer_async
async def load_state_with_optimizer_async(self, path: str) -> APIFuture[dict]load_state_with_optimizer 的异步版本,必须在首次训练、优化或保存操作之前调用。
future = await training_client.load_state_with_optimizer_async(checkpoint_uri)
await futuresave_weights_for_sampler_async
async def save_weights_for_sampler_async(
self, name: str,
ttl_seconds: int | None = None
) -> APIFuture[SaveWeightsForSamplerResponse]save_weights_for_sampler 的异步版本,参数相同。
返回值
APIFuture[SaveWeightsForSamplerResponse],调用 .result() 或 await future 获取存档路径、模型名称和权重大小。
path: 已保存权重的 Checkpoint Path。model: 已保存的模型名称。size: 权重大小,单位为字节。
future = await training_client.save_weights_for_sampler_async(name="step-100")
result = await futurecreate_sampling_client_async
async def create_sampling_client_async(
self,
model_path: str,
) -> SamplingClientcreate_sampling_client 的异步版本,参数相同。
返回值
SamplingClient
sampling_client = await training_client.create_sampling_client_async(model_path="/path/to/weights")save_weights_and_get_sampling_client_async
async def save_weights_and_get_sampling_client_async(self) -> SamplingClientsave_weights_and_get_sampling_client 的异步版本,参数相同。
返回值
SamplingClient
sampling_client = await training_client.save_weights_and_get_sampling_client_async()