[Ascend]Support of piecewise graph compilation for prefill on NPU (#12287)
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
@@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||
from sglang.srt.utils import get_available_gpu_memory, is_npu, log_info_on_rank0
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -303,7 +303,7 @@ class PiecewiseCudaGraphRunner:
|
||||
seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
next_token_logits_buffer=None,
|
||||
orig_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
seq_lens_cpu=torch.tensor([num_tokens]),
|
||||
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
req_to_token_pool=self.model_runner.req_to_token_pool,
|
||||
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
||||
attn_backend=self.model_runner.attn_backend,
|
||||
@@ -316,9 +316,9 @@ class PiecewiseCudaGraphRunner:
|
||||
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
||||
extend_start_loc=torch.tensor([0], device=self.device),
|
||||
extend_prefix_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_seq_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_logprob_start_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_prefix_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
extend_seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
extend_logprob_start_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
positions=torch.arange(num_tokens, device=self.device),
|
||||
global_num_tokens_gpu=None,
|
||||
global_num_tokens_for_logprob_gpu=None,
|
||||
@@ -347,7 +347,7 @@ class PiecewiseCudaGraphRunner:
|
||||
)
|
||||
|
||||
def _cache_loc_dtype(self):
|
||||
return torch.int64
|
||||
return torch.int64 if not is_npu() else torch.int32
|
||||
|
||||
def can_run(self, forward_batch: ForwardBatch):
|
||||
num_tokens = len(forward_batch.input_ids)
|
||||
@@ -432,7 +432,7 @@ class PiecewiseCudaGraphRunner:
|
||||
seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
next_token_logits_buffer=None,
|
||||
orig_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
seq_lens_cpu=torch.tensor([num_tokens]),
|
||||
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
req_to_token_pool=self.model_runner.req_to_token_pool,
|
||||
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
||||
attn_backend=self.model_runner.attn_backend,
|
||||
@@ -445,9 +445,9 @@ class PiecewiseCudaGraphRunner:
|
||||
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
||||
extend_start_loc=torch.tensor([0], device=self.device),
|
||||
extend_prefix_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_seq_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_logprob_start_lens_cpu=torch.tensor([num_tokens]),
|
||||
extend_prefix_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
extend_seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
extend_logprob_start_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||
positions=positions,
|
||||
global_num_tokens_gpu=None,
|
||||
global_num_tokens_for_logprob_gpu=None,
|
||||
|
||||
Reference in New Issue
Block a user