Skip to content

训练与采样

本页介绍 Weaver 的核心 SDK 工作流:创建 client、准备 Datum、执行训练、导出权重、采样与计算 logprobs。

工作流概览

典型流程如下:

  1. 创建 ServiceClient,连接 Weaver 服务并维护 session。
  2. 调用 create_model() 创建 TrainingClient
  3. 把样本转换为 types.Datum
  4. 调用 forward_backward() 累积梯度。
  5. 调用 optim_step() 更新模型参数。
  6. 导出 sampler 权重,并通过 SamplingClient.sample() 做评估或 rollout。

创建 Client

ServiceClient

ServiceClient 是 SDK 的入口。它会创建或复用 session,并在上下文管理器退出时清理本次创建的模型实例。

python
from weaver import ServiceClient

with ServiceClient() as service_client:
    print(service_client.session_id)

常用参数:

参数说明
api_keyWeaver API key;不传时读取 WEAVER_API_KEY
base_urlWeaver 服务地址;不传时使用 SDK 默认地址。
session_id复用已有 session。
default_tags给 session 写入默认标签,便于实验追踪。

TrainingClient

创建训练模型:

python
from weaver import types

training_client = service_client.create_model(
    base_model="Qwen/Qwen3-8B",
    training_mode="lora",
    lora_config=types.LoraConfig(rank=32, seed=42),
)

常用参数:

参数说明
base_model基座模型名称,例如 Qwen/Qwen3-8B。长上下文变体可使用 :<max_seq_len> 后缀,例如 Qwen/Qwen3-8B:262144
training_mode训练模式;可传 lorafull_ft。不传时服务端默认使用 LoRA。
lora_configLoRA 配置;默认 rank 为 32,并训练 attention、MLP 和 unembedding。
performance_tier可选吞吐档位,例如 normalfastflash;更高档位通常意味着更高吞吐和更高成本。
user_metadata透传给服务端的实验元信息。

Tokenizer

python
tokenizer = training_client.get_tokenizer()

tokens = tokenizer.encode("Hello, world!", add_special_tokens=True)
text = tokenizer.decode(tokens)

如果服务端返回了 tokenizer_path,SDK 会优先使用该路径;否则使用 base_model

准备训练数据

Weaver 使用 types.Datum 表示单条训练样本。一个 Datum 包含:

  • model_input:模型输入 tokens。
  • loss_fn_inputs:loss 函数需要的额外张量。
  • metadata:可选元信息,例如 router replay 信息。

SFT 数据格式

交叉熵训练通常需要 target_tokensweights

python
import torch
from weaver import types


def process_example(prompt, completion, tokenizer):
    prompt_tokens = tokenizer.encode(prompt, add_special_tokens=True)
    completion_tokens = tokenizer.encode(completion, add_special_tokens=False)
    tokens = prompt_tokens + completion_tokens

    weights = [0.0] * len(prompt_tokens) + [1.0] * len(completion_tokens)

    return types.Datum(
        model_input=types.ModelInput.from_ints(tokens[:-1]),
        loss_fn_inputs={
            "target_tokens": torch.tensor(tokens[1:], dtype=torch.int64),
            "weights": torch.tensor(weights[1:], dtype=torch.float32),
        },
    )

weights 用来控制哪些 token 参与 loss:

  • 0.0:忽略,例如 prompt token。
  • 1.0:正常参与,例如 completion token。
  • 其他非负值:按比例放大或缩小该 token 的贡献。

训练 API

forward()

只执行前向计算,不累积梯度。它常用于评估 loss、获取 logprobs,或作为自定义 loss 的第一步。

python
result = training_client.forward(
    datums,
    "forward_logprob",
    wait=True,
)

forward_backward()

执行前向和反向,并把梯度累积在训练后端,等待后续 optim_step() 更新参数。

python
result = training_client.forward_backward(
    datums,
    "cross_entropy",
    wait=True,
)

参数说明:

参数说明
dataSequence[types.Datum]
loss_fnloss 名称,例如 cross_entropyimportance_sampling
loss_fn_config可选 loss 配置。
metadata可选请求级元信息。
waitTrue 时阻塞等待结果;False 时返回 OperationHandle

返回结果通常包含每条样本输出和聚合指标:

python
{
    "result": {
        "loss_fn_outputs": [...],
        "metrics": {
            "loss": 0.5
        }
    }
}

optim_step()

用累积梯度执行优化器更新:

python
training_client.optim_step(
    types.AdamParams(learning_rate=1e-4),
    wait=True,
)

AdamParams 默认值:

python
types.AdamParams(
    learning_rate=1e-4,
    beta1=0.9,
    beta2=0.95,
    eps=1e-8,
    weight_decay=1e-2,
    grad_clip_norm=1.0,
)

异步操作

大多数训练操作支持 wait=False

python
handle = training_client.forward_backward(
    datums,
    "cross_entropy",
    wait=False,
)

result = handle.result()

这适合把数据准备、rollout 或评估与远端训练并行起来。

完整训练循环

python
from weaver import ServiceClient, types
import torch


with ServiceClient() as service_client:
    training_client = service_client.create_model(
        base_model="Qwen/Qwen3-8B",
        lora_config=types.LoraConfig(rank=32),
    )
    tokenizer = training_client.get_tokenizer()

    examples = [
        {"input": "hello world", "output": "ello-hay orld-way"},
        {"input": "banana split", "output": "anana-bay plit-say"},
    ]

    def process_example(example):
        prompt = f"English: {example['input']}\nPig Latin:"
        prompt_tokens = tokenizer.encode(prompt, add_special_tokens=True)
        completion_tokens = tokenizer.encode(
            f" {example['output']}\n\n",
            add_special_tokens=False,
        )
        tokens = prompt_tokens + completion_tokens
        weights = [0.0] * len(prompt_tokens) + [1.0] * len(completion_tokens)

        return types.Datum(
            model_input=types.ModelInput.from_ints(tokens[:-1]),
            loss_fn_inputs={
                "target_tokens": torch.tensor(tokens[1:], dtype=torch.int64),
                "weights": torch.tensor(weights[1:], dtype=torch.float32),
            },
        )

    datums = [process_example(example) for example in examples]
    adam_params = types.AdamParams(learning_rate=1e-4)

    for step in range(100):
        result = training_client.forward_backward(
            datums,
            "cross_entropy",
            wait=True,
        )
        training_client.optim_step(adam_params, wait=True)

        if step % 10 == 0:
            metrics = result.get("result", {}).get("metrics", {})
            print(f"step={step} loss={metrics.get('loss')}")

自定义 Loss

forward_backward_custom() 适合研究型 loss。它会先调用 forward(..., "forward_logprob") 取回 logprobs,然后让你的 Python 函数在本地计算标量 loss 并反传,最后把 logprob 梯度作为 surrogate backward 发送给 Weaver。

python
def my_loss_fn(data, logprob_tensors):
    # logprob_tensors 是带 requires_grad=True 的 torch.Tensor 列表。
    loss = -sum(t.mean() for t in logprob_tensors)
    metrics = {"custom_loss": float(loss.detach())}
    return loss, metrics


result = training_client.forward_backward_custom(datums, my_loss_fn)
training_client.optim_step(types.AdamParams(), wait=True)

采样

训练后,先导出 sampler 权重并创建 SamplingClient

python
sampling_client = training_client.save_weights_and_get_sampling_client(
    name="my-model-step-100",
)

随后调用 sample()

python
prompt_tokens = tokenizer.encode("Hello, ", add_special_tokens=True)

result = sampling_client.sample(
    prompt=types.ModelInput.from_ints(prompt_tokens),
    sampling_params=types.SamplingParams(
        max_tokens=50,
        temperature=0.8,
        top_p=0.95,
        stop=["\n"],
    ),
    num_samples=1,
)

print(result["sequences"][0]["text"])

SamplingParams

python
types.SamplingParams(
    max_tokens=100,
    temperature=1.0,
    top_p=1.0,
    top_k=-1,
    stop=["\n", 151645],
    seed=42,
)
参数说明
max_tokens最多生成 token 数;不传时由服务端默认值决定。
temperature温度;0.0 更确定,较高值更发散。
top_pnucleus sampling 参数。
top_ktop-k sampling;-1 表示使用服务端默认行为。
stop停止条件,可混合字符串和 token id。
seed / sampling_seed采样随机种子。

sample() 还支持:

  • include_prompt_logprobs
  • topk_prompt_logprobs
  • return_sampling_mask
  • return_old_logprob
  • return_moe_topk_indices

这些选项常用于 RL rollout、PPO/GRPO 训练或 MoE router replay。

计算 Logprobs

python
tokens = tokenizer.encode("Hello, world!", add_special_tokens=True)

logprobs = sampling_client.compute_logprobs(
    prompt=types.ModelInput.from_ints(tokens),
)

print(logprobs)

返回值与 prompt token 对齐:长度等于 prompt token 数量,第一项为 None,因为第一个 token 没有前文条件概率。

下一步

Weaver API 中文文档