テンサル並列線形レイヤー:ColumnParallelLinear / QKVParallelLinear / RowParallelLinear
役割
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_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 の完全なコピーを 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_dimでnarrowします (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 メタ情報:
ColumnParallelLinearのweight_loader_v2はshard_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)。
主要ファイル
LinearMethodBase:141-180— 量子化メソッドの抽象基底クラス、create_weights/apply/process_weights_after_loadingの 3 つのインターフェースを定義。LinearBase:231-288— すべての線形レイヤーの親クラス、update_param_tp_statusがtp_rank/tp_sizeを各BasevLLMParameterに書き込み。ColumnParallelLinear.__init__:423-491— 出力次元で切り、output_size_per_partition = divide(output_size, tp_size)。ColumnParallelLinear weight_loader:520-549— v1 はoutput_dimで narrow、v2 はparam.load_column_parallel_weightに委譲。ColumnParallelLinear.forward:551-569—quant_method.apply→ オプションでtensor_model_parallel_all_gather。QKVParallelLinear.__init__:942-1021— GQA の KV head 複製を処理、Q/K/V の 3 セグメントのoutput_sizesをまとめてoutput_sizeに。QKVParallelLinear shard mapping:1023-1047—_get_shard_offset_mapping/_get_shard_size_mapping、weight_loader に param のどのセグメントに書くかを伝える。QKVParallelLinear fused checkpoint:1049-1097— Phi-3 のようなディスク上で融合済みの QKV を自動分割し、セグメントごとにweight_loader_v2を呼び出します。RowParallelLinear.__init__:1572-1626— 入力次元で切り、input_size_per_partition = divide(input_size, tp_size)。RowParallelLinear.forward:1672-1698— 入力が並列化されていない場合はまず split、GEMM 後にtensor_model_parallel_all_reduce。MergedColumnParallelLinear:580-636— 複数の列並列レイヤーの融合 (QKV の汎用版)、output_sizesリストが各セグメントを決定。
データフロー
以下は ColumnParallelLinear.weight_loader のコアロジックです。ディスク上の完全なテンサルを受け取り、tp_rank * shard_size のオフセットで自 rank が必要なセグメントを切り出し、copy_ でパラメータに書き込みます:
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_mismatchはweight_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_loaderはreshape(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 のセグメントに切り出すことです。