Improve linear.py to load sharded weights & remove the dependency of Parameters from vllm (#2784)

Co-authored-by: SangBin Cho rkooo567@gmail.com
This commit is contained in:
Lianmin Zheng
2025-01-07 23:29:10 -08:00
committed by GitHub
co-authored by SangBin Cho rkooo567@gmail.com
parent 694e41925e
commit 8a6906127a
15 changed files with 655 additions and 88 deletions
+165 -57
View File
@@ -18,14 +18,15 @@ from vllm.distributed import (
# workaround
from vllm.model_executor.layers.linear import LinearBase
from vllm.model_executor.parameter import (
from sglang.srt.layers.parameter import (
BasevLLMParameter,
PackedColumnParameter,
PackedvLLMParameter,
PerTensorScaleParameter,
RowvLLMParameter,
_ColumnvLLMParameter,
)
from sglang.srt.layers.quantization.base_config import (
QuantizationConfig,
QuantizeMethodBase,
@@ -94,6 +95,62 @@ def adjust_scalar_to_fused_array(param, loaded_weight, shard_id):
return param[shard_id], loaded_weight
def load_column_qkv_weight(
self, loaded_weight, num_heads, shard_id, shard_offset, shard_size, tp_rank
):
if (
isinstance(self, (PackedColumnParameter, PackedvLLMParameter))
and self.output_dim == self.packed_dim
):
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
shard_offset=shard_offset, shard_size=shard_size
)
param_data = self.data
shard_id = tp_rank if shard_id == "q" else tp_rank // num_heads
param_data = param_data.narrow(self.output_dim, shard_offset, shard_size)
loaded_weight = loaded_weight.narrow(
self.output_dim, shard_id * shard_size, shard_size
)
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
def load_column_parallel_weight(
self, loaded_weight: torch.Tensor, tp_rank, use_presharded_weights: bool = False
):
if isinstance(self, _ColumnvLLMParameter):
if not use_presharded_weights:
shard_size = self.data.shape[self.output_dim]
loaded_weight = loaded_weight.narrow(
self.output_dim, tp_rank * shard_size, shard_size
)
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
else:
self.data.copy_(loaded_weight)
def load_row_parallel_weight(
self, loaded_weight: torch.Tensor, tp_rank, use_presharded_weights: bool = False
):
if isinstance(self, RowvLLMParameter):
if not use_presharded_weights:
shard_size = self.data.shape[self.input_dim]
loaded_weight = loaded_weight.narrow(
self.input_dim, tp_rank * shard_size, shard_size
)
if len(loaded_weight.shape) == 0:
loaded_weight = loaded_weight.reshape(1)
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
else:
self.data.copy_(loaded_weight)
class LinearMethodBase(QuantizeMethodBase):
"""Base class for different (maybe quantized) linear methods."""
@@ -287,6 +344,8 @@ class ColumnParallelLinear(LinearBase):
quant_config: Optional[QuantizationConfig] = None,
output_sizes: Optional[List[int]] = None,
prefix: str = "",
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
):
super().__init__(
input_size, output_size, skip_bias_add, params_dtype, quant_config, prefix
@@ -295,7 +354,11 @@ class ColumnParallelLinear(LinearBase):
self.gather_output = gather_output
# Divide the weight matrix along the last dimension.
tp_size = get_tensor_model_parallel_world_size()
if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank()
if tp_size is None:
tp_size = get_tensor_model_parallel_world_size()
self.tp_rank, self.tp_size = tp_rank, tp_size
assert self.quant_method is not None
self.output_size_per_partition = divide(self.output_size, tp_size)
self.output_partition_sizes = [self.output_size_per_partition]
@@ -336,7 +399,6 @@ class ColumnParallelLinear(LinearBase):
self.register_parameter("bias", None)
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank()
output_dim = getattr(param, "output_dim", None)
# Special case for GGUF
@@ -356,7 +418,7 @@ class ColumnParallelLinear(LinearBase):
# no need to narrow here
if output_dim is not None and not use_bitsandbytes_4bit:
shard_size = param_data.shape[output_dim]
start_idx = tp_rank * shard_size
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
@@ -364,7 +426,9 @@ class ColumnParallelLinear(LinearBase):
if len(loaded_weight.shape) == 0:
loaded_weight = loaded_weight.reshape(1)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
def weight_loader_v2(self, param: Parameter, loaded_weight: torch.Tensor):
@@ -373,7 +437,7 @@ class ColumnParallelLinear(LinearBase):
if len(loaded_weight.shape) == 0:
assert loaded_weight.numel() == 1
loaded_weight = loaded_weight.reshape(1)
param.load_column_parallel_weight(loaded_weight=loaded_weight)
load_column_parallel_weight(param, loaded_weight, self.tp_rank)
def forward(self, input_):
bias = self.bias if not self.skip_bias_add else None
@@ -393,7 +457,7 @@ class ColumnParallelLinear(LinearBase):
s = f"in_features={self.input_size}"
s += f", output_features={self.output_size_per_partition}"
s += f", bias={self.bias is not None}"
s += f", tp_size={get_tensor_model_parallel_world_size()}"
s += f", tp_size={self.tp_size}"
s += f", gather_output={self.gather_output}"
return s
@@ -431,10 +495,18 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
use_presharded_weights: bool = False,
):
self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size()
if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank()
if tp_size is None:
tp_size = get_tensor_model_parallel_world_size()
self.tp_rank, self.tp_size = tp_rank, tp_size
assert all(output_size % tp_size == 0 for output_size in output_sizes)
self.use_presharded_weights = use_presharded_weights
super().__init__(
input_size=input_size,
output_size=sum(output_sizes),
@@ -444,6 +516,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
params_dtype=params_dtype,
quant_config=quant_config,
prefix=prefix,
tp_rank=tp_rank,
tp_size=tp_size,
)
def weight_loader(
@@ -463,12 +537,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
return
if is_gguf_weight:
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
output_dim = getattr(param, "output_dim", None)
shard_size = loaded_weight.size(output_dim) // tp_size
start_idx = tp_rank * shard_size
shard_size = loaded_weight.size(output_dim) // self.tp_size
start_idx = self.tp_rank * shard_size
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
@@ -494,7 +565,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
param_data, loaded_weight, 0
)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
return
current_shard_offset = 0
@@ -522,11 +595,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
return
assert loaded_shard_id < len(self.output_sizes)
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
if output_dim is not None:
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // self.tp_size
shard_size = self.output_sizes[loaded_shard_id] // self.tp_size
# Special case for quantization.
# If quantized, we need to adjust the offset and size to account
# for the packing.
@@ -545,10 +616,10 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
shard_offset = loaded_weight.shape[output_dim] * loaded_shard_id
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
start_idx = tp_rank * shard_size
start_idx = self.tp_rank * shard_size
# bitsandbytes loads the weights of the specific portion
# no need to narrow here
if not use_bitsandbytes_4bit:
if not use_bitsandbytes_4bit and not self.use_presharded_weights:
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
# Special case for AQLM codebooks.
elif is_metadata:
@@ -572,7 +643,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
"the same for all partitions."
)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
def _load_fused_module_from_checkpoint(
@@ -629,26 +702,27 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
assert loaded_shard_id < len(self.output_sizes)
tp_size = get_tensor_model_parallel_world_size()
if isinstance(param, BlockQuantScaleParameter):
weight_block_size = self.quant_method.quant_config.weight_block_size
block_n, _ = weight_block_size[0], weight_block_size[1]
shard_offset = (
(sum(self.output_sizes[:loaded_shard_id]) + block_n - 1) // block_n
) // tp_size
) // self.tp_size
shard_size = (
(self.output_sizes[loaded_shard_id] + block_n - 1) // block_n // tp_size
(self.output_sizes[loaded_shard_id] + block_n - 1)
// block_n
// self.tp_size
)
else:
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // self.tp_size
shard_size = self.output_sizes[loaded_shard_id] // self.tp_size
param.load_merged_column_weight(
loaded_weight=loaded_weight,
shard_id=loaded_shard_id,
shard_offset=shard_offset,
shard_size=shard_size,
use_presharded_weights=self.use_presharded_weights,
)
@@ -689,6 +763,8 @@ class QKVParallelLinear(ColumnParallelLinear):
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
):
self.hidden_size = hidden_size
self.head_size = head_size
@@ -697,7 +773,11 @@ class QKVParallelLinear(ColumnParallelLinear):
total_num_kv_heads = total_num_heads
self.total_num_kv_heads = total_num_kv_heads
# Divide the weight matrix along the last dimension.
tp_size = get_tensor_model_parallel_world_size()
if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank()
if tp_size is None:
tp_size = get_tensor_model_parallel_world_size()
self.tp_rank, self.tp_size = tp_rank, tp_size
self.num_heads = divide(self.total_num_heads, tp_size)
if tp_size >= self.total_num_kv_heads:
self.num_kv_heads = 1
@@ -724,6 +804,8 @@ class QKVParallelLinear(ColumnParallelLinear):
params_dtype=params_dtype,
quant_config=quant_config,
prefix=prefix,
tp_rank=tp_rank,
tp_size=tp_size,
)
def _get_shard_offset_mapping(self, loaded_shard_id: str):
@@ -814,13 +896,24 @@ class QKVParallelLinear(ColumnParallelLinear):
shard_offset = (shard_offset + block_n - 1) // block_n
shard_size = (shard_size + block_n - 1) // block_n
param.load_qkv_weight(
loaded_weight=loaded_weight,
num_heads=self.num_kv_head_replicas,
shard_id=loaded_shard_id,
shard_offset=shard_offset,
shard_size=shard_size,
)
if isinstance(param, _ColumnvLLMParameter):
load_column_qkv_weight(
param,
loaded_weight,
num_heads=self.num_kv_head_replicas,
shard_id=loaded_shard_id,
shard_offset=shard_offset,
shard_size=shard_size,
tp_rank=self.tp_rank,
)
else:
param.load_qkv_weight(
loaded_weight=loaded_weight,
num_heads=self.num_kv_head_replicas,
shard_id=loaded_shard_id,
shard_offset=shard_offset,
shard_size=shard_size,
)
def weight_loader(
self,
@@ -840,12 +933,9 @@ class QKVParallelLinear(ColumnParallelLinear):
return
if is_gguf_weight:
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
output_dim = getattr(param, "output_dim", None)
shard_size = loaded_weight.size(output_dim) // tp_size
start_idx = tp_rank * shard_size
shard_size = loaded_weight.size(output_dim) // self.tp_size
start_idx = self.tp_rank * shard_size
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
@@ -872,7 +962,9 @@ class QKVParallelLinear(ColumnParallelLinear):
param_data, loaded_weight, 0
)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
return
shard_offsets = [
@@ -934,7 +1026,6 @@ class QKVParallelLinear(ColumnParallelLinear):
self.weight_loader(param, loaded_weight_shard, shard_id)
return
tp_rank = get_tensor_model_parallel_rank()
assert loaded_shard_id in ["q", "k", "v"]
# If output dim is defined, use the default loading process.
@@ -984,9 +1075,9 @@ class QKVParallelLinear(ColumnParallelLinear):
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
if loaded_shard_id == "q":
shard_id = tp_rank
shard_id = self.tp_rank
else:
shard_id = tp_rank // self.num_kv_head_replicas
shard_id = self.tp_rank // self.num_kv_head_replicas
start_idx = shard_id * shard_size
# bitsandbytes loads the weights of the specific portion
@@ -1014,7 +1105,9 @@ class QKVParallelLinear(ColumnParallelLinear):
"for all partitions."
)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
@@ -1055,6 +1148,9 @@ class RowParallelLinear(LinearBase):
reduce_results: bool = True,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
use_presharded_weights: bool = False,
):
super().__init__(
input_size, output_size, skip_bias_add, params_dtype, quant_config, prefix
@@ -1064,10 +1160,14 @@ class RowParallelLinear(LinearBase):
self.reduce_results = reduce_results
# Divide the weight matrix along the last dimension.
self.tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank()
if tp_size is None:
tp_size = get_tensor_model_parallel_world_size()
self.tp_rank, self.tp_size = tp_rank, tp_size
self.input_size_per_partition = divide(input_size, self.tp_size)
assert self.quant_method is not None
self.use_presharded_weights = use_presharded_weights
self.quant_method.create_weights(
layer=self,
@@ -1101,8 +1201,6 @@ class RowParallelLinear(LinearBase):
self.register_parameter("bias", None)
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
input_dim = getattr(param, "input_dim", None)
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
@@ -1116,15 +1214,19 @@ class RowParallelLinear(LinearBase):
if is_gguf_weight and isinstance(param, UninitializedParameter):
weight_shape = list(loaded_weight.shape)
if input_dim:
weight_shape[input_dim] = weight_shape[input_dim] // tp_size
weight_shape[input_dim] = weight_shape[input_dim] // self.tp_size
param.materialize(tuple(weight_shape), dtype=loaded_weight.dtype)
param_data = param.data
# bitsandbytes loads the weights of the specific portion
# no need to narrow here
if input_dim is not None and not use_bitsandbytes_4bit:
if (
input_dim is not None
and not use_bitsandbytes_4bit
and not self.use_presharded_weights
):
shard_size = param_data.shape[input_dim]
start_idx = tp_rank * shard_size
start_idx = self.tp_rank * shard_size
loaded_weight = loaded_weight.narrow(input_dim, start_idx, shard_size)
# Special case for loading scales off disk, which often do not
@@ -1132,7 +1234,9 @@ class RowParallelLinear(LinearBase):
if len(loaded_weight.shape) == 0:
loaded_weight = loaded_weight.reshape(1)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"{param_data.shape=}, {loaded_weight.shape=}"
param_data.copy_(loaded_weight)
def weight_loader_v2(self, param: BasevLLMParameter, loaded_weight: torch.Tensor):
@@ -1143,17 +1247,21 @@ class RowParallelLinear(LinearBase):
assert loaded_weight.numel() == 1
loaded_weight = loaded_weight.reshape(1)
param.load_row_parallel_weight(loaded_weight=loaded_weight)
load_row_parallel_weight(
param,
loaded_weight,
self.tp_rank,
use_presharded_weights=self.use_presharded_weights,
)
def forward(self, input_):
if self.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.tp_size
)
input_parallel = splitted_input[tp_rank].contiguous()
input_parallel = splitted_input[self.tp_rank].contiguous()
# Matrix multiply.
assert self.quant_method is not None