[2/N] Added the core structure of elastic EP and the eplb algorithm with faulty rank (#10606)
Co-authored-by: Xun Sun <UNIDY2002@outlook.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Xun Sun
Shangming Cai
parent
e028af6998
commit
904655c5fd
@@ -0,0 +1,74 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ServerArgs
|
||||
from sglang.srt.utils import is_cpu, is_cuda
|
||||
|
||||
|
||||
@dataclass
|
||||
class ElasticEPState:
|
||||
active_ranks: Optional[torch.Tensor]
|
||||
last_active_ranks: Optional[torch.Tensor]
|
||||
active_ranks_cpu: Optional[torch.Tensor]
|
||||
|
||||
def is_active_equal_last(self) -> bool:
|
||||
return torch.equal(self.active_ranks, self.last_active_ranks)
|
||||
|
||||
def sync_active_to_cpu(self):
|
||||
if self.active_ranks is not None:
|
||||
self.active_ranks_cpu = self.active_ranks.detach().cpu().clone()
|
||||
|
||||
def snapshot_active_to_last(self):
|
||||
if self.active_ranks is not None:
|
||||
self.last_active_ranks = self.active_ranks.clone()
|
||||
|
||||
|
||||
class ElasticEPStateManager:
|
||||
_instance: Optional[ElasticEPState] = None
|
||||
|
||||
@classmethod
|
||||
def instance(cls) -> ElasticEPState:
|
||||
return cls._instance
|
||||
|
||||
@classmethod
|
||||
def init(cls, server_args: ServerArgs):
|
||||
if cls._instance is not None:
|
||||
return cls._instance
|
||||
|
||||
if server_args.elastic_ep_backend is not None:
|
||||
cls._instance = cls._build_state(ep_size=None, device=None)
|
||||
return cls._instance
|
||||
|
||||
@staticmethod
|
||||
def _select_device() -> torch.device:
|
||||
if is_cuda():
|
||||
return torch.device("cuda")
|
||||
elif is_cpu():
|
||||
return torch.device("cpu")
|
||||
else:
|
||||
raise NotImplementedError("Only CUDA and CPU support elastic ep now.")
|
||||
|
||||
@classmethod
|
||||
def _build_state(
|
||||
cls, *, ep_size: Optional[int] = None, device: Optional[torch.device] = None
|
||||
) -> ElasticEPState:
|
||||
|
||||
active = cls.healthy_rank_state(ep_size=ep_size, device=device)
|
||||
return ElasticEPState(
|
||||
active_ranks=active,
|
||||
last_active_ranks=active.clone(),
|
||||
active_ranks_cpu=active.detach().cpu().clone(),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def healthy_rank_state(
|
||||
cls, *, ep_size: Optional[int] = None, device: Optional[torch.device] = None
|
||||
) -> torch.Tensor:
|
||||
size = ep_size if ep_size is not None else torch.distributed.get_world_size()
|
||||
dev = device if device is not None else cls._select_device()
|
||||
|
||||
return torch.ones(size, dtype=torch.int32, device=dev)
|
||||
Reference in New Issue
Block a user