Skip to content

Commit 6ccccad

Browse files
authored
[Lint]Style: Convert vllm-ascend/ to ruff format(Batch vllm-project#5) (vllm-project#5996)
### What this PR does / why we need it? **Scope of Changes**: | File Path | | :--- | | `.../distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py` | | `vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/backend/backend.py` | | ` .../distributed/kv_transfer/kv_pool/ascend_store/backend/memcache_backend.py` | | ` .../distributed/kv_transfer/kv_pool/ascend_store/backend/mooncake_backend.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/config_data.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/kv_transfer.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_scheduler.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py` | | ` .../distributed/kv_transfer/kv_pool/cpu_offload/cpu_kv_cache_manager.py` | | ` .../distributed/kv_transfer/kv_pool/cpu_offload/cpu_offload_connector.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/cpu_offload/metadata.py` | | ` vllm_ascend/distributed/kv_transfer/kv_pool/ucm_connector.py` | | ` vllm_ascend/distributed/kv_transfer/utils/mooncake_transfer_engine.py` | | ` vllm_ascend/distributed/kv_transfer/utils/utils.py` | | ` vllm_ascend/kv_offload/cpu_npu.py` | | ` vllm_ascend/kv_offload/npu.py` | | ` vllm_ascend/lora/lora_ops.py` | | ` vllm_ascend/lora/punica_npu.py` | | ` vllm_ascend/lora/utils.py` | ### Does this PR introduce _any_ user-facing change? ### How was this patch tested? - vLLM version: v0.13.0 - vLLM main: vllm-project@2c24bc6 --------- Signed-off-by: MrZ20 <2609716663@qq.com> Signed-off-by: SILONG ZENG <2609716663@qq.com>
1 parent 7faa687 commit 6ccccad

21 files changed

Lines changed: 865 additions & 1033 deletions

mypy.ini

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
[mypy]
22
; warn_return_any = True
33
warn_unused_configs = True
4+
; disable errors about unchecked annotations for now.
5+
disable_error_code = annotation-unchecked
46

57
; Suppress all missing import errors from torch_npu for mypy.
68
[mypy-torch_npu.*]
@@ -31,4 +33,4 @@ ignore_missing_imports = True
3133
ignore_missing_imports = True
3234

3335
[mypy-ucm.*]
34-
ignore_missing_imports = True
36+
ignore_missing_imports = True

pyproject.toml

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -51,11 +51,6 @@ line-length = 120
5151
# Folder to be modified
5252
exclude = [
5353
"tests/**",
54-
# (5)
55-
"vllm_ascend/distributed/kv_transfer/kv_pool/**",
56-
"vllm_ascend/distributed/kv_transfer/utils/**",
57-
"vllm_ascend/kv_offload/**",
58-
"vllm_ascend/lora/**",
5954
# (7)
6055
"vllm_ascend/quantization/**",
6156
"vllm_ascend/sample/*.py",

vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py

Lines changed: 28 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,10 @@
11
import threading
2-
from typing import Any, Optional
2+
from typing import Any
33

44
import torch
55
import zmq
66
from vllm.config import VllmConfig
7-
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
8-
KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole)
7+
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole
98
from vllm.forward_context import ForwardContext
109
from vllm.logger import logger
1110
from vllm.utils.network_utils import make_zmq_socket
@@ -17,40 +16,35 @@
1716
from vllm.v1.serial_utils import MsgpackDecoder
1817

1918
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_scheduler import (
20-
KVPoolScheduler, get_zmq_rpc_path_lookup)
21-
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_worker import \
22-
KVPoolWorker
19+
KVPoolScheduler,
20+
get_zmq_rpc_path_lookup,
21+
)
22+
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_worker import KVPoolWorker
2323

2424

2525
class AscendStoreConnector(KVConnectorBase_V1):
26-
27-
def __init__(self,
28-
vllm_config: VllmConfig,
29-
role: KVConnectorRole,
30-
kv_cache_config: Optional[KVCacheConfig] = None):
31-
super().__init__(vllm_config=vllm_config,
32-
role=role,
33-
kv_cache_config=kv_cache_config)
26+
def __init__(self, vllm_config: VllmConfig, role: KVConnectorRole, kv_cache_config: KVCacheConfig | None = None):
27+
super().__init__(vllm_config=vllm_config, role=role, kv_cache_config=kv_cache_config)
3428
self.kv_role = vllm_config.kv_transfer_config.kv_role
3529

36-
self.use_layerwise = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
37-
"use_layerwise", False)
30+
self.use_layerwise = vllm_config.kv_transfer_config.kv_connector_extra_config.get("use_layerwise", False)
3831
self.consumer_is_to_put = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
39-
"consumer_is_to_put", False)
32+
"consumer_is_to_put", False
33+
)
4034

4135
connector_name = vllm_config.kv_transfer_config.kv_connector
4236
if connector_name == "MooncakeConnectorStoreV1":
4337
logger.warning(
44-
"It is recommended to use the AscendStoreConnector, as the MoonCakeStoreConnector will be removed in the future."
38+
"It is recommended to use the AscendStoreConnector, "
39+
"as the MoonCakeStoreConnector will be removed in the future."
4540
)
4641

4742
self.kv_caches: dict[str, torch.Tensor] = {}
4843

4944
self.sended_but_unfinished_reqs: set[str] = set()
5045

5146
if role == KVConnectorRole.SCHEDULER:
52-
self.connector_scheduler = KVPoolScheduler(vllm_config,
53-
self.use_layerwise)
47+
self.connector_scheduler = KVPoolScheduler(vllm_config, self.use_layerwise)
5448
else:
5549
self.connector_worker = KVPoolWorker(
5650
vllm_config,
@@ -59,27 +53,19 @@ def __init__(self,
5953

6054
assert self.connector_worker is not None
6155
if vllm_config.parallel_config.rank == 0:
62-
self.lookup_server = LookupKeyServer(self.connector_worker,
63-
vllm_config,
64-
self.use_layerwise)
56+
self.lookup_server = LookupKeyServer(self.connector_worker, vllm_config, self.use_layerwise)
6557

6658
############################################################
6759
# Scheduler Side Methods
6860
############################################################
6961

70-
def get_num_new_matched_tokens(
71-
self, request: "Request",
72-
num_computed_tokens: int) -> tuple[int, bool]:
62+
def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> tuple[int, bool]:
7363
assert self.connector_scheduler is not None
74-
return self.connector_scheduler.get_num_new_matched_tokens(
75-
request, num_computed_tokens)
64+
return self.connector_scheduler.get_num_new_matched_tokens(request, num_computed_tokens)
7665

77-
def update_state_after_alloc(self, request: "Request",
78-
blocks: "KVCacheBlocks",
79-
num_external_tokens: int):
66+
def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int):
8067
assert self.connector_scheduler is not None
81-
return self.connector_scheduler.update_state_after_alloc(
82-
request, blocks, num_external_tokens)
68+
return self.connector_scheduler.update_state_after_alloc(request, blocks, num_external_tokens)
8369

8470
def build_connector_meta(
8571
self,
@@ -92,7 +78,7 @@ def request_finished(
9278
self,
9379
request: "Request",
9480
block_ids: list[int],
95-
) -> tuple[bool, Optional[dict[str, Any]]]:
81+
) -> tuple[bool, dict[str, Any] | None]:
9682
assert self.connector_scheduler is not None
9783
return self.connector_scheduler.request_finished(request, block_ids)
9884

@@ -103,8 +89,7 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
10389
assert self.connector_worker is not None
10490
self.connector_worker.register_kv_caches(kv_caches)
10591

106-
def start_load_kv(self, forward_context: "ForwardContext",
107-
**kwargs) -> None:
92+
def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
10893
assert self.connector_worker is not None
10994
self.connector_worker.start_load_kv(self._get_connector_metadata())
11095

@@ -113,8 +98,9 @@ def wait_for_layer_load(self, layer_name: str) -> None:
11398
return
11499
self.connector_worker.wait_for_layer_load()
115100

116-
def save_kv_layer(self, layer_name: str, kv_layer: torch.Tensor,
117-
attn_metadata: "AttentionMetadata", **kwargs) -> None:
101+
def save_kv_layer(
102+
self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", **kwargs
103+
) -> None:
118104
if not self.use_layerwise:
119105
return
120106

@@ -133,17 +119,16 @@ def wait_for_save(self):
133119

134120
self.connector_worker.wait_for_save(self._get_connector_metadata())
135121

136-
def get_finished(self,
137-
finished_req_ids: set[str]) -> tuple[set[str], set[str]]:
122+
def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str], set[str]]:
138123
"""Get the finished recving and sending requests."""
139124
assert self.connector_worker is not None
140125
done_sending, done_recving = self.connector_worker.get_finished(
141-
finished_req_ids, self._get_connector_metadata())
126+
finished_req_ids, self._get_connector_metadata()
127+
)
142128
return done_sending, done_recving
143129

144130

145131
class LookupKeyServer:
146-
147132
def __init__(
148133
self,
149134
pool_worker: KVPoolWorker,
@@ -171,8 +156,7 @@ def process_request():
171156
token_len = int.from_bytes(all_frames[0], byteorder="big")
172157
hash_frames = all_frames[1:]
173158
hashes_str = self.decoder.decode(hash_frames)
174-
result = self.pool_worker.lookup_scheduler(
175-
token_len, hashes_str, self.use_layerwise)
159+
result = self.pool_worker.lookup_scheduler(token_len, hashes_str, self.use_layerwise)
176160
response = result.to_bytes(4, "big")
177161
self.socket.send(response)
178162

vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/backend/backend.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,13 +4,15 @@
44

55

66
class Backend(ABC):
7-
7+
@abstractmethod
88
def __init__(self, parallel_config: ParallelConfig):
99
pass
1010

11+
@abstractmethod
1112
def set_device(self):
1213
pass
1314

15+
@abstractmethod
1416
def register_buffer(self, ptrs: list[int], lengths: list[int]):
1517
pass
1618

@@ -19,11 +21,9 @@ def exists(self, keys: list[str]) -> list[int]:
1921
pass
2022

2123
@abstractmethod
22-
def put(self, keys: list[str], addrs: list[list[int]],
23-
sizes: list[list[int]]):
24+
def put(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
2425
pass
2526

2627
@abstractmethod
27-
def get(self, keys: list[str], addrs: list[list[int]],
28-
sizes: list[list[int]]):
28+
def get(self, keys: list[str], addrs: list[list[int]], sizes: list[list[int]]):
2929
pass

vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/backend/memcache_backend.py

Lines changed: 11 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,7 @@
55
from vllm.config import ParallelConfig
66
from vllm.logger import logger
77

8-
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import \
9-
Backend
8+
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
109
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
1110

1211

@@ -18,29 +17,24 @@ class MmcDirect(Enum):
1817

1918

2019
class MemcacheBackend(Backend):
21-
2220
def __init__(self, parallel_config: ParallelConfig):
2321
try:
2422
from memcache_hybrid import DistributedObjectStore # type: ignore
2523
except ImportError as e:
2624
raise ImportError(
2725
"Please install memcache by following the instructions at "
2826
"https://gitee.com/ascend/memfabric_hybrid " # noqa: E501
29-
"to run vLLM with MemcacheConnector.") from e
27+
"to run vLLM with MemcacheConnector."
28+
) from e
3029
try:
3130
soc_version = get_ascend_device_type()
3231
if soc_version in {AscendDeviceType.A2}:
3332
import torch
3433
from vllm.distributed import get_world_group
34+
3535
tmp_tensor = torch.zeros(1, device="npu")
36-
output_tensor_list = [
37-
torch.empty_like(tmp_tensor)
38-
for _ in range(torch.distributed.get_world_size())
39-
]
40-
torch.distributed.all_gather(
41-
output_tensor_list,
42-
tmp_tensor,
43-
group=get_world_group().device_group)
36+
output_tensor_list = [torch.empty_like(tmp_tensor) for _ in range(torch.distributed.get_world_size())]
37+
torch.distributed.all_gather(output_tensor_list, tmp_tensor, group=get_world_group().device_group)
4438
self.rank = parallel_config.rank
4539
self.store = DistributedObjectStore()
4640
res = self.store.init(self.rank)
@@ -54,8 +48,7 @@ def __init__(self, parallel_config: ParallelConfig):
5448
logger.error("Configuration loading failed: %s", e)
5549
raise
5650
except Exception as exc:
57-
logger.error(
58-
"An error occurred while loading the configuration: %s", exc)
51+
logger.error("An error occurred while loading the configuration: %s", exc)
5952
raise
6053

6154
def set_device(self):
@@ -73,22 +66,18 @@ def register_buffer(self, ptrs: list[int], sizes: list[int]):
7366
def exists(self, keys: list[str]) -> list[int]:
7467
return self.store.batch_is_exist(keys)
7568

76-
def get(self, key: list[str], addr: list[list[int]],
77-
size: list[list[int]]):
69+
def get(self, key: list[str], addr: list[list[int]], size: list[list[int]]):
7870
try:
79-
res = self.store.batch_get_into_layers(key, addr, size,
80-
MmcDirect.COPY_G2L.value)
71+
res = self.store.batch_get_into_layers(key, addr, size, MmcDirect.COPY_G2L.value)
8172
for value in res:
8273
if value != 0:
8374
logger.error(f"Failed to get key {key},res:{res}")
8475
except Exception as e:
8576
logger.error(f"Failed to get key {key}. {e}")
8677

87-
def put(self, key: list[str], addr: list[list[int]],
88-
size: list[list[int]]):
78+
def put(self, key: list[str], addr: list[list[int]], size: list[list[int]]):
8979
try:
90-
res = self.store.batch_put_from_layers(key, addr, size,
91-
MmcDirect.COPY_L2G.value)
80+
res = self.store.batch_put_from_layers(key, addr, size, MmcDirect.COPY_L2G.value)
9281
for value in res:
9382
if value != 0:
9483
logger.error(f"Failed to get key {key},res:{res}")

0 commit comments

Comments
 (0)