[NPU]: Optimize xgrammar token bitmask on NPU with AscendC (#24133)
This commit is contained in:
@@ -1,33 +0,0 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def apply_token_bitmask_inplace_torch(
|
||||
logits: torch.Tensor,
|
||||
bitmask: torch.Tensor,
|
||||
) -> None:
|
||||
"""Backend-agnostic torch fallback for packed-bitmask application.
|
||||
|
||||
This path is currently used as a fallback on NPU in xgrammar backend.
|
||||
"""
|
||||
vocab_size = logits.shape[-1]
|
||||
bitmask_cpu = bitmask.detach().cpu()
|
||||
token_ids = torch.arange(vocab_size, device="cpu", dtype=torch.int32)
|
||||
word_idx = token_ids // 32
|
||||
bit_idx = token_ids % 32
|
||||
words = bitmask_cpu[:, word_idx].to(torch.int32)
|
||||
allowed = ((words >> bit_idx) & 1).to(torch.bool)
|
||||
allowed = allowed.to(logits.device, non_blocking=True)
|
||||
logits.masked_fill_(~allowed, float("-inf"))
|
||||
@@ -35,13 +35,11 @@ from sglang.srt.constrained.base_grammar_backend import (
|
||||
GrammarStats,
|
||||
InvalidGrammarObject,
|
||||
)
|
||||
from sglang.srt.constrained.torch_ops.bitmask_ops import (
|
||||
apply_token_bitmask_inplace_torch,
|
||||
)
|
||||
from sglang.srt.constrained.utils import is_legacy_structural_tag
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
if _is_hip:
|
||||
from sgl_kernel import apply_token_bitmask_inplace_cuda
|
||||
else:
|
||||
@@ -118,7 +116,9 @@ class XGrammarGrammar(BaseGrammarObject):
|
||||
else:
|
||||
apply_token_bitmask_inplace_triton(logits, vocab_mask)
|
||||
elif logits.device.type == "npu":
|
||||
apply_token_bitmask_inplace_torch(logits, vocab_mask)
|
||||
import sgl_kernel_npu # noqa: F401
|
||||
|
||||
torch.ops.npu.apply_token_bitmask(logits, vocab_mask)
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported device: {logits.device.type}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user