Convert formatting to use ruff instead of yapf + isort (#26247)

Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
This commit is contained in:
Harry Mellor
2025-10-05 15:06:22 +01:00
committed by GitHub
parent 17edd8a807
commit d6953beb91
1508 changed files with 115244 additions and 94146 deletions

View File

@@ -14,15 +14,21 @@ from typing_extensions import deprecated
import vllm.envs as envs
from vllm.config import VllmConfig
from vllm.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from vllm.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from vllm.logger import init_logger
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
from vllm.multimodal import NestedTensors
from vllm.sequence import IntermediateTensors
from vllm.utils import (cdiv, direct_register_custom_op,
get_cuda_view_from_cpu_tensor, is_pin_memory_available,
is_uva_available)
from vllm.utils import (
cdiv,
direct_register_custom_op,
get_cuda_view_from_cpu_tensor,
is_pin_memory_available,
is_uva_available,
)
logger = init_logger(__name__)
@@ -65,12 +71,16 @@ class WeightsMapper:
def apply(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> Iterable[tuple[str, torch.Tensor]]:
return ((out_name, data) for name, data in weights
if (out_name := self._map_name(name)) is not None)
return (
(out_name, data)
for name, data in weights
if (out_name := self._map_name(name)) is not None
)
def apply_list(self, values: list[str]) -> list[str]:
return [
out_name for name in values
out_name
for name in values
if (out_name := self._map_name(name)) is not None
]
@@ -129,17 +139,20 @@ class AutoWeightsLoader:
self,
weights: Iterable[tuple[str, torch.Tensor]],
) -> Iterable[tuple[str, Iterable[tuple[str, torch.Tensor]]]]:
weights_by_parts = ((weight_name.split(".", 1), weight_data)
for weight_name, weight_data in weights)
weights_by_parts = (
(weight_name.split(".", 1), weight_data)
for weight_name, weight_data in weights
)
for prefix, group in itertools.groupby(weights_by_parts,
key=lambda x: x[0][0]):
for prefix, group in itertools.groupby(weights_by_parts, key=lambda x: x[0][0]):
yield (
prefix,
# Because maxsplit=1 in weight_name.split(...),
# the length of `parts` must either be 1 or 2
(("" if len(parts) == 1 else parts[1], weights_data)
for parts, weights_data in group),
(
("" if len(parts) == 1 else parts[1], weights_data)
for parts, weights_data in group
),
)
def _get_qualname(self, prefix: str, rest: str) -> str:
@@ -151,8 +164,9 @@ class AutoWeightsLoader:
return ".".join((prefix, rest))
def _can_skip(self, qualname: str) -> bool:
return (any(qualname.startswith(p) for p in self.skip_prefixes)
or any(substr in qualname for substr in self.skip_substrs))
return any(qualname.startswith(p) for p in self.skip_prefixes) or any(
substr in qualname for substr in self.skip_substrs
)
def _can_ignore_unexpected(self, qualname: str) -> bool:
iup = (qualname.startswith(p) for p in self.ignore_unexpected_prefixes)
@@ -181,24 +195,26 @@ class AutoWeightsLoader:
raise ValueError(
f"Attempted to load nested weight '{weight_qualname}' "
f"into a single parameter '{base_prefix}'")
f"into a single parameter '{base_prefix}'"
)
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, weight_data)
logger.debug("Loaded weight %s with shape %s", weight_qualname,
param.shape)
logger.debug("Loaded weight %s with shape %s", weight_qualname, param.shape)
yield weight_qualname
def _add_loadable_non_param_tensors(self, module: nn.Module,
child_params: dict[str, torch.Tensor]):
def _add_loadable_non_param_tensors(
self, module: nn.Module, child_params: dict[str, torch.Tensor]
):
"""
Add tensor names that are not in the model params that may be in the
safetensors, e.g., batch normalization stats.
"""
if isinstance(module, (
if isinstance(
module,
(
nn.BatchNorm1d,
nn.BatchNorm2d,
nn.BatchNorm3d,
@@ -206,10 +222,10 @@ class AutoWeightsLoader:
nn.LazyBatchNorm2d,
nn.LazyBatchNorm3d,
nn.SyncBatchNorm,
)):
),
):
module_state_dict = module.state_dict()
for stat_name in ("running_mean", "running_var",
"num_batches_tracked"):
for stat_name in ("running_mean", "running_var", "num_batches_tracked"):
child_params[stat_name] = module_state_dict[stat_name]
def _load_module(
@@ -229,8 +245,8 @@ class AutoWeightsLoader:
loaded_params = module_load_weights(weights)
if loaded_params is None:
logger.warning(
"Unable to collect loaded parameters "
"for module %s", module)
"Unable to collect loaded parameters for module %s", module
)
else:
yield from map(
lambda x: self._get_qualname(base_prefix, x),
@@ -253,17 +269,18 @@ class AutoWeightsLoader:
continue
yield from self._load_module(prefix,
child_modules[child_prefix],
child_weights)
yield from self._load_module(
prefix, child_modules[child_prefix], child_weights
)
elif child_prefix in child_params:
if self._can_skip(prefix):
logger.debug("Skipping param %s", prefix)
continue
yield from self._load_param(prefix, child_params[child_prefix],
child_weights)
yield from self._load_param(
prefix, child_params[child_prefix], child_weights
)
else:
can_skip_module = self._can_skip(prefix + ".")
can_skip_param = self._can_skip(prefix)
@@ -279,8 +296,10 @@ class AutoWeightsLoader:
continue
msg = (f"There is no module or parameter named '{prefix}' "
f"in {type(self.module).__name__}")
msg = (
f"There is no module or parameter named '{prefix}' "
f"in {type(self.module).__name__}"
)
raise ValueError(msg)
def load_weights(
@@ -292,8 +311,9 @@ class AutoWeightsLoader:
if mapper is not None:
weights = mapper.apply(weights)
# filter out weights with first-prefix/substr to skip in name
weights = ((name, weight) for name, weight in weights
if not self._can_skip(name))
weights = (
(name, weight) for name, weight in weights if not self._can_skip(name)
)
autoloaded_weights = set(self._load_module("", self.module, weights))
return autoloaded_weights
@@ -317,20 +337,17 @@ def init_vllm_registered_model(
hf_config = vllm_config.model_config.hf_config
if hf_config is not None:
vllm_config = vllm_config.with_hf_config(hf_config,
architectures=architectures)
vllm_config = vllm_config.with_hf_config(hf_config, architectures=architectures)
return initialize_model(vllm_config=vllm_config, prefix=prefix)
@overload
def flatten_bn(x: torch.Tensor) -> torch.Tensor:
...
def flatten_bn(x: torch.Tensor) -> torch.Tensor: ...
@overload
def flatten_bn(x: list[torch.Tensor]) -> list[torch.Tensor]:
...
def flatten_bn(x: list[torch.Tensor]) -> list[torch.Tensor]: ...
@overload
@@ -338,8 +355,7 @@ def flatten_bn(
x: Union[list[torch.Tensor], torch.Tensor],
*,
concat: Literal[True],
) -> torch.Tensor:
...
) -> torch.Tensor: ...
@overload
@@ -347,8 +363,7 @@ def flatten_bn(
x: Union[list[torch.Tensor], torch.Tensor],
*,
concat: bool = False,
) -> Union[list[torch.Tensor], torch.Tensor]:
...
) -> Union[list[torch.Tensor], torch.Tensor]: ...
def flatten_bn(
@@ -392,8 +407,7 @@ def _embedding_count_expression(embeddings: NestedTensors) -> str:
if isinstance(embeddings, torch.Tensor):
return " x ".join([str(dim) for dim in embeddings.shape[:-1]])
return " + ".join(
_embedding_count_expression(inner) for inner in embeddings)
return " + ".join(_embedding_count_expression(inner) for inner in embeddings)
def _merge_multimodal_embeddings(
@@ -421,8 +435,9 @@ def _merge_multimodal_embeddings(
# NOTE: This can avoid D2H sync (#22105), but fails to
# raise an error if is_multimodal.sum() < len(mm_embeds_flat)
inputs_embeds.masked_scatter_(is_multimodal.unsqueeze(-1),
mm_embeds_flat.to(dtype=input_dtype))
inputs_embeds.masked_scatter_(
is_multimodal.unsqueeze(-1), mm_embeds_flat.to(dtype=input_dtype)
)
except RuntimeError as e:
num_actual_tokens = len(mm_embeds_flat)
num_expected_tokens = is_multimodal.sum().item()
@@ -440,9 +455,11 @@ def _merge_multimodal_embeddings(
return inputs_embeds
@deprecated("`merge_multimodal_embeddings` has been replaced with "
"`SupportsMultiModal.get_input_embeddings` and will be "
"removed in v0.12.")
@deprecated(
"`merge_multimodal_embeddings` has been replaced with "
"`SupportsMultiModal.get_input_embeddings` and will be "
"removed in v0.12."
)
def merge_multimodal_embeddings(
input_ids: torch.Tensor,
inputs_embeds: torch.Tensor,
@@ -477,7 +494,7 @@ def merge_multimodal_embeddings(
if isinstance(placeholder_token_id, list):
is_multimodal = isin_list(input_ids, placeholder_token_id)
else:
is_multimodal = (input_ids == placeholder_token_id)
is_multimodal = input_ids == placeholder_token_id
return _merge_multimodal_embeddings(
inputs_embeds,
@@ -499,9 +516,7 @@ def isin_list(
class LayerFn(Protocol):
def __call__(self, prefix: str) -> torch.nn.Module:
...
def __call__(self, prefix: str) -> torch.nn.Module: ...
class PPMissingLayer(torch.nn.Identity):
@@ -544,8 +559,7 @@ def maybe_offload_to_cpu(module: torch.nn.Module) -> torch.nn.Module:
uva_available = is_uva_available()
if envs.VLLM_USE_V1:
assert uva_available, ("V1 CPU offloading requires"
" uva (pin memory) support")
assert uva_available, "V1 CPU offloading requires uva (pin memory) support"
uva_offloading = True
else:
uva_offloading = False
@@ -560,12 +574,14 @@ def maybe_offload_to_cpu(module: torch.nn.Module) -> torch.nn.Module:
break
# `torch.empty_like` does not support `pin_memory` argument
cpu_data = torch.empty_strided(size=p.data.size(),
stride=p.data.stride(),
dtype=p.data.dtype,
layout=p.data.layout,
device='cpu',
pin_memory=pin_memory)
cpu_data = torch.empty_strided(
size=p.data.size(),
stride=p.data.stride(),
dtype=p.data.dtype,
layout=p.data.layout,
device="cpu",
pin_memory=pin_memory,
)
cpu_data.copy_(p.data)
if not uva_offloading:
p.data = cpu_data
@@ -587,10 +603,7 @@ def maybe_offload_to_cpu(module: torch.nn.Module) -> torch.nn.Module:
k: v.to(device, non_blocking=True)
for k, v in module.state_dict().items()
}
output = functional_call(module,
device_state,
args=args,
kwargs=kwargs)
output = functional_call(module, device_state, args=args, kwargs=kwargs)
module.forward = forward
return output
@@ -609,14 +622,18 @@ def make_layers(
"""
from vllm.distributed.parallel_state import get_pp_group
from vllm.distributed.utils import get_pp_indices
start_layer, end_layer = get_pp_indices(num_hidden_layers,
get_pp_group().rank_in_group,
get_pp_group().world_size)
start_layer, end_layer = get_pp_indices(
num_hidden_layers, get_pp_group().rank_in_group, get_pp_group().world_size
)
modules = torch.nn.ModuleList(
[PPMissingLayer() for _ in range(start_layer)] + [
[PPMissingLayer() for _ in range(start_layer)]
+ [
maybe_offload_to_cpu(layer_fn(prefix=f"{prefix}.{idx}"))
for idx in range(start_layer, end_layer)
] + [PPMissingLayer() for _ in range(end_layer, num_hidden_layers)])
]
+ [PPMissingLayer() for _ in range(end_layer, num_hidden_layers)]
)
return start_layer, end_layer, modules
@@ -636,7 +653,7 @@ def get_pp_missing_layer_names(model: torch.nn.Module) -> list[str]:
# NOTE: the trailing dot is used to match the prefix of the layer.
# without the dot, we could match a layer that is not missing,
# e.g., 'encoder.layer.1' would match 'encoder.layer.11'
missing_layer_names.append(name + '.')
missing_layer_names.append(name + ".")
_model_to_pp_missing_layer_names[model_id] = missing_layer_names
return missing_layer_names
@@ -649,21 +666,22 @@ def is_pp_missing_parameter(name: str, model: torch.nn.Module) -> bool:
return any(
name.startswith(missing_layer_name)
for missing_layer_name in get_pp_missing_layer_names(model))
for missing_layer_name in get_pp_missing_layer_names(model)
)
def make_empty_intermediate_tensors_factory(keys: list[str], hidden_size: int):
def make_empty_intermediate_tensors(
batch_size: int,
dtype: torch.dtype,
device: torch.device,
) -> IntermediateTensors:
return IntermediateTensors({
key:
torch.zeros((batch_size, hidden_size), dtype=dtype, device=device)
for key in keys
})
return IntermediateTensors(
{
key: torch.zeros((batch_size, hidden_size), dtype=dtype, device=device)
for key in keys
}
)
return make_empty_intermediate_tensors
@@ -698,15 +716,20 @@ def extract_layer_index(layer_name: str, num_attn_module: int = 1) -> int:
except ValueError:
continue
if num_attn_module == 1 or "attn" not in layer_name:
assert len(int_vals) == 1, (f"layer name {layer_name} should"
" only contain one integer")
assert len(int_vals) == 1, (
f"layer name {layer_name} should only contain one integer"
)
return int_vals[0]
else:
assert len(int_vals) <= 2, (f"layer name {layer_name} should"
" contain most two integers")
layer_index = int_vals[0] * num_attn_module + int_vals[1] if len(
int_vals) == 2 else int_vals[0]
assert len(int_vals) <= 2, (
f"layer name {layer_name} should contain most two integers"
)
layer_index = (
int_vals[0] * num_attn_module + int_vals[1]
if len(int_vals) == 2
else int_vals[0]
)
return layer_index
@@ -720,19 +743,20 @@ def cast_overflow_tensors(
return tensors
def fast_topk(values: torch.Tensor, topk: int,
dim: int) -> tuple[torch.Tensor, torch.Tensor]:
def fast_topk(
values: torch.Tensor, topk: int, dim: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Optimized topk implementation that uses torch.max for k=1 case.
This function provides better performance for the common case of k=1
by using torch.max instead of the more general torch.topk.
Args:
values: Input tensor to find top-k values from
topk: Number of top values to return (k). Must be > 0.
dim: Dimension along which to compute topk
Returns:
Tuple of (values, indices) where values are the top-k values
and indices are their corresponding indices in the input tensor
@@ -791,5 +815,5 @@ direct_register_custom_op(
op_name="sequence_parallel_chunk_impl",
op_func=sequence_parallel_chunk_impl,
fake_impl=sequence_parallel_chunk_impl_fake,
tags=(torch.Tag.needs_fixed_stride_order, ),
tags=(torch.Tag.needs_fixed_stride_order,),
)