共享 base,不共享 adapter 语义
对线性层 ,LoRA adapter 用 rank 的两矩阵:
等价 merged weight:
单 adapter 离线部署可把 merge 进 ;Multi-LoRA serving 要在每请求动态选择 adapter,不能每轮改写共享 base weight,否则并发请求会互相污染且 merge/copy 成本巨大。生产 kernel通常保留 base GEMM,再按 adapter groups/batched low-rank GEMM 累加 delta。
同一个 continuous batch 中的 adapter ownership
本页 batch 有 6 token rows:
adapter_ids = [finance, base, code, finance, code, base]
按 adapter 分组后各 2 rows:
base X:[2,8] → XW
finance X:[2,8] → XW + (α/r)(XA_fin)B_fin
code X:[2,8] → XW + (α/r)(XA_code)B_code
最后 scatter 回原 batch row order。Runtime 可避免显式 gather/scatter,用 token→adapter offsets 和 grouped GEMM/kernel;逻辑不变量仍是每 row 只应用其 adapter composition。
| 对象 | shape | 所有权 |
|---|---|---|
base W | [K,N] | worker/GPU shared,可能 TP shard/quantized |
adapter A_a,B_a | [K,r_a],[r_a,N] | adapter cache;只在需要的 workers/ranks resident |
| request adapter id | [sequence] 或每 token row metadata | scheduler/request state,不可与别的租户混淆 |
| low-rank intermediate | [M_a,r_a] | adapter group 临时,可 fused |
| output | [M,N] | 恢复原 request/token order |
同一 request 若支持 adapter composition(多个 LoRA 相加),identity 还需包含有序/规范化 composition、weights 与版本;不能只记录第一个 adapter name。
可运行的 grouped batch 等价实验
import platform,torch
torch.manual_seed(241)
DEVICE=torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE=torch.float32
def lora_forward(x,weight,a,b,alpha):
rank=a.shape[1]
return x@weight+(alpha/rank)*((x@a)@b)
def main():
tokens,hidden,output,rank=6,8,5,2
weight=torch.randn(hidden,output,device=DEVICE,dtype=DTYPE)
adapters={
"base":None,
"finance":(torch.randn(hidden,rank,device=DEVICE),
torch.randn(rank,output,device=DEVICE),4.),
"code":(torch.randn(hidden,rank,device=DEVICE),
torch.randn(rank,output,device=DEVICE),2.)}
x=torch.randn(tokens,hidden,device=DEVICE,dtype=DTYPE)
adapter_ids=["finance","base","code","finance","code","base"]
outputs=torch.empty(tokens,output,device=DEVICE); group_sizes={}
for adapter_id in sorted(set(adapter_ids)):
indices=torch.tensor([i for i,name in enumerate(adapter_ids) if name==adapter_id],device=DEVICE)
group=x[indices]
if adapters[adapter_id] is None: group_output=group@weight
else:
a,b,alpha=adapters[adapter_id]
group_output=lora_forward(group,weight,a,b,alpha)
merged=weight+(alpha/rank)*(a@b)
torch.testing.assert_close(group_output,group@merged,rtol=1e-5,atol=1e-5)
outputs[indices]=group_output; group_sizes[adapter_id]=len(indices)
rows=[]
for row,adapter_id in zip(x,adapter_ids):
if adapters[adapter_id] is None: rows.append(row@weight)
else:
a,b,alpha=adapters[adapter_id]
rows.append(lora_forward(row[None,:],weight,a,b,alpha).squeeze(0))
reference=torch.stack(rows)
torch.testing.assert_close(outputs,reference,rtol=1e-5,atol=1e-5)
adapter_params=sum(a.numel()+b.numel() for value in adapters.values()
if value for a,b,_ in [value])
print(f"python={platform.python_version()} torch={torch.__version__}")
print(f"device={DEVICE} dtype={DTYPE} X={tuple(x.shape)} W={tuple(weight.shape)} lora_rank={rank}")
print(f"adapter_ids={adapter_ids} grouped_batch_sizes={group_sizes}")
print(f"grouped_vs_row_reference_max_abs_diff={(outputs-reference).abs().max():.3e}")
print(f"base_parameters={weight.numel()} two_adapter_parameters={adapter_params}")
print(f"per_adapter_delta_elements={hidden*rank+rank*output} full_weight_elements={hidden*output}")
if __name__=="__main__":
with torch.inference_mode(): main()
完整文件位于 examples/multi_lora_serving.py。实际 grouped batch 与逐 row reference 最大差 1.907e-06;每 adapter delta 26 元素,full weight 40 元素。真实大层在 时比例更小,本例只为 shape 可读性。
代码没有量化 base/adapter、TP、不同 ranks 或专用 Punica/S-LoRA kernels,不能做性能结论。它证明 merged/on-the-fly 与 mixed-adapter row ownership。
容量、调度与 kernel 成本
一个 adapter 跨若干 target linear layers 的参数容量:
还要计 scales/zeros(若量化)、metadata、alignment 与 TP shards。若同时 resident 个 adapters,总容量线性增长到近 ;需要 adapter LRU/cache、load queue 和 admission。首个请求可能支付磁盘/CPU→GPU cold load,必须与 steady request 分开报告。
计算方面,base GEMM 共享;每 adapter group 增加 和 。Batch 中 distinct adapters 越多,每组 越小,kernel launch/低利用率越严重。专用 kernels 可把多个 low-rank groups 合并调度,但 rank、target modules、dtype/packing 不同会限制 batching。
Continuous Batching不能只按 sequence count 组 batch,还要考虑 adapter locality;为减少 adapter switches 把同 adapter 请求聚合,可能增加其他请求 queue/TTFT。Scheduler 应在 GPU token/KV budget之外考虑 adapter residency/load bandwidth。
量化、TP 与加载边界
- Base INT4/FP8 + LoRA FP16 是常见 mixed precision;
XW_q + LoRA的 accumulator/scale/epilogue须与 kernel支持一致。 - 直接把 FP16 LoRA merge 到 quantized base 后,需要重新量化或保留更高精度 merged weight;不能对 packed ints 原地相加浮点 delta。
- Column-parallel 要切 ,通常 replicated;row-parallel 的切分/归约契约不同。
- 模型加载要验证 adapter base model revision、target module names、rank、alpha、dtype、tensor shapes 与 hash;加载成功但 target module mapping 错会静默不生效或作用到错误层。
- Backend Dispatch确认实际走 multi-LoRA kernel还是逐 adapter Python/GEMM fallback。
多租户隔离与失败场景
Adapter 文件来自用户时不能作为 pickle 任意反序列化;应使用安全格式、大小/shape 配额、hash/signature 与 allowlisted target modules。租户 A 不应通过 adapter id 猜测、prefix cache 或 timing 访问租户 B 的 adapter/state。
常见失败:
- request 完成后 adapter refcount 未减,resident cache泄漏;
- adapter 被驱逐时仍有 in-flight CUDA kernel 使用 storage;
- batch reorder 后 adapter ids 未同步,token 应用错误 delta;
- prefix/KV cache key 未包含 adapter version;
- base model reload 后旧 adapter 仍标 resident;
- rank/alpha scaling重复或遗漏;
- 一个恶意 adapter 声明巨大 rank/大量 target modules,绕过容量 admission;
- adapter cold load 在关键线程同步执行,阻塞所有 requests。
Request Lifecycle应把 adapter acquisition/release 纳入幂等 cleanup;Scheduler Admission把 adapter bytes/load slots 与 KV blocks共同预算。