[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,
|
GrammarStats,
|
||||||
InvalidGrammarObject,
|
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.constrained.utils import is_legacy_structural_tag
|
||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_hip
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
from sgl_kernel import apply_token_bitmask_inplace_cuda
|
from sgl_kernel import apply_token_bitmask_inplace_cuda
|
||||||
else:
|
else:
|
||||||
@@ -118,7 +116,9 @@ class XGrammarGrammar(BaseGrammarObject):
|
|||||||
else:
|
else:
|
||||||
apply_token_bitmask_inplace_triton(logits, vocab_mask)
|
apply_token_bitmask_inplace_triton(logits, vocab_mask)
|
||||||
elif logits.device.type == "npu":
|
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:
|
else:
|
||||||
raise RuntimeError(f"Unsupported device: {logits.device.type}")
|
raise RuntimeError(f"Unsupported device: {logits.device.type}")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user