Mixed style of chunked prefill (#1013)
This commit is contained in:
@@ -88,11 +88,11 @@ class InputMetadata:
|
||||
self.image_sizes = [r.image_size for r in reqs]
|
||||
self.image_offsets = [
|
||||
(
|
||||
(r.image_offset - len(r.prefix_indices))
|
||||
(r.image_offset - batch.prefix_lens_cpu[i])
|
||||
if r.image_offset is not None
|
||||
else 0
|
||||
)
|
||||
for r in reqs
|
||||
for i, r in enumerate(reqs)
|
||||
]
|
||||
|
||||
def compute_positions(self, batch: ScheduleBatch):
|
||||
@@ -109,8 +109,8 @@ class InputMetadata:
|
||||
self.positions = torch.tensor(
|
||||
np.concatenate(
|
||||
[
|
||||
np.arange(len(req.prefix_indices), len(req.fill_ids))
|
||||
for req in batch.reqs
|
||||
np.arange(batch.prefix_lens_cpu[i], len(req.fill_ids))
|
||||
for i, req in enumerate(batch.reqs)
|
||||
],
|
||||
axis=0,
|
||||
),
|
||||
@@ -123,7 +123,7 @@ class InputMetadata:
|
||||
np.concatenate(
|
||||
[
|
||||
np.arange(
|
||||
len(req.prefix_indices) + position_ids_offsets_cpu[i],
|
||||
batch.prefix_lens_cpu[i] + position_ids_offsets_cpu[i],
|
||||
len(req.fill_ids) + position_ids_offsets_cpu[i],
|
||||
)
|
||||
for i, req in enumerate(batch.reqs)
|
||||
@@ -141,12 +141,13 @@ class InputMetadata:
|
||||
self.extend_seq_lens = self.extend_start_loc = self.extend_no_prefix = None
|
||||
else:
|
||||
extend_lens_cpu = [
|
||||
len(r.fill_ids) - len(r.prefix_indices) for r in batch.reqs
|
||||
len(r.fill_ids) - batch.prefix_lens_cpu[i]
|
||||
for i, r in enumerate(batch.reqs)
|
||||
]
|
||||
self.extend_seq_lens = torch.tensor(extend_lens_cpu, device="cuda")
|
||||
self.extend_start_loc = torch.zeros_like(self.seq_lens)
|
||||
self.extend_start_loc[1:] = torch.cumsum(self.extend_seq_lens[:-1], dim=0)
|
||||
self.extend_no_prefix = all(len(r.prefix_indices) == 0 for r in batch.reqs)
|
||||
self.extend_no_prefix = all(l == 0 for l in batch.prefix_lens_cpu)
|
||||
|
||||
@classmethod
|
||||
def from_schedule_batch(
|
||||
@@ -180,14 +181,8 @@ class InputMetadata:
|
||||
if forward_mode != ForwardMode.DECODE:
|
||||
ret.init_multimuldal_info(batch)
|
||||
|
||||
prefix_lens = None
|
||||
if forward_mode != ForwardMode.DECODE:
|
||||
prefix_lens = torch.tensor(
|
||||
[len(r.prefix_indices) for r in batch.reqs], device="cuda"
|
||||
)
|
||||
|
||||
if model_runner.server_args.disable_flashinfer:
|
||||
ret.init_triton_args(batch, prefix_lens)
|
||||
ret.init_triton_args(batch)
|
||||
|
||||
flashinfer_use_ragged = False
|
||||
if not model_runner.server_args.disable_flashinfer:
|
||||
@@ -198,30 +193,35 @@ class InputMetadata:
|
||||
):
|
||||
flashinfer_use_ragged = True
|
||||
ret.init_flashinfer_handlers(
|
||||
model_runner, prefix_lens, flashinfer_use_ragged
|
||||
model_runner, batch.prefix_lens_cpu, flashinfer_use_ragged
|
||||
)
|
||||
|
||||
return ret
|
||||
|
||||
def init_triton_args(self, batch: ScheduleBatch, prefix_lens):
|
||||
def init_triton_args(self, batch: ScheduleBatch):
|
||||
"""Init auxiliary variables for triton attention backend."""
|
||||
self.triton_max_seq_len = int(torch.max(self.seq_lens))
|
||||
self.triton_prefix_lens = prefix_lens
|
||||
self.triton_start_loc = torch.zeros_like(self.seq_lens, dtype=torch.int32)
|
||||
self.triton_start_loc[1:] = torch.cumsum(self.seq_lens[:-1], dim=0)
|
||||
|
||||
if self.forward_mode == ForwardMode.DECODE:
|
||||
self.triton_max_extend_len = None
|
||||
else:
|
||||
extend_seq_lens = self.seq_lens - prefix_lens
|
||||
self.triton_prefix_lens = torch.tensor(batch.prefix_lens_cpu, device="cuda")
|
||||
extend_seq_lens = self.seq_lens - self.triton_prefix_lens
|
||||
self.triton_max_extend_len = int(torch.max(extend_seq_lens))
|
||||
|
||||
def init_flashinfer_handlers(
|
||||
self,
|
||||
model_runner,
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
flashinfer_use_ragged,
|
||||
):
|
||||
if self.forward_mode != ForwardMode.DECODE:
|
||||
prefix_lens = torch.tensor(prefix_lens_cpu, device="cuda")
|
||||
else:
|
||||
prefix_lens = None
|
||||
|
||||
update_flashinfer_indices(
|
||||
self.forward_mode,
|
||||
model_runner,
|
||||
|
||||
Reference in New Issue
Block a user