mirror of
https://github.com/wassname/vllm.git
synced 2026-09-11 12:51:01 +08:00
[Core] Subclass ModelRunner to support cross-attention & encoder sequences (towards eventual encoder/decoder model support) (#4942)
Co-authored-by: Andrew Feldman <afeld2012@gmail.com> Co-authored-by: Nick Hill <nickhill@us.ibm.com>
This commit is contained in:
co-authored by
Andrew Feldman
Nick Hill
parent
660470e5a3
commit
fd95e026e0
+65
-36
@@ -53,27 +53,30 @@ def create_dummy_prompt_encoder_decoder(
|
||||
block_size = decoder_prompt_length
|
||||
|
||||
# Create dummy prompt sequence with tokens 0...block_size-1
|
||||
# and prompt "0 ... block_size".
|
||||
# and prompt "0 ... block_size". Note that the prompt string
|
||||
# doesn't actually match the tokens
|
||||
decoder_prompt_tokens = list(range(decoder_prompt_length))
|
||||
decoder_prompt_str = " ".join([str(t) for t in decoder_prompt_tokens])
|
||||
|
||||
decoder_prompt = Sequence(int(request_id),
|
||||
inputs={
|
||||
"prompt": decoder_prompt_str,
|
||||
"prompt_token_ids": decoder_prompt_tokens,
|
||||
"multi_modal_data": None,
|
||||
},
|
||||
block_size=block_size)
|
||||
|
||||
encoder_prompt_tokens = list(reversed(list(range(encoder_prompt_length))))
|
||||
encoder_prompt_str = " ".join([str(t) for t in encoder_prompt_tokens])
|
||||
|
||||
inputs = {
|
||||
"prompt": decoder_prompt_str,
|
||||
"prompt_token_ids": decoder_prompt_tokens,
|
||||
"encoder_prompt": encoder_prompt_str,
|
||||
"encoder_prompt_token_ids": encoder_prompt_tokens,
|
||||
"multi_modal_data": None,
|
||||
}
|
||||
|
||||
decoder_prompt = Sequence(int(request_id),
|
||||
inputs=inputs,
|
||||
block_size=block_size,
|
||||
from_decoder_prompt=True)
|
||||
|
||||
encoder_prompt = Sequence(int(request_id),
|
||||
inputs={
|
||||
"prompt": encoder_prompt_str,
|
||||
"prompt_token_ids": encoder_prompt_tokens,
|
||||
"multi_modal_data": None,
|
||||
},
|
||||
block_size=block_size)
|
||||
inputs=inputs,
|
||||
block_size=block_size,
|
||||
from_decoder_prompt=False)
|
||||
seq_group = SequenceGroup(request_id=request_id,
|
||||
seqs=[decoder_prompt],
|
||||
sampling_params=SamplingParams(
|
||||
@@ -139,17 +142,21 @@ def create_seq_group_encoder_decoder(
|
||||
|
||||
prompt_token_ids = [0] * seq_prompt_len
|
||||
|
||||
inputs = {
|
||||
"prompt": "",
|
||||
"prompt_token_ids": prompt_token_ids,
|
||||
"encoder_prompt": "",
|
||||
"encoder_prompt_token_ids": prompt_token_ids,
|
||||
"multi_modal_data": None,
|
||||
}
|
||||
|
||||
seqs = []
|
||||
for seq_id_offset, output_len in enumerate(seq_output_lens):
|
||||
seq = Sequence(
|
||||
seq_id=seq_id_start + seq_id_offset,
|
||||
inputs={
|
||||
"prompt": "",
|
||||
"prompt_token_ids": prompt_token_ids,
|
||||
"multi_modal_data": None,
|
||||
},
|
||||
block_size=16,
|
||||
)
|
||||
# Construct decoder input sequences
|
||||
seq = Sequence(seq_id=seq_id_start + seq_id_offset,
|
||||
inputs=inputs,
|
||||
block_size=16,
|
||||
from_decoder_prompt=True)
|
||||
|
||||
for i in range(output_len):
|
||||
seq.append_token_id(
|
||||
@@ -158,16 +165,11 @@ def create_seq_group_encoder_decoder(
|
||||
)
|
||||
seqs.append(seq)
|
||||
|
||||
# Encoder sequence
|
||||
encoder_seq = Sequence(
|
||||
seq_id=seq_id_start + len(seq_output_lens),
|
||||
inputs={
|
||||
"prompt": "",
|
||||
"prompt_token_ids": prompt_token_ids,
|
||||
"multi_modal_data": None,
|
||||
},
|
||||
block_size=16,
|
||||
)
|
||||
# Encoder input sequence
|
||||
encoder_seq = Sequence(seq_id=seq_id_start + len(seq_output_lens),
|
||||
inputs=inputs,
|
||||
block_size=16,
|
||||
from_decoder_prompt=False)
|
||||
|
||||
return SequenceGroup(request_id=request_id,
|
||||
seqs=seqs,
|
||||
@@ -177,4 +179,31 @@ def create_seq_group_encoder_decoder(
|
||||
|
||||
|
||||
def round_up_to_next_block(seq_len: int, block_size: int) -> int:
|
||||
return (seq_len + block_size - 1) // block_size
|
||||
return (seq_len + block_size - 1) // block_size
|
||||
|
||||
|
||||
# Helper functions for scheduler tests
|
||||
|
||||
|
||||
def get_sequence_groups(scheduler_output):
|
||||
return [s.seq_group for s in scheduler_output.scheduled_seq_groups]
|
||||
|
||||
|
||||
def append_new_token(out, token_id: int):
|
||||
seq_groups = get_sequence_groups(out)
|
||||
for seq_group in seq_groups:
|
||||
for seq in seq_group.get_seqs():
|
||||
seq.append_token_id(token_id, {token_id: Logprob(token_id)})
|
||||
|
||||
|
||||
def schedule_and_update_computed_tokens(scheduler):
|
||||
metas, out = scheduler.schedule()
|
||||
for s, meta in zip(out.scheduled_seq_groups, metas):
|
||||
s.seq_group.update_num_computed_tokens(meta.token_chunk_size)
|
||||
return metas, out
|
||||
|
||||
|
||||
def append_new_token_seq_group(token_chunk_size, seq_group, token_id: int):
|
||||
seq_group.update_num_computed_tokens(token_chunk_size)
|
||||
for seq in seq_group.get_seqs():
|
||||
seq.append_token_id(token_id, {token_id: Logprob(token_id)})
|
||||
|
||||
Reference in New Issue
Block a user