训练与采样
本页介绍 Weaver 的核心 SDK 工作流:创建 client、准备 Datum、执行训练、导出权重、采样与计算 logprobs。
工作流概览
典型流程如下:
- 创建
ServiceClient,连接 Weaver 服务并维护 session。 - 调用
create_model()创建TrainingClient。 - 把样本转换为
types.Datum。 - 调用
forward_backward()累积梯度。 - 调用
optim_step()更新模型参数。 - 导出 sampler 权重,并通过
SamplingClient.sample()做评估或 rollout。
创建 Client
ServiceClient
ServiceClient 是 SDK 的入口。它会创建或复用 session,并在上下文管理器退出时清理本次创建的模型实例。
from weaver import ServiceClient
with ServiceClient() as service_client:
print(service_client.session_id)常用参数:
| 参数 | 说明 |
|---|---|
api_key | Weaver API key;不传时读取 WEAVER_API_KEY。 |
base_url | Weaver 服务地址;不传时使用 SDK 默认地址。 |
session_id | 复用已有 session。 |
default_tags | 给 session 写入默认标签,便于实验追踪。 |
TrainingClient
创建训练模型:
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 | 训练模式;可传 lora 或 full_ft。不传时服务端默认使用 LoRA。 |
lora_config | LoRA 配置;默认 rank 为 32,并训练 attention、MLP 和 unembedding。 |
performance_tier | 可选吞吐档位,例如 normal、fast、flash;更高档位通常意味着更高吞吐和更高成本。 |
user_metadata | 透传给服务端的实验元信息。 |
Tokenizer
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_tokens 和 weights:
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 的第一步。
result = training_client.forward(
datums,
"forward_logprob",
wait=True,
)forward_backward()
执行前向和反向,并把梯度累积在训练后端,等待后续 optim_step() 更新参数。
result = training_client.forward_backward(
datums,
"cross_entropy",
wait=True,
)参数说明:
| 参数 | 说明 |
|---|---|
data | Sequence[types.Datum]。 |
loss_fn | loss 名称,例如 cross_entropy、importance_sampling。 |
loss_fn_config | 可选 loss 配置。 |
metadata | 可选请求级元信息。 |
wait | True 时阻塞等待结果;False 时返回 OperationHandle。 |
返回结果通常包含每条样本输出和聚合指标:
{
"result": {
"loss_fn_outputs": [...],
"metrics": {
"loss": 0.5
}
}
}optim_step()
用累积梯度执行优化器更新:
training_client.optim_step(
types.AdamParams(learning_rate=1e-4),
wait=True,
)AdamParams 默认值:
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:
handle = training_client.forward_backward(
datums,
"cross_entropy",
wait=False,
)
result = handle.result()这适合把数据准备、rollout 或评估与远端训练并行起来。
完整训练循环
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。
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:
sampling_client = training_client.save_weights_and_get_sampling_client(
name="my-model-step-100",
)随后调用 sample():
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
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_p | nucleus sampling 参数。 |
top_k | top-k sampling;-1 表示使用服务端默认行为。 |
stop | 停止条件,可混合字符串和 token id。 |
seed / sampling_seed | 采样随机种子。 |
sample() 还支持:
include_prompt_logprobstopk_prompt_logprobsreturn_sampling_maskreturn_old_logprobreturn_moe_topk_indices
这些选项常用于 RL rollout、PPO/GRPO 训练或 MoE router replay。
计算 Logprobs
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 没有前文条件概率。