Ouzhang's picture
Add files using upload-large-folder tool
d91766b verified
Raw
History Blame Contribute Delete
4.03 kB
"""Req base class and registry."""
from __future__ import annotations
from copy import copy
from itertools import count
from typing import Callable
from diffulex.config import Config
from diffulex.sampling_params import SamplingParams
from diffulex.engine.strategy_registry import DiffulexStrategyRegistry
from diffulex.engine.status import DllmReqStatus
from diffulex.mixin.request_state import ReqStateMixin
class DllmReq(ReqStateMixin):
"""Minimal base class that tracks prompt tokens and cache bookkeeping."""
page_size = 32
counter = count()
def __init__(self, token_ids: list[int], sampling_params: SamplingParams = SamplingParams()):
self.req_id = next(DllmReq.counter)
self.status = DllmReqStatus.WAITING
self.dp_rank = 0
self._dp_owner_assigned = False
self.token_ids = copy(token_ids)
self.last_token = token_ids[-1]
self.num_prompt_tokens = len(token_ids)
self.num_cached_tokens = 0
self.page_table: list[int] = []
self.page_cache_missed: list[bool] = []
self.temperature = sampling_params.temperature
self.max_tokens = sampling_params.max_tokens
self.max_nfe = sampling_params.max_nfe
self.max_repetition_run = sampling_params.max_repetition_run
self.ignore_eos = sampling_params.ignore_eos
self.new_tokens = 0
self.nfe = 0
self.meet_eos = False
self.is_multi_block = False
self._execution_prepared = False
def __len__(self) -> int:
return self.num_tokens
def __getitem__(self, key) -> int:
return self.token_ids[key]
@property
def num_tokens(self) -> int:
return len(self.token_ids)
@property
def is_finished(self) -> bool:
return self.status == DllmReqStatus.FINISHED
@property
def prompt_token_ids(self) -> list[int]:
return self.token_ids[: self.num_prompt_tokens]
@property
def num_pages(self) -> int:
if self.is_multi_block:
# return (self.running_len + self.page_size - 1) // self.page_size
return (self.to_cache_len + self.page_size - 1) // self.page_size
else:
return (self.num_tokens + self.page_size - 1) // self.page_size
@property
def last_page_num_tokens(self) -> int:
return self.num_tokens - (self.num_pages - 1) * self.page_size
def reset_new_tokens(self):
self.new_tokens = 0
@property
def is_execution_prepared(self) -> bool:
return bool(self._execution_prepared)
def mark_execution_prepared(self) -> None:
self._execution_prepared = True
def clear_execution_prepared(self) -> None:
self._execution_prepared = False
def assign_dp_rank(self, dp_rank: int) -> None:
self.dp_rank = dp_rank
self._dp_owner_assigned = True
def page(self, index: int) -> list[int]:
assert 0 <= index < self.num_pages
return self.token_ids[index * self.page_size : (index + 1) * self.page_size]
ReqFactory = Callable[[list[int], SamplingParams, Config], DllmReq]
class AutoReq(DiffulexStrategyRegistry):
"""Registry-driven factory for req implementations."""
@classmethod
def create(
cls,
config: Config,
token_ids: list[int],
sampling_params: SamplingParams = SamplingParams(),
) -> DllmReq:
cls._MODULE_MAPPING: dict[str, ReqFactory]
candidates: list[str] = []
if config.decoding_strategy:
candidates.append(config.decoding_strategy)
candidates.append(cls._DEFAULT_KEY)
for key in candidates:
factory = cls._MODULE_MAPPING.get(key)
if factory is not None:
return factory(token_ids, sampling_params, config)
available = ", ".join(cls.available_modules()) or "<none>"
raise ValueError(
"No req registered for decoding_strategy="
f"'{config.decoding_strategy}'. Available reqs: {available}."
)