Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
67 changes: 59 additions & 8 deletions examples/bench.py
Original file line number Diff line number Diff line change
Expand Up @@ -292,10 +292,44 @@ def resize_benchmark_prompt(
return prefix_ids + repeat_prompt(content_ids, content_length) + suffix_ids


def convert_minicpmv_inputs(model, processed_inputs):
"""Convert one MiniCPM-V image into reusable low-level input tensors."""
def convert_multimodal_inputs(model, processor, processed_inputs):
"""Convert one image into reusable low-level input tensors."""
if processed_inputs is None:
return {}

if model.model_type == "videonsa":
import torch

pixel_values = processed_inputs.get("pixel_values")
image_bound = processed_inputs.get("image_bound")
image_grid_thw = processed_inputs.get("image_grid_thw")
if pixel_values is None or image_bound is None or image_grid_thw is None:
raise ValueError("VideoNSA image preprocessing returned incomplete inputs")

valid_bounds = image_bound[0]
valid_bounds = valid_bounds[valid_bounds[:, 1] > valid_bounds[:, 0]]
if len(valid_bounds) != len(image_grid_thw):
raise ValueError(
"VideoNSA image token ranges do not match the preprocessed images"
)

expected_patches = sum(int(grid.prod().item()) for grid in image_grid_thw)
if len(pixel_values) != expected_patches:
raise ValueError("VideoNSA image patch count does not match image_grid_thw")

pixel_dtype = getattr(processor, "pixel_values_dtype", torch.bfloat16)
return {
"pixel_values": infinicore.from_torch(
pixel_values.to(dtype=pixel_dtype).contiguous()
),
"image_bound": infinicore.from_torch(
valid_bounds.unsqueeze(0).to(torch.int64).contiguous()
),
"tgt_sizes": infinicore.from_torch(
image_grid_thw.to(torch.int64).contiguous()
),
}

if model.model_type != "minicpmv":
raise ValueError(
f"--image is not supported by bench.py for model_type={model.model_type!r}"
Expand Down Expand Up @@ -398,6 +432,7 @@ def __init__(
self.processor = processor
self.tokenizer = tokenizer
self.prompt_token_segments = prompt_token_segments
self.processed_multimodal_inputs = processed_multimodal_inputs
self.multimodal_inputs = {}
self.pp = pp

Expand Down Expand Up @@ -478,8 +513,8 @@ def __init__(
weight_load_mode=weight_load_mode,
pre_transpose=pre_transpose,
)
self.multimodal_inputs = convert_minicpmv_inputs(
model, processed_multimodal_inputs
self.multimodal_inputs = convert_multimodal_inputs(
model, self.processor, processed_multimodal_inputs
)

# ---------------------------------------------------------------------------- #
Expand Down Expand Up @@ -519,16 +554,32 @@ def __init__(
self.weight_load_mode = weight_load_mode
self.skip_load = skip_load

def get_multimodal_inputs(self, batch_size: int):
def get_multimodal_inputs(self, batch_size: int, prompt_ids=None):
if not self.multimodal_inputs:
return {}

return {
inputs = {
"pixel_values": [self.multimodal_inputs["pixel_values"]] * batch_size,
"image_bound": [self.multimodal_inputs["image_bound"]] * batch_size,
"tgt_sizes": [self.multimodal_inputs["tgt_sizes"]] * batch_size,
"image_req_ids": list(range(batch_size)),
}
if self.model.model_type == "videonsa":
if prompt_ids is None:
raise ValueError("VideoNSA multimodal generation requires prompt IDs")
position_ids = self.processor._prompt_mrope_positions(
prompt_ids, self.processed_multimodal_inputs
)
if any(len(axis) != len(prompt_ids) for axis in position_ids):
raise ValueError("VideoNSA mRoPE positions do not match the prompt")
position_id_delta = (
max(max(axis) for axis in position_ids) + 1 - len(prompt_ids)
)
inputs.update(
prompt_position_ids=position_ids,
position_id_delta=position_id_delta,
)
return inputs

@property
def uses_pipeline_parallel(self) -> bool:
Expand Down Expand Up @@ -617,7 +668,7 @@ def run(
temperature=temperature,
stop_on_eos=False,
),
**self.get_multimodal_inputs(batch_size),
**self.get_multimodal_inputs(batch_size, input_ids),
_measure_and_log_time=True,
)
t2 = time.time()
Expand Down Expand Up @@ -902,7 +953,7 @@ def run(
top_p=cfg.top_p,
stop_on_eos=False,
),
**test.get_multimodal_inputs(warmup_batch),
**test.get_multimodal_inputs(warmup_batch, warmup_prompt_ids),
_measure_and_log_time=False,
)

Expand Down
55 changes: 47 additions & 8 deletions python/infinilm/infer_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,9 @@ def _infer_position_id_axes(hf_config: dict) -> int:
raise ValueError("position_id_axes must be positive")
return axes

rope_parameters = text_config.get("rope_parameters") or {}
rope_parameters = (
text_config.get("rope_parameters") or text_config.get("rope_scaling") or {}
)
mrope_section = rope_parameters.get("mrope_section")
if isinstance(mrope_section, (list, tuple)) and mrope_section:
return len(mrope_section)
Expand Down Expand Up @@ -502,6 +504,8 @@ def generate(
image_bound=None,
tgt_sizes=None,
image_req_ids=None,
prompt_position_ids=None,
position_id_delta=0,
_measure_and_log_time=False,
):
eos_token_id = self.eos_token_id
Expand All @@ -517,17 +521,31 @@ def generate(
"When `batch_size > 1`, `max_new_tokens` must be specified."
)

if prompt_position_ids is not None:
if not self.enable_paged_attn:
raise ValueError(
"Custom prompt position IDs currently require paged attention"
)
if len(prompt_position_ids) != self.position_id_axes:
raise ValueError(
f"Expected {self.position_id_axes} position ID axes, got "
f"{len(prompt_position_ids)}"
)
if any(len(axis) != initial_seqlen for axis in prompt_position_ids):
raise ValueError("Prompt position IDs must match the input length")

if _measure_and_log_time:
time_measurements = []

block_tables = None
max_blocks_per_batch = 0
mamba_state_indices = None
if self.has_mamba_cache:
if not self.enable_paged_attn:
if self.has_mamba_cache and not self.enable_paged_attn:
if self.model_type != "mamba":
raise RuntimeError(
"Low-level generate for mamba-cache models currently requires paged attention"
)
elif self.has_mamba_cache:
mamba_pool_size = max(2, self.get_cache_config().num_blocks() // 4)
if batch_size > mamba_pool_size - 1:
raise RuntimeError(
Expand Down Expand Up @@ -560,11 +578,32 @@ def generate(

if self.enable_paged_attn:
input_ids = input_ids.view([1, batch_size * seq_len])
position_ids_list = (
list(range(past_seq_len, past_seq_len + seq_len)) * batch_size
)
if self.position_id_axes > 1:
position_ids_list = [position_ids_list] * self.position_id_axes
if prompt_position_ids is not None:
if iter == 0:
position_ids_list = [
list(axis) * batch_size for axis in prompt_position_ids
]
else:
decode_positions = (
list(
range(
past_seq_len + position_id_delta,
past_seq_len + position_id_delta + seq_len,
)
)
* batch_size
)
position_ids_list = [
decode_positions for _ in range(self.position_id_axes)
]
else:
position_ids_list = (
list(range(past_seq_len, past_seq_len + seq_len)) * batch_size
)
if self.position_id_axes > 1:
position_ids_list = [
position_ids_list for _ in range(self.position_id_axes)
]
position_ids = infinicore.from_list(
position_ids_list, dtype=infinicore.int64
)
Expand Down
Loading