Introduce LLM class for offline inference (#115)

This commit is contained in:
Woosuk Kwon
2023-05-21 17:04:18 -07:00
committed by GitHub
parent f746ced08d
commit 655a5e48df
9 changed files with 222 additions and 81 deletions

View File

@@ -35,7 +35,7 @@ class RequestOutput:
prompt: str,
prompt_token_ids: List[int],
outputs: List[CompletionOutput],
done: bool = False,
done: bool,
) -> None:
self.request_id = request_id
self.prompt = prompt
@@ -43,8 +43,8 @@ class RequestOutput:
self.outputs = outputs
self.done = done
@staticmethod
def from_seq_group(seq_group: SequenceGroup) -> "RequestOutput":
@classmethod
def from_seq_group(cls, seq_group: SequenceGroup) -> "RequestOutput":
# Get the top-n sequences.
n = seq_group.sampling_params.n
seqs = seq_group.get_seqs()
@@ -70,8 +70,8 @@ class RequestOutput:
# Every sequence in the sequence group should have the same prompt.
prompt = top_n_seqs[0].prompt
prompt_token_ids = top_n_seqs[0].data.prompt_token_ids
return RequestOutput(seq_group.request_id, prompt, prompt_token_ids,
outputs, seq_group.is_finished())
return cls(seq_group.request_id, prompt, prompt_token_ids, outputs,
seq_group.is_finished())
def __repr__(self) -> str:
return (f"RequestOutput(request_id={self.request_id}, "