From 8b3977adf885272289e8434a9ea5ed47b9d95960 Mon Sep 17 00:00:00 2001 From: wooway777 Date: Tue, 18 Aug 2026 09:13:26 +0000 Subject: [PATCH] fix: resupport bench mamaba and nsa --- examples/bench.py | 67 +++++++++++++++++++++++++++++---- python/infinilm/infer_engine.py | 55 +++++++++++++++++++++++---- 2 files changed, 106 insertions(+), 16 deletions(-) diff --git a/examples/bench.py b/examples/bench.py index 5e1ecfd37..5235dc922 100644 --- a/examples/bench.py +++ b/examples/bench.py @@ -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}" @@ -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 @@ -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 ) # ---------------------------------------------------------------------------- # @@ -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: @@ -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() @@ -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, ) diff --git a/python/infinilm/infer_engine.py b/python/infinilm/infer_engine.py index a9b195b08..5e4bc8caa 100644 --- a/python/infinilm/infer_engine.py +++ b/python/infinilm/infer_engine.py @@ -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) @@ -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 @@ -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( @@ -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 )