Skip to content

Capas lineales con tensor parallelism: ColumnParallelLinear / QKVParallelLinear / RowParallelLinear

源码版本v0.25.1

Responsabilidades

El tensor parallelism (TP) de vLLM parte un único nn.Linear en varias GPUs, cada una calcula una parte, y en el forward / backward el resultado completo se reensambla con all-gather / all-reduce. Esta capa de abstracción reside en tres clases centrales de vllm/model_executor/layers/linear.py: ColumnParallelLinear parte a lo largo de la dimensión de salida, RowParallelLinear a lo largo de la dimensión de entrada, y QKVParallelLinear gestiona los tres segmentos Q/K/V con distinto número de heads en attention (sobre todo en GQA / MQA). Las tres heredan de LinearBase y delegan la llamada GEMM real a quant_method.apply (linear.py:551-569), de modo que la misma lógica de paralelismo sirve para backends sin cuantizar, fp8, AWQ, GPTQ, etc.

Al instanciar cada capa lineal TP se calculan el tp_rank / tp_size de este rank, se divide output_size entre tp_size para obtener output_size_per_partition (linear.py:437-442) y se escribe esa información de partición en los atributos output_dim / input_dim de los parámetros; el weight_loader usa luego ese atributo para decidir qué segmento cortar. QKVParallelLinear añade encima la replicación de heads KV para GQA: cuando tp_size >= total_num_kv_heads, cada rank mantiene una copia completa de los heads KV (num_kv_head_replicas = tp_size / total_num_kv_heads) (linear.py:994-1001).

Motivación de diseño

¿Por qué introducir encima de LinearBase dos mecanismos, weight_loader / weight_loader_v2?

  • Corte distinto: ColumnParallelLinear corta por la dimensión de salida (cada rank toma un segmento de [A_1, A_2, ...]), RowParallelLinear por la dimensión de entrada (A = [A_1; A_2; ...], X = [X_1, X_2, ...]), por lo que los weight_loader de cada uno hacen narrow según output_dim o input_dim (linear.py:530-533)(linear.py:1650-1653).
  • v2 va por objetos de parámetros: el nuevo weight_loader_v2 envuelve el tensor en un BasevLLMParameter y llama directamente a param.load_column_parallel_weight / param.load_row_parallel_weight, encapsulando "qué segmento cortar" dentro del propio parámetro (linear.py:543-549)(linear.py:1663-1670); los métodos de cuantización pueden elegir v2 (WEIGHT_LOADER_V2_SUPPORTED) o la ruta vieja.
  • Fusión de los tres segmentos Q/K/V: QKVParallelLinear concatena los pesos de Q/K/V en la dimensión de salida formando una matriz grande; output_sizes = [q, k, v] calcula cada segmento con divide(..., tp_size) por separado (linear.py:1003-1008); MergedColumnParallelLinear es la versión generalizada, capaz de fusionar cualquier conjunto de capas column-parallel.
  • Replicación de KV en GQA: en QKVParallelLinear.__init__ se ejecuta 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) para que cada rank pueda calcular un head KV completo y el forward posterior no necesite sincronización extra.
  • Metadatos de shard: el weight_loader_v2 de ColumnParallelLinear también reconoce el shard_id (en QKV es "q" / "k" / "v") y según el shard id decide a qué offset del param escribir (linear.py:1023-1047).
  • bias solo en rank 0: en RowParallelLinear.forward, bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias (linear.py:1685-1688); sumar el bias después del reduce duplicaría, así que solo se añade en rank 0.
  • TP desactivable: con disable_tp=True se fuerza tp_rank=0 / tp_size=1 y la capa degenera en un Linear normal, útil para ciertos experts no paralelos de MoE (linear.py:438-439).

Archivos clave

Flujo de datos

Lo siguiente es la lógica central de ColumnParallelLinear.weight_loader: tras recibir el tensor completo desde el disco, lo corta por el offset tp_rank * shard_size para quedarse con el segmento de este rank y lo copia con copy_ al parámetro:

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 es la imagen espejo: cambia la dimensión del narrow de output_dim a input_dim (linear.py:1650-1653). DefaultModelLoader los invoca indirectamente al cargar los pesos vía model.load_weights —cada implementación de modelo tiene algo parecido a loaded_weights = param.weight_loader(param, loaded_weight).

En el forward, ColumnParallelLinear hace all-gather (con gather_output=True concatena el output_parallel de cada rank en un output completo) y RowParallelLinear hace all-reduce (con reduce_results=True suma los resultados del GEMM fragmentado); cuando se combinan ambos, ahorran una sincronización: QKVParallelLinear (column) → attention → o_proj (row) es la combinación clásica column-then-row.

Límites y fallos

  • shard_id inválido en QKV: validate_shard_id restringe el shard id de QKV a "q" / "k" / "v" o None; cualquier otro valor lanza ValueError directamente (linear.py:1023-1029).
  • RowParallelLinear con reduce_results=False + bias: no hacer reduce y a la vez añadir bias dentro del GEMM duplicaría el bias en cada rank; el código lanza ValueError (linear.py:1622-1626).
  • Forma de bloque FP8 incompatible: _maybe_allow_fp8_block_shape_mismatch relaja allow_fp8_block_shape_mismatch cuando weight_block_size no divide algún output_partition_sizes (linear.py:493-516); en caso contrario el kernel rechaza la forma.
  • bitsandbytes salta el narrow: con use_bitsandbytes_4bit, el checkpoint ya viene partido por rank, así que se hace copy_ directo sin narrow (linear.py:524-527).
  • scales como escalar de dimensión cero: checkpoints como AutoFP8 guardan las scales como tensores de dimensión cero; weight_loader las reemplaza con reshape(1) como fallback (linear.py:537-538).
  • capas row con input_is_parallel=False: si la entrada no está partida, RowParallelLinear.forward hace split_tensor_along_last_dim para tomar su segmento (linear.py:1676-1682).

Resumen

Las tres clases lineales paralelas forman el "lenguaje de partición de pesos" del TP de vLLM; el cálculo concreto lo gestiona quant_method. El contacto con carga de capas cuantizadas es: ColumnParallelLinear.__init__ llama a self.quant_method.create_weights(...) para insertar el weight_loader (v1 o v2) en cada parámetro; el contacto con DefaultModelLoader es que model.load_weights invoca weight_loader para cada tensor y corta el peso completo del disco en el segmento correspondiente a este rank.

Véase la documentación oficial: vLLM 文档 · README