@@ -551,6 +551,7 @@ def dummy_run(
551551 query_start_loc = self .query_start_loc .gpu [: num_reqs + 1 ],
552552 query_start_loc_cpu = self .query_start_loc .cpu [: num_reqs + 1 ],
553553 seq_lens_cpu = self .runner .optimistic_seq_lens_cpu ,
554+ _seq_lens_cpu = self .runner .optimistic_seq_lens_cpu ,
554555 seq_lens_cpu_upper_bound = self .runner .optimistic_seq_lens_cpu ,
555556 seq_lens = self .runner .seq_lens [:num_reqs ],
556557 num_reqs = num_reqs ,
@@ -559,11 +560,14 @@ def dummy_run(
559560 max_query_len = self .num_speculative_tokens + 1 ,
560561 num_computed_tokens_cpu = num_computed_tokens_cpu ,
561562 actual_seq_lengths_q = self .runner .actual_seq_lengths_q ,
562- block_table_tensor = self .runner .input_batch .block_table [0 ].get_device_tensor ()[:num_reqs ],
563+ block_table_tensor = self .runner .input_batch .block_table [self .kv_cache_gid ].get_device_tensor ()[
564+ :num_reqs
565+ ],
563566 # This is used to hold a position.
564- slot_mapping = self .runner .input_batch .block_table [0 ].slot_mapping .gpu ,
565- slot_mapping_cpu = self .runner .input_batch .block_table [0 ].slot_mapping .cpu ,
567+ slot_mapping = self .runner .input_batch .block_table [self . kv_cache_gid ].slot_mapping .gpu ,
568+ slot_mapping_cpu = self .runner .input_batch .block_table [self . kv_cache_gid ].slot_mapping .cpu ,
566569 positions = self .runner .positions ,
570+ positions_cpu = self .runner ._dsa_positions_cpu_buf if self .use_compress else None ,
567571 attn_state = self .runner .attn_state ,
568572 decode_token_per_req = self .runner .decode_token_per_req ,
569573 is_prefilling = torch .zeros (num_reqs , dtype = torch .bool ),
@@ -575,20 +579,46 @@ def dummy_run(
575579
576580 assert len (self .draft_attn_groups ) > 0
577581 builder = self .draft_attn_groups [0 ].get_metadata_builder ()
582+ kv_cache_spec = self .draft_attn_groups [0 ].kv_cache_spec
578583 # update the tensor's address for each step.
579584 for draft_index in range (self .num_speculative_tokens ):
580585 common_attn_metadata = self .shallow_copy_metadata (common_attn_metadata )
586+ extra_attn_metadata_args : dict = {}
587+ if self .use_compress :
588+ extra_attn_metadata_args = dict (
589+ prefill_ratio_to_sas_metadata = dict (),
590+ decode_ratio_to_sas_metadata = dict (),
591+ common_ratio_to_sas_metadata = dict (),
592+ block_size = kv_cache_spec .block_size ,
593+ )
581594 # Set the real slot_mapping.
595+ slot_mapping_lens = common_attn_metadata .slot_mapping .shape [0 ]
596+ self .slot_mapping_group [draft_index ][:slot_mapping_lens ].copy_ (common_attn_metadata .slot_mapping )
597+ self .slot_mapping_group [draft_index ][slot_mapping_lens :].fill_ (PADDING_SLOT_ID )
582598 common_attn_metadata .slot_mapping = self .slot_mapping_group [draft_index ]
599+ self .seq_lens_group [draft_index ][:num_reqs ].copy_ (common_attn_metadata .seq_lens )
600+ self .seq_lens_group [draft_index ][num_reqs :].fill_ (0 )
583601 common_attn_metadata .seq_lens = self .seq_lens_group [draft_index ][:num_reqs ]
602+ self .query_start_loc_group [draft_index ][: num_reqs + 1 ].copy_ (common_attn_metadata .query_start_loc )
603+ self .query_start_loc_group [draft_index ][num_reqs + 1 :].fill_ (0 )
584604 common_attn_metadata .query_start_loc = self .query_start_loc_group [draft_index ][: num_reqs + 1 ]
585605 if self .pcp_size * self .dcp_size > 1 and draft_index > 0 :
586606 assert self .block_table_tensor_clone is not None , "block_table_tensor_clone is not init"
587607 common_attn_metadata .block_table_tensor = self .block_table_tensor_clone [:num_reqs ]
588- attn_metadata_eagle = builder .build_for_graph_capture (
589- common_attn_metadata ,
590- AscendAttentionState .SpecDecoding if self .method == "mtp" else AscendAttentionState .ChunkedPrefill ,
591- )
608+ if not self .use_compress or draft_index == 0 :
609+ attn_metadata_eagle = builder .build_for_graph_capture (
610+ common_attn_metadata ,
611+ AscendAttentionState .SpecDecoding
612+ if self .method == "mtp"
613+ else AscendAttentionState .ChunkedPrefill ,
614+ ** extra_attn_metadata_args ,
615+ )
616+ else :
617+ attn_metadata_eagle = builder .build_for_drafting (
618+ common_attn_metadata ,
619+ draft_index ,
620+ ** extra_attn_metadata_args ,
621+ )
592622 per_layer_attn_metadata = dict ()
593623 for layer_name in self .attn_layer_names :
594624 per_layer_attn_metadata [layer_name ] = attn_metadata_eagle
@@ -744,16 +774,20 @@ def _propose(
744774 # is run in eager mode currently, which means `_pad_query_start_loc_for_fia` is not called,
745775 # while draft model is run in graph model, which means we should pad the `query_start_loc`.
746776 # Need to be fixed in the future.
777+ num_reqs = common_attn_metadata .query_start_loc .shape [0 ]
778+ self .query_start_loc .gpu [:num_reqs ].copy_ (common_attn_metadata .query_start_loc )
779+ self .query_start_loc .cpu [:num_reqs ].copy_ (common_attn_metadata .query_start_loc_cpu )
747780 num_reqs_padded = self .runner ._pad_query_start_loc_for_fia (
781+ self .query_start_loc ,
748782 num_input_tokens ,
749783 batch_descriptor .num_reqs if batch_descriptor .num_reqs is not None else common_attn_metadata .num_reqs ,
750784 common_attn_metadata .num_reqs ,
751785 aclgraph_runtime_mode ,
752786 batch_descriptor .num_reqs ,
753787 )
754788 common_attn_metadata .num_reqs = num_reqs_padded
755- common_attn_metadata .query_start_loc = self .runner . query_start_loc .gpu [: num_reqs_padded + 1 ]
756- common_attn_metadata .query_start_loc_cpu = self .runner . query_start_loc .cpu [: num_reqs_padded + 1 ]
789+ common_attn_metadata .query_start_loc = self .query_start_loc .gpu [: num_reqs_padded + 1 ]
790+ common_attn_metadata .query_start_loc_cpu = self .query_start_loc .cpu [: num_reqs_padded + 1 ]
757791 slicing_length = (
758792 num_reqs_padded * self .decode_threshold if self .pcp_size * self .dcp_size > 1 else num_reqs_padded
759793 )
0 commit comments