Skip to content

テンサル並列線形レイヤー:ColumnParallelLinear / QKVParallelLinear / RowParallelLinear

源码版本v0.25.1

役割

vLLM のテンサル並列 (tensor parallelism, TP) は 1 つの nn.Linear を複数の GPU に分割して各々が一部を計算し、前向き・逆向きで all-gather / all-reduce で完全な結果を組み立てます。この抽象レイヤーは vllm/model_executor/layers/linear.py の 3 つのコアクラスにあります:ColumnParallelLinear は出力次元で切り、RowParallelLinear は入力次元で切り、QKVParallelLinear は attention の Q/K/V の 3 つのセグメントの異なる head 数 (特に GQA / MQA) を処理します。これらはすべて LinearBase を継承し、実際の GEMM 呼び出しは quant_method.apply に委譲します (linear.py:551-569)。そのため同じ並列ロジックで未量子化、fp8、AWQ、GPTQ など各種バックエンドを動かせます。

各 TP 線形レイヤーはインスタンス化時に自 rank の tp_rank / tp_size を計算し、output_sizetp_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 の完全なコピーを 1 つ持ちます (num_kv_head_replicas = tp_size / total_num_kv_heads) (linear.py:994-1001)。

設計動機

なぜ LinearBase の上に weight_loader / weight_loader_v2 の 2 つの仕組みを導入するのか?

  • 切り方が異なる:ColumnParallelLinear は出力次元で切り (各 rank が [A_1, A_2, ...] の 1 つのセグメントを持つ)、RowParallelLinear は入力次元で切ります (A = [A_1; A_2; ...]X = [X_1, X_2, ...])。そのため 2 つの weight_loader はそれぞれ output_dim または input_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 の 3 セグメント融合:QKVParallelLinear は Q/K/V の重みを出力次元で 1 つの大きな行列に結合します。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_v2shard_id (QKV では "q" / "k" / "v") も識別し、shard id に応じて param のどの offset に書き込むかを決めます (linear.py:1023-1047)。
  • bias は rank 0 のみ:RowParallelLinear.forward の中で bias_ = 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)。これら 2 つの関数は DefaultModelLoader が重みをロードするときに model.load_weights を経由して間接的に呼び出します。各モデル実装には loaded_weights = param.weight_loader(param, loaded_weight) のような書き方があります。

前向き計算時、ColumnParallelLinear は all-gather を使い (gather_output=True のとき各 rank の output_parallel を完全な出力に結合)、RowParallelLinear は all-reduce を使います (reduce_results=True のとき分割 GEMM の結果を sum)。両者を組み合わせると同期を 1 回省けます: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 が 1 つずつ bias を足してしまい結果が不正になります。コード内では raise ValueError します (linear.py:1622-1626)。
  • FP8 block shape の不一致:_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 が 0 次元スカラ:AutoFP8 などの checkpoint は scale を 0 次元テンサルで保存し、weight_loaderreshape(1) でフォールバックします (linear.py:537-538)。
  • input_is_parallel=False の row レイヤー:入力が未分割のとき、RowParallelLinear.forward 自身が split_tensor_along_last_dim で自 rank の分を取得します (linear.py:1676-1682)。

まとめ

3 つの並列線形クラスが vLLM TP の「重みスライス言語」を構成し、具体的な計算方法は quant_method に任せます。量子化レイヤーのロード との接続は ColumnParallelLinear.__init__self.quant_method.create_weights(...) を呼び出して weight_loader (v1 または v2) を各パラメータに取り付けること、DefaultModelLoader との接続は model.load_weights でテンサルごとに weight_loader を呼び出してディスク上の完全な重みを自 rank のセグメントに切り出すことです。

公式資料: vLLM 文档 · README