Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions simdistserve/base/request.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
E_DO_PREFILL = "do_prefill"
E_WAIT_DECODE = "wait_decode"
E_DO_DECODE = "do_decode"
E_WAIT_KVCACHE_MIGRATION = "wait_kvcache_migration"
E_DO_KVCACHE_MIGRATION = "do_kvcache_migration"
E_FINISH_PREFILL = "finish_prefill"
E_FINISH_DECODE = "finish_decode"
E_EXIT_SYSTEM = "exit_system"
Expand Down Expand Up @@ -61,6 +63,11 @@ def __init__(
# set this value if a request belongs to a particular chunk
# The last worker in the pipeline unset this value at a chunk's end.
self.chunk_id = None
# after the request is finished prefill, `kvcache_generated` should be set to `True`.
self.kvcache_generated = False
self.prefill_is_done = False
self.kvcache_is_transferred = False
self.prefill_worker = None
Comment thread
Toseic marked this conversation as resolved.
Outdated

@property
def current_context_len(self):
Expand Down Expand Up @@ -88,6 +95,12 @@ def wait_decode(self, wid=None):

def do_decode(self, wid=None):
self._log_event(E_DO_DECODE, wid=wid)

def wait_kvcache_migration(self, wid=None):
self._log_event(E_WAIT_KVCACHE_MIGRATION, wid=wid)

def do_kvcache_migration(self, wid=None):
self._log_event(E_DO_KVCACHE_MIGRATION, wid=wid)

def _reset_chunked_prefill_metadata(self):
"""Reset the metadata of chunked prefill."""
Expand All @@ -111,6 +124,7 @@ def finish_prefill(self, is_finished_one_round=False, wid=None, next_wid=None):
# Reset counter to 0
# TODO: Should we do self.counter += 1?
self.counter = 0
self.prefill_is_done = True
# Hack to ensure "wait_decode" appears at least once.
self.wait_decode(wid=next_wid)
if not self.should_finish():
Expand Down
205 changes: 182 additions & 23 deletions simdistserve/base/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,8 @@ def __init__(
enable_chunked_prefill=False,
prefill_max_tokens=10 ** 7,
decode_max_tokens=10 ** 7,
free_mem_slots_num=69230, # TODO: (Refactor) This is a magic number. Should be a configuration.
per_token_kvcache_transfertime=0.01, # TODO: (Refactor) This is a magic number. Should be a configuration.
Comment thread
Toseic marked this conversation as resolved.
Outdated
decode_back_pressure: float = 0.9,
engine_type: Literal["distserve", "vllm"] = "distserve",
):
Expand Down Expand Up @@ -94,11 +96,19 @@ def __init__(

self.prefill_queue: 'deque[Request]' = deque()
self.decode_queue: 'deque[Request]' = deque()
# Transfer kv-cache to other workers, every request will be push to this queue after prefill,
# and waitng for some worker to receive it and decode.
self.migrate_queue: 'deque[Request]' = deque()
self._prefill_ips: int = 0 # Elements in progress for prefill
self._decode_ips: int = 0 # Elements in progress for decode
self._wakeup_event = env.event()
self.log: 'list[tuple[float, str, int, int, int, list[int], list[int]]]' = []

self.free_mem_slots_num = free_mem_slots_num # free slots for kv-cache
self.max_mem_slot_num = free_mem_slots_num
self.mem_slot_lower_bound = 0.1 * free_mem_slots_num # 10% of free slots
self.per_token_kvcache_transfertime = per_token_kvcache_transfertime

# Simulate scheduler delay in terms of number of decode rounds.
self._prefill_sched_delay: int = 0
self.engine_type = engine_type
Expand All @@ -108,10 +118,6 @@ def __init__(
def is_first_in_pipeline(self):
return self.pipe_rank == 0

@property
def has_back_pressure(self) -> bool:
threshold = int(self.decode_max_batch_size * self.decode_back_pressure)
return sum(r.current_context_len for r in self.decode_queue) > threshold

def __repr__(self):
return f"Worker {self.wid}"
Expand All @@ -129,11 +135,15 @@ def _log_event(self, event, num_tokens: int = 0, prefill_bs=0, decode_bs=0,

def run(self):
while True:
self.check_migrate_queue()
if not (self.prefill_queue or self.decode_queue):
yield self._wakeup_event

if self.prefill_queue and not self.has_back_pressure:
yield from self.do_prefill()
if self.prefill_queue :
if self.mem_is_enough():
yield from self.do_prefill()
else:
yield self.env.timeout(0.1) # avoid dead lock
else:
yield from self.do_decode()

Expand Down Expand Up @@ -182,28 +192,74 @@ def forward_decode(self, items: Union['Request', Iterable['Request']], to_schedu
return

def _enter_decodes(self, remaining_tok_in_batch: int) -> 'List[Request]':
# decode_max_tokens

# Acceptable decode requests is capped by the remaining allowed tokens in this batch.
# TODO: Hack: Must revert this to use the max token given
# watermark = 0.9
# decode_max_tokens = self.decode_max_tokens * watermark
decode_max_tokens = 50000
_decode_len = min(remaining_tok_in_batch, len(self.decode_queue))
decode_reqs = []
for i in range(_decode_len):
req = self.decode_queue[0]
if (req.current_context_len + 1) > decode_max_tokens:
break
decode_max_tokens -= (req.current_context_len + 1)
decode_reqs.append(self.decode_queue.popleft())
decode_reqs: 'List[Request]' = []

# if memory is not enough, only schedule the requests that have been decoded before
# because their kv-cache is already alloced, new requests' kv-cache cost too much free memory
decode_all_kinds_requests = self.mem_is_enough()

# request is given up if
# 1. request needs kv-cache migrate but memory is less than the mem_slot_lower_bound
# 2. available memory is not enough
requests_give_up = deque()

if self.mem_is_enough() or decode_all_kinds_requests:
available_slots = self.free_mem_slots_num // 2 # avoid free_mem_slots_num is used in one batch
else:
available_slots = self.free_mem_slots_num // 128 # batch size control

for _ in range(_decode_len):
left_req = self.decode_queue[0]
self.decode_queue.popleft()
if not left_req.kvcache_is_transferred: # requests just finished prefill
if not decode_all_kinds_requests:
requests_give_up.append(left_req) # needs kv-cache migrate, give up
continue
elif available_slots > left_req.current_context_len : # memory is enough, migrate kv-cache
available_slots -= left_req.current_context_len
decode_reqs.append(left_req)
else: # memory is not enough, give up the request
requests_give_up.append(left_req)
else: # common requests
if available_slots <= 1:
requests_give_up.append(left_req)
continue
decode_reqs.append(left_req)
available_slots -= 1



# put the requests kicked back to the queue from front
while len(requests_give_up) > 0:
self.decode_queue.appendleft(requests_give_up.pop())
assert len(self.decode_queue) + len(decode_reqs) == _decode_len

# kv-cache transfer
migrate_time = 0
for r in list(decode_reqs):
if not r.kvcache_is_transferred:
if r.prefill_worker.wid != self.wid: # if the kv-cache is not in the worker, then migrate
migrate_time += self.per_token_kvcache_transfertime * r.current_context_len
self.migrate_alloc_kvcache([r,])
r.prefill_worker.wakeup()

yield self.env.timeout(migrate_time)
self.decode_alloc_kvcache(decode_reqs)
Comment thread
Toseic marked this conversation as resolved.
Outdated

for r in decode_reqs:
r.do_decode(wid=self.wid)
return decode_reqs

def _enter_prefill(self) -> 'List[Request]':
result: 'List[Request]' = []

# check if free_slot_num touches the lower bound
if not self.mem_is_enough():
return result

available_slots = self.free_mem_slots_num

# Limit the maximum prefill requests to handle.
max_request_size = min(self.prefill_max_batch_size, len(self.prefill_queue))

Expand All @@ -216,7 +272,10 @@ def _enter_prefill(self) -> 'List[Request]':
candidate: 'Request' = self.prefill_queue[0]
if candidate.chunk_id != chunk_id:
break
if available_slots < candidate.current_prefill_lens: # TODO: not sure if this is correct
break
result.append(self.prefill_queue.popleft())
available_slots -= candidate.current_prefill_lens
pass

else:
Expand Down Expand Up @@ -248,6 +307,9 @@ def _enter_prefill(self) -> 'List[Request]':
break
pass

if available_slots < sched_size:
break
available_slots -= sched_size
# Candidate is picked. Now fill in the chunked-prefill information.
candidate.current_prefill_lens = sched_size
candidate.remain_prefill_lens -= sched_size
Expand All @@ -259,12 +321,22 @@ def _enter_prefill(self) -> 'List[Request]':
pass
for i in result:
i.do_prefill(wid=self.wid)

if result:
self.prefill_alloc_kvcache(result)

return result

def _exit_prefill(self, prefill_items: List['Request']):
# if a request finished prefill, it should be migrated to other workers
requests_need_migrate = []

for item in prefill_items:
next_wid = self.next_worker.wid if self.next_worker else None
item.finish_prefill(is_finished_one_round=self.is_last_in_pipeline, wid=self.wid, next_wid=next_wid)
if item.prefill_is_done:
item.prefill_worker = self
requests_need_migrate.append(item)
if not self.is_last_in_pipeline or (item.remain_prefill_lens > 0):
# Finish one chunk of prefill. Now forward to the next worker
# (or head of worker) to do the rest of the parts.
Expand All @@ -276,20 +348,28 @@ def _exit_prefill(self, prefill_items: List['Request']):
# ... just a sanity check to avoid any infinite loop.
continue
self.forward_decode(item, to_scheduler=(not self.should_request_stay))
self.migrate_kvcache(requests_need_migrate)
return

def _exit_decode(self, decode_reqs):
def _exit_decode(self, decode_reqs: 'List[Request]'):
finished_requests = [] # if the request is finished, its kv-cache should be freed

if not decode_reqs:
return
next_wid = self.next_worker.wid if self.next_worker else None
for r in decode_reqs:
r.finish_decode(is_finished_one_round=self.is_last_in_pipeline, next_wid=next_wid)
if r._terminated:
finished_requests.append(r)
next_decode_batch = tuple(r for r in decode_reqs if not r.should_finish())
self.decode_free_kvcache(finished_requests)
self.forward_decode(next_decode_batch)
return

def do_prefill(self):
prefill_items: 'List[Request]' = self._enter_prefill()
if not prefill_items:
return
if self.enable_chunked_prefill:
remaining_tok_in_batch = self.prefill_max_tokens - sum(x.current_prefill_lens for x in prefill_items)
decode_reqs = self._enter_decodes(remaining_tok_in_batch)
Expand Down Expand Up @@ -332,8 +412,11 @@ def do_prefill(self):
return

def do_decode(self):
decode_reqs = self._enter_decodes(self.decode_max_tokens)
batch_size = len(decode_reqs)
decode_reqs = yield self.env.process(self._enter_decodes(self.decode_max_tokens))
batch_size = len(list(decode_reqs))
if batch_size == 0:
return

self._log_event(
"do_decode", num_tokens=batch_size, decode_bs=batch_size,
decode_len_list=[x.current_context_len for x in decode_reqs],
Expand All @@ -351,3 +434,79 @@ def do_decode(self):
return

pass

def check_migrate_queue(self):
# check if requests' kv-cache is received by other workers, if yes then free the slots
_migrate_queue_len = len(self.migrate_queue)
if len(self.migrate_queue) == 0:
return
migrated_requests = []
stay_requests = []
for i in range(len(self.migrate_queue)):
r = self.migrate_queue.popleft()
if r.kvcache_is_transferred:
migrated_requests.append(r)
else:
stay_requests.append(r)
for r in stay_requests:
self.migrate_queue.append(r)
assert len(migrated_requests) + len(stay_requests) == _migrate_queue_len
if migrated_requests:
self.prefill_free_kvcache(migrated_requests)

def mem_is_enough(self, requests: List['Request']=[]) -> bool:
assert self.free_mem_slots_num <= self.max_mem_slot_num
assert self.free_mem_slots_num >= 0

if requests and all(r.counter > 0 for r in requests):
return True

return self.free_mem_slots_num >= self.mem_slot_lower_bound


def migrate_kvcache(self, requests: 'list[Request]') -> None:
# called by prefill worker, push the requests to the migrate queue,
# and waiting for the decode worker to receive it
# TODO: if the request's output_len == 1, decode won't happen, the kv-cache should be freed
for i in requests:
self.migrate_queue.append(i)
i.wait_kvcache_migration(wid=self.wid)

def prefill_alloc_kvcache(self, requests: 'list[Request]') -> bool:
# allocate slots for kv-cache
if not self.mem_is_enough():
return False
for i in requests:
i.kvcache_generated = True
self.free_mem_slots_num -= sum([request.current_prefill_lens for request in requests])
self._log_event('prefill_alloc_kvcache')
return True


def prefill_free_kvcache(self, requests: 'list[Request]') -> None:
# called by prefill worker, free the slots immediately after kv-cache migration finished
self.free_mem_slots_num += sum([request.prefill_lens for request in requests])
self._log_event('prefill_free_kvcache')
return

def migrate_alloc_kvcache(self, requests: 'list[Request]') -> bool:
# called by the decode worker, prepare for the kv-cache migration
assert not any([request.kvcache_is_transferred for request in requests]), "The kv-cache is already transferred."
self.free_mem_slots_num -= sum([(request.prefill_lens) for request in requests])
for r in requests:
r.kvcache_is_transferred = True

def decode_alloc_kvcache(self, requests: 'list[Request]') -> bool:
# for each request, decode once cost one slot
self.free_mem_slots_num -= len(requests)
self._log_event('decode_alloc_kvcache')
return True

def decode_free_kvcache(self, requests: 'list[Request]') -> None:
# free the slots immediately after decoding finished
self.free_mem_slots_num += sum([(request.current_context_len) for request in requests])
self._log_event('decode_free_kvcache')
return

def __del__(self):
assert self.free_mem_slots_num == self.max_mem_slot_num, f"worker:{self.wid} free_mem_slots_num: {self.free_mem_slots_num}, max_mem_slot_num: {self.max_mem_slot_num}"