Abstraction for spec worker and code cleanup (#11643)
This commit is contained in:
@@ -15,6 +15,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
@@ -54,7 +55,140 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TpModelWorker:
|
||||
class BaseTpWorker(ABC):
|
||||
@abstractmethod
|
||||
def forward_batch_generation(self, forward_batch: ForwardBatch):
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def model_runner(self) -> ModelRunner:
|
||||
pass
|
||||
|
||||
@property
|
||||
def sliding_window_size(self) -> Optional[int]:
|
||||
return self.model_runner.sliding_window_size
|
||||
|
||||
@property
|
||||
def is_hybrid(self) -> bool:
|
||||
return self.model_runner.is_hybrid is not None
|
||||
|
||||
def get_tokens_per_layer_info(self):
|
||||
return (
|
||||
self.model_runner.full_max_total_num_tokens,
|
||||
self.model_runner.swa_max_total_num_tokens,
|
||||
)
|
||||
|
||||
def get_pad_input_ids_func(self):
|
||||
return getattr(self.model_runner.model, "pad_input_ids", None)
|
||||
|
||||
def get_tp_group(self):
|
||||
return self.model_runner.tp_group
|
||||
|
||||
def get_attention_tp_group(self):
|
||||
return self.model_runner.attention_tp_group
|
||||
|
||||
def get_attention_tp_cpu_group(self):
|
||||
return getattr(self.model_runner.attention_tp_group, "cpu_group", None)
|
||||
|
||||
def get_memory_pool(self):
|
||||
return (
|
||||
self.model_runner.req_to_token_pool,
|
||||
self.model_runner.token_to_kv_pool_allocator,
|
||||
)
|
||||
|
||||
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||
success, message = self.model_runner.update_weights_from_disk(
|
||||
recv_req.model_path, recv_req.load_format
|
||||
)
|
||||
return success, message
|
||||
|
||||
def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput):
|
||||
success, message = self.model_runner.init_weights_update_group(
|
||||
recv_req.master_address,
|
||||
recv_req.master_port,
|
||||
recv_req.rank_offset,
|
||||
recv_req.world_size,
|
||||
recv_req.group_name,
|
||||
recv_req.backend,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def destroy_weights_update_group(self, recv_req: DestroyWeightsUpdateGroupReqInput):
|
||||
success, message = self.model_runner.destroy_weights_update_group(
|
||||
recv_req.group_name,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def init_weights_send_group_for_remote_instance(
|
||||
self, recv_req: InitWeightsSendGroupForRemoteInstanceReqInput
|
||||
):
|
||||
success, message = (
|
||||
self.model_runner.init_weights_send_group_for_remote_instance(
|
||||
recv_req.master_address,
|
||||
recv_req.ports,
|
||||
recv_req.group_rank,
|
||||
recv_req.world_size,
|
||||
recv_req.group_name,
|
||||
recv_req.backend,
|
||||
)
|
||||
)
|
||||
return success, message
|
||||
|
||||
def send_weights_to_remote_instance(
|
||||
self, recv_req: SendWeightsToRemoteInstanceReqInput
|
||||
):
|
||||
success, message = self.model_runner.send_weights_to_remote_instance(
|
||||
recv_req.master_address,
|
||||
recv_req.ports,
|
||||
recv_req.group_name,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def update_weights_from_distributed(
|
||||
self, recv_req: UpdateWeightsFromDistributedReqInput
|
||||
):
|
||||
success, message = self.model_runner.update_weights_from_distributed(
|
||||
recv_req.names, recv_req.dtypes, recv_req.shapes, recv_req.group_name
|
||||
)
|
||||
return success, message
|
||||
|
||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||
|
||||
monkey_patch_torch_reductions()
|
||||
success, message = self.model_runner.update_weights_from_tensor(
|
||||
named_tensors=MultiprocessingSerializer.deserialize(
|
||||
recv_req.serialized_named_tensors[self.tp_rank]
|
||||
),
|
||||
load_format=recv_req.load_format,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
|
||||
parameter = self.model_runner.get_weights_by_name(
|
||||
recv_req.name, recv_req.truncate_size
|
||||
)
|
||||
return parameter
|
||||
|
||||
def load_lora_adapter(self, recv_req: LoadLoRAAdapterReqInput):
|
||||
result = self.model_runner.load_lora_adapter(recv_req.to_ref())
|
||||
return result
|
||||
|
||||
def unload_lora_adapter(self, recv_req: UnloadLoRAAdapterReqInput):
|
||||
result = self.model_runner.unload_lora_adapter(recv_req.to_ref())
|
||||
return result
|
||||
|
||||
def can_run_lora_batch(self, lora_ids: list[str]) -> bool:
|
||||
return self.model_runner.lora_manager.validate_lora_batch(lora_ids)
|
||||
|
||||
def forward_batch_embedding(self, model_worker_batch: ModelWorkerBatch):
|
||||
forward_batch = ForwardBatch.init_new(model_worker_batch, self.model_runner)
|
||||
logits_output, _ = self.model_runner.forward(forward_batch)
|
||||
embeddings = logits_output.embeddings
|
||||
return embeddings
|
||||
|
||||
|
||||
class TpModelWorker(BaseTpWorker):
|
||||
"""A tensor parallel model worker."""
|
||||
|
||||
def __init__(
|
||||
@@ -92,7 +226,7 @@ class TpModelWorker:
|
||||
is_draft_model=is_draft_worker,
|
||||
)
|
||||
|
||||
self.model_runner = ModelRunner(
|
||||
self._model_runner = ModelRunner(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
gpu_id=gpu_id,
|
||||
@@ -171,6 +305,10 @@ class TpModelWorker:
|
||||
self.enable_overlap = not server_args.disable_overlap_schedule
|
||||
self.hicache_layer_transfer_counter = None
|
||||
|
||||
@property
|
||||
def model_runner(self) -> ModelRunner:
|
||||
return self._model_runner
|
||||
|
||||
def register_hicache_layer_transfer_counter(self, counter: LayerDoneCounter):
|
||||
self.hicache_layer_transfer_counter = counter
|
||||
|
||||
@@ -193,38 +331,6 @@ class TpModelWorker:
|
||||
self.model_runner.token_to_kv_pool.size,
|
||||
)
|
||||
|
||||
@property
|
||||
def sliding_window_size(self) -> Optional[int]:
|
||||
return self.model_runner.sliding_window_size
|
||||
|
||||
@property
|
||||
def is_hybrid(self) -> bool:
|
||||
return self.model_runner.is_hybrid is not None
|
||||
|
||||
def get_tokens_per_layer_info(self):
|
||||
return (
|
||||
self.model_runner.full_max_total_num_tokens,
|
||||
self.model_runner.swa_max_total_num_tokens,
|
||||
)
|
||||
|
||||
def get_pad_input_ids_func(self):
|
||||
return getattr(self.model_runner.model, "pad_input_ids", None)
|
||||
|
||||
def get_tp_group(self):
|
||||
return self.model_runner.tp_group
|
||||
|
||||
def get_attention_tp_group(self):
|
||||
return self.model_runner.attention_tp_group
|
||||
|
||||
def get_attention_tp_cpu_group(self):
|
||||
return getattr(self.model_runner.attention_tp_group, "cpu_group", None)
|
||||
|
||||
def get_memory_pool(self):
|
||||
return (
|
||||
self.model_runner.req_to_token_pool,
|
||||
self.model_runner.token_to_kv_pool_allocator,
|
||||
)
|
||||
|
||||
def forward_batch_generation(
|
||||
self,
|
||||
model_worker_batch: ModelWorkerBatch,
|
||||
@@ -313,93 +419,3 @@ class TpModelWorker:
|
||||
pp_hidden_states_proxy_tensors=pp_proxy_tensors,
|
||||
can_run_cuda_graph=can_run_cuda_graph,
|
||||
)
|
||||
|
||||
def forward_batch_embedding(self, model_worker_batch: ModelWorkerBatch):
|
||||
forward_batch = ForwardBatch.init_new(model_worker_batch, self.model_runner)
|
||||
logits_output, _ = self.model_runner.forward(forward_batch)
|
||||
embeddings = logits_output.embeddings
|
||||
return embeddings
|
||||
|
||||
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||
success, message = self.model_runner.update_weights_from_disk(
|
||||
recv_req.model_path, recv_req.load_format
|
||||
)
|
||||
return success, message
|
||||
|
||||
def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput):
|
||||
success, message = self.model_runner.init_weights_update_group(
|
||||
recv_req.master_address,
|
||||
recv_req.master_port,
|
||||
recv_req.rank_offset,
|
||||
recv_req.world_size,
|
||||
recv_req.group_name,
|
||||
recv_req.backend,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def destroy_weights_update_group(self, recv_req: DestroyWeightsUpdateGroupReqInput):
|
||||
success, message = self.model_runner.destroy_weights_update_group(
|
||||
recv_req.group_name,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def init_weights_send_group_for_remote_instance(
|
||||
self, recv_req: InitWeightsSendGroupForRemoteInstanceReqInput
|
||||
):
|
||||
success, message = (
|
||||
self.model_runner.init_weights_send_group_for_remote_instance(
|
||||
recv_req.master_address,
|
||||
recv_req.ports,
|
||||
recv_req.group_rank,
|
||||
recv_req.world_size,
|
||||
recv_req.group_name,
|
||||
recv_req.backend,
|
||||
)
|
||||
)
|
||||
return success, message
|
||||
|
||||
def send_weights_to_remote_instance(
|
||||
self, recv_req: SendWeightsToRemoteInstanceReqInput
|
||||
):
|
||||
success, message = self.model_runner.send_weights_to_remote_instance(
|
||||
recv_req.master_address,
|
||||
recv_req.ports,
|
||||
recv_req.group_name,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def update_weights_from_distributed(
|
||||
self, recv_req: UpdateWeightsFromDistributedReqInput
|
||||
):
|
||||
success, message = self.model_runner.update_weights_from_distributed(
|
||||
recv_req.names, recv_req.dtypes, recv_req.shapes, recv_req.group_name
|
||||
)
|
||||
return success, message
|
||||
|
||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||
|
||||
monkey_patch_torch_reductions()
|
||||
success, message = self.model_runner.update_weights_from_tensor(
|
||||
named_tensors=MultiprocessingSerializer.deserialize(
|
||||
recv_req.serialized_named_tensors[self.tp_rank]
|
||||
),
|
||||
load_format=recv_req.load_format,
|
||||
)
|
||||
return success, message
|
||||
|
||||
def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
|
||||
parameter = self.model_runner.get_weights_by_name(
|
||||
recv_req.name, recv_req.truncate_size
|
||||
)
|
||||
return parameter
|
||||
|
||||
def load_lora_adapter(self, recv_req: LoadLoRAAdapterReqInput):
|
||||
result = self.model_runner.load_lora_adapter(recv_req.to_ref())
|
||||
return result
|
||||
|
||||
def unload_lora_adapter(self, recv_req: UnloadLoRAAdapterReqInput):
|
||||
result = self.model_runner.unload_lora_adapter(recv_req.to_ref())
|
||||
return result
|
||||
|
||||
def can_run_lora_batch(self, lora_ids: list[str]) -> bool:
|
||||
return self.model_runner.lora_manager.validate_lora_batch(lora_ids)
|
||||
|
||||
Reference in New Issue
Block a user