Skip to content

張量並行線性層:ColumnParallelLinear / QKVParallelLinear / RowParallelLinear

源码版本v0.25.1

職責

vLLM 的張量並行 (tensor parallelism, TP) 把單個 nn.Linear 拆成多個 GPU 各算一部分,前後向用 all-gather / all-reduce 拼回完整結果。這一層抽象落在 vllm/model_executor/layers/linear.py 的三個核心類上:ColumnParallelLinear 沿輸出維切、RowParallelLinear 沿輸入維切、QKVParallelLinear 處理 attention 裡 Q/K/V 三段不同的 head 數(尤其是 GQA / MQA)。它們都繼承自 LinearBase,真正的 GEMM 呼叫委託給 quant_method.apply(linear.py:551-569),所以同一套並行邏輯能跑未量化、fp8、AWQ、GPTQ 等各種後端。

每個 TP 線性層實例化時都會算出本 rank 的 tp_rank / tp_size,把 output_size 除以 tp_size 得到 output_size_per_partition(linear.py:437-442),並把這個分區信息寫到參數的 output_dim / input_dim 屬性上,後續 weight_loader 憑這個屬性決定切哪一段。QKVParallelLinear 在此基礎上額外處理 GQA 的 KV head 複製——當 tp_size >= total_num_kv_heads 時,每個 rank 持有一整份 KV head 的副本(num_kv_head_replicas = tp_size / total_num_kv_heads),(linear.py:994-1001)。

設計動機

為什麼要在 LinearBase 之上引入 weight_loader / weight_loader_v2 兩套機制?

  • 切法不同:ColumnParallelLinear 在輸出維切(每個 rank 拿 [A_1, A_2, ...] 的一段),RowParallelLinear 在輸入維切(A = [A_1; A_2; ...],X = [X_1, X_2, ...]),所以兩類的 weight_loader 分別按 output_diminput_dimnarrow(linear.py:530-533)(linear.py:1650-1653)。
  • v2 走參數物件:新版 weight_loader_v2 把張量包裝成 BasevLLMParameter,直接調 param.load_column_parallel_weight / param.load_row_parallel_weight,把"切哪段"封裝在參數物件內部(linear.py:543-549)(linear.py:1663-1670),量化方法可以選擇走 v2(WEIGHT_LOADER_V2_SUPPORTED)或老路。
  • Q/K/V 三段融合:QKVParallelLinear 把 Q/K/V 的權重在輸出維拼成一個大矩陣,output_sizes = [q, k, v] 每段單獨算 divide(..., tp_size)(linear.py:1003-1008),MergedColumnParallelLinear 是它的通用版,任何幾個列並行的層都能融合。
  • GQA KV 複製:QKVParallelLinear.__init__if tp_size >= total_num_kv_heads: self.num_kv_heads = 1; self.num_kv_head_replicas = divide(tp_size, total_num_kv_heads)(linear.py:996-1001),保證每個 rank 都能算完整 KV head,後續 forward 不需要額外同步。
  • Shard 元信息:ColumnParallelLinearweight_loader_v2 還要識別 shard_id(QKV 裡是 "q" / "k" / "v"),根據 shard id 決定寫入 param 的哪一段 offset(linear.py:1023-1047)。
  • bias 只在 rank 0 加:RowParallelLinear.forwardbias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias(linear.py:1685-1688),reduce 之後再加 bias 會重複,所以只在 rank 0 加。
  • 可關閉 TP:disable_tp=True 時讓 tp_rank=0 / tp_size=1,該層就退化成普通 Linear,MoE 的某些非並行 expert 用得上(linear.py:438-439)。

關鍵檔案

資料流

下面這段是 ColumnParallelLinear.weight_loader 的核心邏輯——拿到磁盤上的完整張量後,按 tp_rank * shard_size 偏移切出本 rank 需要的那一段,再 copy_ 進參數:

python
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
    output_dim = getattr(param, "output_dim", None)

    is_sharded_weight = getattr(param, "is_sharded_weight", False)
    use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
    # bitsandbytes loads the weights of the specific portion
    # no need to narrow
    is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit

    param_data = param.data
    if output_dim is not None and not is_sharded_weight:
        shard_size = param_data.shape[output_dim]
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)

    # Special case for loading scales off disk, which often do not
    # have a shape (such as in the case of AutoFP8).
    if len(loaded_weight.shape) == 0:
        loaded_weight = loaded_weight.reshape(1)

    assert param_data.shape == loaded_weight.shape
    param_data.copy_(loaded_weight)

(linear.py:520-541)RowParallelLinear.weight_loader 鏡像這段,只是 narrow 的維度從 output_dim 換成 input_dim(linear.py:1650-1653)。這兩個函數被 DefaultModelLoader 在載入權重時通過 model.load_weights 間接呼叫——每個模型實現裡都有類似 loaded_weights = param.weight_loader(param, loaded_weight) 的寫法。

前向時,ColumnParallelLinear 走 all-gather(gather_output=True 時把每個 rank 的 output_parallel 拼成完整 output),RowParallelLinear 走 all-reduce(reduce_results=True 時把分塊 GEMM 結果求和);兩者搭配使用時正好可以省掉一次同步:QKVParallelLinear(column)→ attention → o_proj(row)是一次經典的 column-then-row 組合。

邊界與失敗

  • QKV shard_id 非法:validate_shard_id 限制 QKV 的 shard id 只能是 "q" / "k" / "v"None,其他值直接 raise ValueError(linear.py:1023-1029)。
  • RowParallelLinear reduce_results=False + 加 bias:不 reduce 又在 GEMM 里加 bias 會導致每個 rank 都加了一份,結果錯誤,代碼裡 raise ValueError(linear.py:1622-1626)。
  • FP8 塊形狀不匹配:_maybe_allow_fp8_block_shape_mismatchweight_block_size 不能整除某個 output_partition_sizes 時,放開 allow_fp8_block_shape_mismatch(linear.py:493-516),否則 kernel 會拒絕。
  • bitsandbytes 跳過 narrow:use_bitsandbytes_4bit 時 checkpoint 已經是按 rank 分好的,直接 copy_ 不能 narrow(linear.py:524-527)。
  • scales 是零維標量:AutoFP8 等 checkpoint 把 scale 存成零維張量,weight_loaderreshape(1) 兜底(linear.py:537-538)。
  • input_is_parallel=False 的 row 層:輸入未切分時,RowParallelLinear.forward 自己 split_tensor_along_last_dim 取本 rank 那一份(linear.py:1676-1682)。

小結

三個並行線性類構成了 vLLM TP 的"權重切片語言",具體怎麼算交給 quant_method。跟 量化層載入 的銜接是:ColumnParallelLinear.__init__ 裡調 self.quant_method.create_weights(...),把 weight_loader(v1 或 v2)塞進每個參數;跟 DefaultModelLoader 的銜接是 model.load_weights 裡逐張量調 weight_loader,把磁盤上的完整權重切成本 rank 那一段。

對照官方資料:vLLM 文件 · README