0
vllm_ascend/ops/triton/linearnorm/__init__.py
Normal file
0
vllm_ascend/ops/triton/linearnorm/__init__.py
Normal file
427
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_mrope.py
Normal file
427
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_mrope.py
Normal file
@@ -0,0 +1,427 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
|
||||
import torch
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import extract_slice, get_vectorcore_num, insert_slice
|
||||
|
||||
|
||||
@triton.jit(
|
||||
do_not_specialize=["num_tokens", "front_core_num", "num_tokens_each_front_core", "num_tokens_each_tail_core"]
|
||||
)
|
||||
def split_qkv_rmsnorm_mrope_kernel(
|
||||
in_qkv_ptr: torch.Tensor,
|
||||
q_weight_ptr: torch.Tensor,
|
||||
q_bias_ptr: torch.Tensor,
|
||||
k_weight_ptr: torch.Tensor,
|
||||
k_bias_ptr: torch.Tensor,
|
||||
cos_sin_ptr: torch.Tensor,
|
||||
out_q_ptr: torch.Tensor,
|
||||
out_k_ptr: torch.Tensor,
|
||||
out_v_ptr: torch.Tensor,
|
||||
out_gate_ptr: torch.Tensor,
|
||||
num_tokens,
|
||||
front_core_num,
|
||||
num_tokens_each_front_core,
|
||||
num_tokens_each_tail_core,
|
||||
num_q_heads: tl.constexpr,
|
||||
num_kv_heads: tl.constexpr,
|
||||
head_size: tl.constexpr,
|
||||
q_size: tl.constexpr,
|
||||
kv_size: tl.constexpr,
|
||||
eps: tl.constexpr,
|
||||
mrope_section_t,
|
||||
mrope_section_h,
|
||||
mrope_section_w,
|
||||
has_bias: tl.constexpr,
|
||||
is_interleaved: tl.constexpr,
|
||||
rope_dim: tl.constexpr,
|
||||
half_rope_dim: tl.constexpr,
|
||||
IS_PARTIAL_ROPE: tl.constexpr,
|
||||
gate_size: tl.constexpr,
|
||||
):
|
||||
block_idx = tl.program_id(0)
|
||||
|
||||
loop_num = num_tokens_each_front_core
|
||||
if block_idx >= front_core_num:
|
||||
loop_num = num_tokens_each_tail_core
|
||||
|
||||
block_offset = num_tokens_each_front_core * block_idx
|
||||
if block_idx >= front_core_num:
|
||||
block_offset = (
|
||||
num_tokens_each_front_core * front_core_num + (block_idx - front_core_num) * num_tokens_each_tail_core
|
||||
)
|
||||
|
||||
q_rmsnorm_weight = tl.load(q_weight_ptr + tl.arange(0, head_size))
|
||||
k_rmsnorm_weight = tl.load(k_weight_ptr + tl.arange(0, head_size))
|
||||
|
||||
if has_bias:
|
||||
q_bias = tl.load(q_bias_ptr + tl.arange(0, head_size))
|
||||
k_bias = tl.load(k_bias_ptr + tl.arange(0, head_size))
|
||||
|
||||
for index in range(loop_num):
|
||||
## load ##
|
||||
# q
|
||||
in_q_offset = in_qkv_ptr + (block_offset + index) * (q_size + gate_size + 2 * kv_size)
|
||||
if gate_size > 0:
|
||||
in_q_gate_tensor = (
|
||||
tl.load(in_q_offset + tl.arange(0, q_size + gate_size))
|
||||
.to(tl.float32)
|
||||
.reshape(num_q_heads, head_size * 2)
|
||||
)
|
||||
in_q_tensor = extract_slice(
|
||||
in_q_gate_tensor,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_q_heads, head_size),
|
||||
strides=(1, 1),
|
||||
)
|
||||
in_gate_tensor = extract_slice(
|
||||
in_q_gate_tensor,
|
||||
offsets=(0, head_size),
|
||||
sizes=(num_q_heads, head_size),
|
||||
strides=(1, 1),
|
||||
).reshape(q_size)
|
||||
else:
|
||||
in_q_tensor = tl.load(in_q_offset + tl.arange(0, q_size)).to(tl.float32).reshape(num_q_heads, head_size)
|
||||
|
||||
# k
|
||||
in_k_offset = in_q_offset + q_size + gate_size
|
||||
in_k_tensor = tl.load(in_k_offset + tl.arange(0, kv_size)).to(tl.float32).reshape(num_kv_heads, head_size)
|
||||
# v
|
||||
in_v_offset = in_k_offset + kv_size
|
||||
in_v_tensor = tl.load(in_v_offset + tl.arange(0, kv_size))
|
||||
|
||||
# cos, sin
|
||||
cos_offsets = tl.arange(0, half_rope_dim)
|
||||
if is_interleaved:
|
||||
h_mask = ((cos_offsets % 3) == 1) & (cos_offsets <= 3 * mrope_section_h)
|
||||
w_mask = ((cos_offsets % 3) == 2) & (cos_offsets <= 3 * mrope_section_w)
|
||||
t_mask = ~(h_mask | w_mask)
|
||||
else:
|
||||
t_mask = cos_offsets < mrope_section_t
|
||||
h_mask = (mrope_section_t - 1 < cos_offsets) & (cos_offsets < mrope_section_t + mrope_section_h)
|
||||
w_mask = (mrope_section_t + mrope_section_h - 1 < cos_offsets) & (
|
||||
cos_offsets < mrope_section_t + mrope_section_h + mrope_section_w
|
||||
)
|
||||
|
||||
t_cos_offset = cos_sin_ptr + (block_offset + index) * rope_dim
|
||||
h_cos_offset = t_cos_offset + num_tokens * rope_dim
|
||||
w_cos_offset = h_cos_offset + num_tokens * rope_dim
|
||||
|
||||
t_sin_offset = cos_sin_ptr + (block_offset + index) * rope_dim + half_rope_dim
|
||||
h_sin_offset = t_sin_offset + num_tokens * rope_dim
|
||||
w_sin_offset = h_sin_offset + num_tokens * rope_dim
|
||||
|
||||
t_cos_tensor = tl.load(t_cos_offset + cos_offsets, mask=t_mask, other=0)
|
||||
h_cos_tensor = tl.load(h_cos_offset + cos_offsets, mask=h_mask, other=0)
|
||||
w_cos_tensor = tl.load(w_cos_offset + cos_offsets, mask=w_mask, other=0)
|
||||
t_sin_tensor = tl.load(t_sin_offset + cos_offsets, mask=t_mask, other=0)
|
||||
h_sin_tensor = tl.load(h_sin_offset + cos_offsets, mask=h_mask, other=0)
|
||||
w_sin_tensor = tl.load(w_sin_offset + cos_offsets, mask=w_mask, other=0)
|
||||
|
||||
cos_tensor = (t_cos_tensor + h_cos_tensor + w_cos_tensor).to(tl.float32).reshape(1, half_rope_dim)
|
||||
cos_tensor = tl.broadcast_to(cos_tensor, (2, half_rope_dim)).reshape(1, rope_dim)
|
||||
|
||||
sin_tensor = (t_sin_tensor + h_sin_tensor + w_sin_tensor).to(tl.float32).reshape(1, half_rope_dim)
|
||||
sin_tensor = tl.broadcast_to(sin_tensor, (2, half_rope_dim)).reshape(1, rope_dim)
|
||||
|
||||
## compute ##
|
||||
# q-rmsnorm
|
||||
squares = in_q_tensor * in_q_tensor
|
||||
variances = tl.sum(squares, axis=1) / head_size
|
||||
reciprocal_std = (1 / tl.sqrt(variances + eps)).reshape(num_q_heads, 1)
|
||||
q_normalized = in_q_tensor * reciprocal_std
|
||||
q_normalized = q_normalized * q_rmsnorm_weight
|
||||
if has_bias:
|
||||
q_normalized = q_normalized + q_bias
|
||||
|
||||
# k-rmsnorm
|
||||
squares = in_k_tensor * in_k_tensor
|
||||
variances = tl.sum(squares, axis=1) / head_size
|
||||
reciprocal_std = (1 / tl.sqrt(variances + eps)).reshape(num_kv_heads, 1)
|
||||
k_normalized = in_k_tensor * reciprocal_std
|
||||
k_normalized = k_normalized * k_rmsnorm_weight
|
||||
if has_bias:
|
||||
k_normalized = k_normalized + k_bias
|
||||
|
||||
# q-mrope
|
||||
x1 = extract_slice(
|
||||
q_normalized,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_q_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
x2 = extract_slice(
|
||||
q_normalized,
|
||||
offsets=(0, half_rope_dim),
|
||||
sizes=(num_q_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
cat_x = tl.zeros((num_q_heads, rope_dim), dtype=tl.float32)
|
||||
cat_x = insert_slice(
|
||||
cat_x,
|
||||
-x2,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_q_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
cat_x = insert_slice(
|
||||
cat_x,
|
||||
x1,
|
||||
offsets=(0, half_rope_dim),
|
||||
sizes=(num_q_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
if IS_PARTIAL_ROPE:
|
||||
orig_qk = extract_slice(
|
||||
q_normalized,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_q_heads, rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
else:
|
||||
orig_qk = q_normalized
|
||||
roped_q = cat_x * sin_tensor + orig_qk * cos_tensor
|
||||
|
||||
# k-mrope
|
||||
y1 = extract_slice(
|
||||
k_normalized,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_kv_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
y2 = extract_slice(
|
||||
k_normalized,
|
||||
offsets=(0, half_rope_dim),
|
||||
sizes=(num_kv_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
cat_y = tl.zeros((num_kv_heads, rope_dim), dtype=tl.float32)
|
||||
cat_y = insert_slice(
|
||||
cat_y,
|
||||
-y2,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_kv_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
cat_y = insert_slice(
|
||||
cat_y,
|
||||
y1,
|
||||
offsets=(0, half_rope_dim),
|
||||
sizes=(num_kv_heads, half_rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
if IS_PARTIAL_ROPE:
|
||||
orig_qk = extract_slice(
|
||||
k_normalized,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_kv_heads, rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
else:
|
||||
orig_qk = k_normalized
|
||||
roped_k = cat_y * sin_tensor + orig_qk * cos_tensor
|
||||
|
||||
if IS_PARTIAL_ROPE:
|
||||
q_normalized = insert_slice(
|
||||
q_normalized,
|
||||
roped_q,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_q_heads, rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
k_normalized = insert_slice(
|
||||
k_normalized,
|
||||
roped_k,
|
||||
offsets=(0, 0),
|
||||
sizes=(num_kv_heads, rope_dim),
|
||||
strides=(1, 1),
|
||||
)
|
||||
else:
|
||||
q_normalized = roped_q
|
||||
k_normalized = roped_k
|
||||
|
||||
## store ##
|
||||
# out_q
|
||||
out_q_offset = out_q_ptr + (block_offset + index) * q_size
|
||||
out_q_indices = tl.arange(0, q_size)
|
||||
tl.store(out_q_offset + out_q_indices, q_normalized.reshape(q_size))
|
||||
|
||||
# out_k
|
||||
out_k_offset = out_k_ptr + (block_offset + index) * kv_size
|
||||
out_k_indices = tl.arange(0, kv_size)
|
||||
tl.store(out_k_offset + out_k_indices, k_normalized.reshape(kv_size))
|
||||
|
||||
# out_v
|
||||
out_v_offset = out_v_ptr + (block_offset + index) * kv_size
|
||||
tl.store(out_v_offset + tl.arange(0, kv_size), in_v_tensor)
|
||||
|
||||
# out_gate
|
||||
if gate_size > 0:
|
||||
out_gate_offset = out_gate_ptr + (block_offset + index) * gate_size
|
||||
tl.store(out_gate_offset + tl.arange(0, gate_size), in_gate_tensor)
|
||||
|
||||
|
||||
def triton_split_qkv_rmsnorm_mrope(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
cos_sin: torch.Tensor,
|
||||
num_q_heads: int,
|
||||
num_kv_heads: int,
|
||||
head_size: int,
|
||||
eps: float,
|
||||
mrope_section: list[int],
|
||||
is_interleaved: bool,
|
||||
rope_dim: int | None = None,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
has_gate: bool = False,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
core_num = get_vectorcore_num()
|
||||
|
||||
q_size = num_q_heads * head_size
|
||||
kv_size = num_kv_heads * head_size
|
||||
num_tokens = qkv.shape[0]
|
||||
|
||||
gate_size = q_size if has_gate else 0
|
||||
|
||||
if rope_dim is None:
|
||||
rope_dim = head_size
|
||||
IS_PARTIAL_ROPE = rope_dim != head_size
|
||||
|
||||
front_core_num = core_num
|
||||
if num_tokens % core_num != 0:
|
||||
front_core_num = num_tokens % core_num
|
||||
|
||||
num_tokens_each_front_core = (num_tokens + core_num - 1) // core_num
|
||||
|
||||
tail_core_num = 0
|
||||
if num_tokens > core_num:
|
||||
tail_core_num = core_num - front_core_num
|
||||
|
||||
num_tokens_each_tail_core = num_tokens // core_num
|
||||
|
||||
q_output = torch.empty(num_tokens, q_size, device=qkv.device, dtype=qkv.dtype)
|
||||
k_output = torch.empty(num_tokens, kv_size, device=qkv.device, dtype=qkv.dtype)
|
||||
v_output = torch.empty(num_tokens, kv_size, device=qkv.device, dtype=qkv.dtype)
|
||||
gate_output = torch.empty(num_tokens, gate_size, device=qkv.device, dtype=qkv.dtype)
|
||||
|
||||
total_core = front_core_num + tail_core_num
|
||||
block_dim = core_num
|
||||
if total_core < core_num:
|
||||
block_dim = total_core
|
||||
|
||||
has_bias = q_bias is not None
|
||||
|
||||
split_qkv_rmsnorm_mrope_kernel[(block_dim,)](
|
||||
qkv,
|
||||
q_weight,
|
||||
q_bias,
|
||||
k_weight,
|
||||
k_bias,
|
||||
cos_sin,
|
||||
q_output,
|
||||
k_output,
|
||||
v_output,
|
||||
gate_output,
|
||||
num_tokens,
|
||||
front_core_num,
|
||||
num_tokens_each_front_core,
|
||||
num_tokens_each_tail_core,
|
||||
num_q_heads,
|
||||
num_kv_heads,
|
||||
head_size,
|
||||
q_size,
|
||||
kv_size,
|
||||
eps,
|
||||
mrope_section[0],
|
||||
mrope_section[1],
|
||||
mrope_section[2],
|
||||
has_bias,
|
||||
is_interleaved,
|
||||
rope_dim,
|
||||
rope_dim // 2,
|
||||
IS_PARTIAL_ROPE,
|
||||
gate_size,
|
||||
)
|
||||
|
||||
return q_output, k_output, v_output, gate_output
|
||||
|
||||
|
||||
def triton_split_qkv_rmsnorm_mrope_fake(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
cos_sin: torch.Tensor,
|
||||
num_q_heads: int,
|
||||
num_kv_heads: int,
|
||||
head_size: int,
|
||||
eps: float,
|
||||
mrope_section: list[int],
|
||||
is_interleaved: bool,
|
||||
rope_dim: int | None = None,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
has_gate: bool = False,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
num_tokens = qkv.shape[0]
|
||||
q_size = num_q_heads * head_size
|
||||
kv_size = num_kv_heads * head_size
|
||||
gate_size = q_size if has_gate else 0
|
||||
|
||||
q_output = torch.empty(
|
||||
num_tokens,
|
||||
q_size,
|
||||
device=qkv.device,
|
||||
dtype=qkv.dtype,
|
||||
)
|
||||
|
||||
k_output = torch.empty(
|
||||
num_tokens,
|
||||
kv_size,
|
||||
device=qkv.device,
|
||||
dtype=qkv.dtype,
|
||||
)
|
||||
|
||||
v_output = torch.empty(
|
||||
num_tokens,
|
||||
kv_size,
|
||||
device=qkv.device,
|
||||
dtype=qkv.dtype,
|
||||
)
|
||||
|
||||
gate_output = torch.empty(
|
||||
num_tokens,
|
||||
gate_size,
|
||||
device=qkv.device,
|
||||
dtype=qkv.dtype,
|
||||
)
|
||||
|
||||
return q_output, k_output, v_output, gate_output
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="triton_split_qkv_rmsnorm_mrope",
|
||||
op_func=triton_split_qkv_rmsnorm_mrope,
|
||||
fake_impl=triton_split_qkv_rmsnorm_mrope_fake,
|
||||
mutates_args=[],
|
||||
dispatch_key="PrivateUse1",
|
||||
)
|
||||
390
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_rope.py
Normal file
390
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_rope.py
Normal file
@@ -0,0 +1,390 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
import torch
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import extract_slice, get_element, get_vectorcore_num, insert_slice
|
||||
|
||||
|
||||
@triton.jit
|
||||
def split_qkv_rmsnorm_rope_kernel(
|
||||
input_gm_ptr,
|
||||
q_gm_ptr,
|
||||
k_gm_ptr,
|
||||
v_gm_ptr,
|
||||
q_weight_ptr,
|
||||
q_bias_ptr,
|
||||
k_weight_ptr,
|
||||
k_bias_ptr,
|
||||
batch_size,
|
||||
q_hidden_size: tl.constexpr,
|
||||
kv_hidden_size: tl.constexpr,
|
||||
total_hidden_size: tl.constexpr,
|
||||
eps: tl.constexpr,
|
||||
BIAS: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
ROPE_DIM: tl.constexpr,
|
||||
HALF_ROPE_DIM: tl.constexpr,
|
||||
IS_PARTIAL_ROPE: tl.constexpr,
|
||||
num_vectorcore: tl.constexpr,
|
||||
batch_size_per_iter_per_vec: tl.constexpr,
|
||||
qk_head_nums_per_iter_per_vec: tl.constexpr,
|
||||
q_head_num: tl.constexpr,
|
||||
kv_head_num: tl.constexpr,
|
||||
qk_head_num_sum: tl.constexpr,
|
||||
v_batch_size_per_iter_per_vec: tl.constexpr,
|
||||
positions_gm_ptr,
|
||||
cos_sin_cache_gm_ptr,
|
||||
):
|
||||
row_pid = tl.program_id(0)
|
||||
|
||||
q_weight_values = tl.load(q_weight_ptr + tl.arange(0, HEAD_DIM))
|
||||
k_weight_values = tl.load(k_weight_ptr + tl.arange(0, HEAD_DIM))
|
||||
|
||||
batch_size_per_vec = tl.cdiv(batch_size, num_vectorcore)
|
||||
iter_num_per_vec = tl.cdiv(batch_size_per_vec, batch_size_per_iter_per_vec)
|
||||
v_iter_num_per_vec = tl.cdiv(batch_size_per_vec, v_batch_size_per_iter_per_vec)
|
||||
input_batch_offset = row_pid * batch_size_per_vec
|
||||
mblk_idx = tl.arange(0, batch_size_per_iter_per_vec) + input_batch_offset
|
||||
nblk_idx = tl.arange(0, q_hidden_size + kv_hidden_size)
|
||||
nmask = nblk_idx < total_hidden_size
|
||||
|
||||
input_batch_offset_end = min(input_batch_offset + batch_size_per_vec, batch_size)
|
||||
|
||||
pos_indices = input_batch_offset + tl.arange(0, batch_size_per_iter_per_vec)
|
||||
output_q_nblk_idx = tl.arange(0, q_hidden_size)
|
||||
output_q_nmask = output_q_nblk_idx < q_hidden_size
|
||||
output_kv_nblk_idx = tl.arange(0, kv_hidden_size)
|
||||
output_kv_nmask = output_kv_nblk_idx < kv_hidden_size
|
||||
sin_cos_range = tl.arange(0, ROPE_DIM)
|
||||
cos_sin_cache_offset = cos_sin_cache_gm_ptr + sin_cos_range
|
||||
|
||||
for iter in tl.range(iter_num_per_vec):
|
||||
pos_offset = iter * batch_size_per_iter_per_vec
|
||||
x = tl.load(
|
||||
positions_gm_ptr + pos_indices + pos_offset, mask=(pos_indices + pos_offset) < input_batch_offset_end
|
||||
)
|
||||
mmask = (mblk_idx + pos_offset) < input_batch_offset_end
|
||||
mask = (mmask[:, None]) & (nmask[None, :])
|
||||
idx = (mblk_idx + pos_offset)[:, None] * total_hidden_size + nblk_idx[None, :]
|
||||
values_tmp1 = tl.load(input_gm_ptr + idx, mask=mask).reshape(qk_head_nums_per_iter_per_vec, HEAD_DIM)
|
||||
if BIAS:
|
||||
q_bias_values = tl.load(q_bias_ptr + tl.arange(0, HEAD_DIM))
|
||||
k_bias_values = tl.load(k_bias_ptr + tl.arange(0, HEAD_DIM))
|
||||
|
||||
values_tmp3 = tl.zeros((batch_size_per_iter_per_vec, ROPE_DIM), dtype=tl.bfloat16)
|
||||
for i in tl.range(batch_size_per_iter_per_vec):
|
||||
pos = get_element(x, (i,))
|
||||
values_tmp3 = insert_slice(
|
||||
values_tmp3.reshape(batch_size_per_iter_per_vec, ROPE_DIM),
|
||||
tl.load(pos * ROPE_DIM + cos_sin_cache_offset[:, None]).reshape(1, ROPE_DIM),
|
||||
offsets=(i, 0),
|
||||
sizes=(1, ROPE_DIM),
|
||||
strides=(1, 1),
|
||||
)
|
||||
values_tmp3 = values_tmp3.reshape(batch_size_per_iter_per_vec, 1, ROPE_DIM)
|
||||
cos = extract_slice(
|
||||
values_tmp3,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, 1, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
sin = extract_slice(
|
||||
values_tmp3,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, 1, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
normalized_values = values_tmp1.to(tl.float32)
|
||||
normalized_values = normalized_values * normalized_values
|
||||
normalized_values = tl.sum(normalized_values, axis=1) / HEAD_DIM
|
||||
normalized_values = 1 / tl.sqrt(normalized_values + eps).reshape(qk_head_nums_per_iter_per_vec, 1)
|
||||
normalized_values = values_tmp1 * normalized_values
|
||||
|
||||
normalized_values_tmp = extract_slice(
|
||||
normalized_values.reshape(batch_size_per_iter_per_vec, qk_head_num_sum, HEAD_DIM),
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HEAD_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
if BIAS:
|
||||
normalized_values_tmp = (normalized_values_tmp * q_weight_values + q_bias_values).to(tl.bfloat16)
|
||||
else:
|
||||
normalized_values_tmp = (normalized_values_tmp * q_weight_values).to(tl.bfloat16)
|
||||
|
||||
# q rope
|
||||
values_tmp = tl.zeros((batch_size_per_iter_per_vec, q_head_num, ROPE_DIM), dtype=tl.bfloat16)
|
||||
x1 = extract_slice(
|
||||
normalized_values_tmp,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
x2 = extract_slice(
|
||||
normalized_values_tmp,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp = insert_slice(
|
||||
values_tmp,
|
||||
x1 * cos - x2 * sin,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp = insert_slice(
|
||||
values_tmp,
|
||||
x2 * cos + x1 * sin,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
q_output_idx = output_q_nblk_idx[None, :] + (mblk_idx + pos_offset)[:, None] * q_hidden_size
|
||||
mask = (mmask[:, None]) & (output_q_nmask[None, :])
|
||||
if IS_PARTIAL_ROPE:
|
||||
normalized_values_tmp = insert_slice(
|
||||
normalized_values_tmp,
|
||||
values_tmp,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(
|
||||
q_gm_ptr + q_output_idx,
|
||||
normalized_values_tmp.reshape(batch_size_per_iter_per_vec, q_hidden_size),
|
||||
mask=mask,
|
||||
)
|
||||
else:
|
||||
tl.store(
|
||||
q_gm_ptr + q_output_idx,
|
||||
values_tmp.reshape(batch_size_per_iter_per_vec, q_hidden_size),
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
# k rope
|
||||
normalized_values_tmp1 = extract_slice(
|
||||
normalized_values.reshape(batch_size_per_iter_per_vec, qk_head_num_sum, HEAD_DIM),
|
||||
offsets=(0, q_head_num, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HEAD_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
if BIAS:
|
||||
normalized_values_tmp1 = (normalized_values_tmp1 * k_weight_values + k_bias_values).to(tl.bfloat16)
|
||||
else:
|
||||
normalized_values_tmp1 = (normalized_values_tmp1 * k_weight_values).to(tl.bfloat16)
|
||||
|
||||
values_tmp2 = tl.zeros((batch_size_per_iter_per_vec, kv_head_num, ROPE_DIM), dtype=tl.bfloat16)
|
||||
|
||||
x1 = extract_slice(
|
||||
normalized_values_tmp1,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
x2 = extract_slice(
|
||||
normalized_values_tmp1,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp2 = insert_slice(
|
||||
values_tmp2,
|
||||
x1 * cos - x2 * sin,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp2 = insert_slice(
|
||||
values_tmp2,
|
||||
x2 * cos + x1 * sin,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
kv_output_idx = output_kv_nblk_idx[None, :] + (mblk_idx + pos_offset)[:, None] * kv_hidden_size
|
||||
mask = (mmask[:, None]) & (output_kv_nmask[None, :])
|
||||
if IS_PARTIAL_ROPE:
|
||||
normalized_values_tmp1 = insert_slice(
|
||||
normalized_values_tmp1,
|
||||
values_tmp2,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(
|
||||
k_gm_ptr + kv_output_idx,
|
||||
normalized_values_tmp1.reshape(batch_size_per_iter_per_vec, kv_hidden_size),
|
||||
mask=mask,
|
||||
)
|
||||
else:
|
||||
tl.store(
|
||||
k_gm_ptr + kv_output_idx,
|
||||
values_tmp2.reshape(batch_size_per_iter_per_vec, kv_hidden_size),
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
mblk_idx = tl.arange(0, v_batch_size_per_iter_per_vec) + input_batch_offset
|
||||
nblk_idx = tl.arange(q_hidden_size + kv_hidden_size, total_hidden_size)
|
||||
nmask = nblk_idx < total_hidden_size
|
||||
out_nblk_idx = tl.arange(0, kv_hidden_size)
|
||||
out_nmask = out_nblk_idx < kv_hidden_size
|
||||
|
||||
for _ in tl.range(v_iter_num_per_vec):
|
||||
mmask = mblk_idx < input_batch_offset_end
|
||||
mask = (mmask[:, None]) & (nmask[None, :])
|
||||
idx = mblk_idx[:, None] * total_hidden_size + nblk_idx[None, :]
|
||||
values = tl.load(input_gm_ptr + idx, mask=mask)
|
||||
out_idx = mblk_idx[:, None] * kv_hidden_size + out_nblk_idx[None, :]
|
||||
out_mask = (mmask[:, None]) & (out_nmask[None, :])
|
||||
tl.store(v_gm_ptr + out_idx, values, mask=out_mask)
|
||||
mblk_idx += v_batch_size_per_iter_per_vec
|
||||
|
||||
|
||||
def split_qkv_rmsnorm_rope_impl(
|
||||
input: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
eps: float,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# get available vector core
|
||||
num_vectorcore = get_vectorcore_num()
|
||||
rope_dim = cos_sin_cache.shape[-1]
|
||||
batch_size = input.shape[0]
|
||||
BIAS = q_bias is not None
|
||||
IS_PARTIAL_ROPE = rope_dim != head_dim
|
||||
# Q + K + V
|
||||
total_hidden_size = q_hidden_size + kv_hidden_size * 2
|
||||
|
||||
q_output = torch.empty(batch_size, q_hidden_size, device=input.device, dtype=input.dtype)
|
||||
k_output = torch.empty(batch_size, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
v_output = torch.empty(batch_size, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
|
||||
q_head_num = q_hidden_size // head_dim
|
||||
kv_head_num = kv_hidden_size // head_dim
|
||||
|
||||
# set number of line loading from GM data is x
|
||||
# x*(q_head_num + kv_head_num)*HEAD_DIM: values_tmp
|
||||
# 2x*(q_head_num + kv_head_num)*HEAD_DIM: normalized_values(float32)
|
||||
# x*ROPE_DIM*2 : cos/sin
|
||||
# x*q_head_num*HEAD_DIM*2: normalized_values_tmp
|
||||
# x*q_head_num*ROPE_DIM*(0.5) (not IS_PARTIAL_ROPE) x*q_head_num*ROPE_DIM*(0.5): y
|
||||
UB_SIZE = 87040 # 85K = 85 * 1024
|
||||
# the factor is the sum of elements number
|
||||
if IS_PARTIAL_ROPE:
|
||||
factor = 5 * q_hidden_size + 3 * kv_hidden_size + rope_dim * 4 + q_head_num * rope_dim
|
||||
batch_size_per_iter_per_vec = int(UB_SIZE / input.element_size()) // factor
|
||||
else:
|
||||
factor = 5 * q_hidden_size + 3 * kv_hidden_size + rope_dim * 2 + q_head_num * rope_dim // 2
|
||||
batch_size_per_iter_per_vec = int(UB_SIZE / input.element_size()) // factor
|
||||
batch_size_per_iter_per_vec = max(1, batch_size_per_iter_per_vec)
|
||||
qk_head_num_sum = int(q_head_num + kv_head_num)
|
||||
qk_head_nums_per_iter_per_vec = batch_size_per_iter_per_vec * qk_head_num_sum
|
||||
|
||||
grid = (num_vectorcore, 1, 1)
|
||||
# v tiling
|
||||
v_batch_size_per_iter_per_vec = UB_SIZE / torch.bfloat16.itemsize // (kv_hidden_size + 1)
|
||||
|
||||
split_qkv_rmsnorm_rope_kernel[grid](
|
||||
input,
|
||||
q_output,
|
||||
k_output,
|
||||
v_output,
|
||||
q_weight,
|
||||
q_bias,
|
||||
k_weight,
|
||||
k_bias,
|
||||
batch_size,
|
||||
q_hidden_size,
|
||||
kv_hidden_size,
|
||||
total_hidden_size,
|
||||
eps,
|
||||
BIAS,
|
||||
head_dim,
|
||||
rope_dim,
|
||||
rope_dim // 2,
|
||||
IS_PARTIAL_ROPE,
|
||||
num_vectorcore,
|
||||
int(batch_size_per_iter_per_vec),
|
||||
int(qk_head_nums_per_iter_per_vec),
|
||||
q_head_num,
|
||||
kv_head_num,
|
||||
qk_head_num_sum,
|
||||
int(v_batch_size_per_iter_per_vec),
|
||||
positions,
|
||||
cos_sin_cache,
|
||||
)
|
||||
return q_output, k_output, v_output
|
||||
|
||||
|
||||
def split_qkv_rmsnorm_rope_impl_fake(
|
||||
input: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
eps: float,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# Fake implementation for shape inference during Dynamo/AOT tracing.
|
||||
# Note: sin and cos are not used in shape computation, but must be present in signature.
|
||||
batch_size = input.shape[0]
|
||||
q_output = torch.empty(
|
||||
batch_size,
|
||||
int(q_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
k_output = torch.empty(
|
||||
batch_size,
|
||||
int(kv_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
v_output = torch.empty(
|
||||
batch_size,
|
||||
int(kv_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
return q_output, k_output, v_output
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="qkv_rmsnorm_rope",
|
||||
op_func=split_qkv_rmsnorm_rope_impl,
|
||||
fake_impl=split_qkv_rmsnorm_rope_impl_fake,
|
||||
mutates_args=[],
|
||||
dispatch_key="PrivateUse1",
|
||||
)
|
||||
386
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_rope_simt.py
Normal file
386
vllm_ascend/ops/triton/linearnorm/split_qkv_rmsnorm_rope_simt.py
Normal file
@@ -0,0 +1,386 @@
|
||||
import torch
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import extract_slice, get_vectorcore_num, insert_slice
|
||||
|
||||
|
||||
@triton.jit
|
||||
def precompute_rope_cos_sin_kernel(
|
||||
positions_gm_ptr,
|
||||
cos_sin_cache_gm_ptr,
|
||||
out_cos_sin_gm_ptr,
|
||||
batch_size,
|
||||
ROPE_DIM: tl.constexpr,
|
||||
num_vectorcore: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr = 128,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
batch_per_prog = tl.cdiv(batch_size, num_vectorcore)
|
||||
start = pid * batch_per_prog
|
||||
end = tl.minimum(start + batch_per_prog, batch_size)
|
||||
|
||||
sin_cos_range = tl.arange(0, ROPE_DIM)
|
||||
|
||||
for off in range(start, end, BLOCK_SIZE):
|
||||
batch_off = off + tl.arange(0, BLOCK_SIZE)
|
||||
mask = batch_off < end
|
||||
|
||||
pos = tl.load(positions_gm_ptr + batch_off, mask=mask)
|
||||
|
||||
offset = pos[:, None] * ROPE_DIM + sin_cos_range[None, :]
|
||||
sin_cos_val = tl.load(cos_sin_cache_gm_ptr + offset).to(tl.float32)
|
||||
|
||||
out_offset = batch_off[:, None] * ROPE_DIM + sin_cos_range[None, :]
|
||||
out_mask = (batch_off[:, None] < end) & (sin_cos_range[None, :] < ROPE_DIM)
|
||||
tl.store(out_cos_sin_gm_ptr + out_offset, sin_cos_val, mask=out_mask)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def split_qkv_rmsnorm_rope_simt_kernel(
|
||||
input_gm_ptr,
|
||||
q_gm_ptr,
|
||||
k_gm_ptr,
|
||||
v_gm_ptr,
|
||||
q_weight_ptr,
|
||||
q_bias_ptr,
|
||||
k_weight_ptr,
|
||||
k_bias_ptr,
|
||||
cos_sin_precomputed_ptr,
|
||||
batch_size,
|
||||
q_hidden_size: tl.constexpr,
|
||||
kv_hidden_size: tl.constexpr,
|
||||
total_hidden_size: tl.constexpr,
|
||||
eps: tl.constexpr,
|
||||
BIAS: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
ROPE_DIM: tl.constexpr,
|
||||
HALF_ROPE_DIM: tl.constexpr,
|
||||
IS_PARTIAL_ROPE: tl.constexpr,
|
||||
num_vectorcore: tl.constexpr,
|
||||
batch_size_per_iter_per_vec: tl.constexpr,
|
||||
v_batch_size_per_iter_per_vec: tl.constexpr,
|
||||
qk_head_nums_per_iter_per_vec: tl.constexpr,
|
||||
q_head_num: tl.constexpr,
|
||||
kv_head_num: tl.constexpr,
|
||||
qk_head_num_sum: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
|
||||
batch_per_prog = tl.cdiv(batch_size, num_vectorcore)
|
||||
start = pid * batch_per_prog
|
||||
end = tl.minimum(start + batch_per_prog, batch_size)
|
||||
|
||||
q_weight_values = tl.load(q_weight_ptr + tl.arange(0, HEAD_DIM)).to(tl.float32)
|
||||
k_weight_values = tl.load(k_weight_ptr + tl.arange(0, HEAD_DIM)).to(tl.float32)
|
||||
|
||||
output_q_nblk_idx = tl.arange(0, q_hidden_size)
|
||||
output_q_nmask = output_q_nblk_idx < q_hidden_size
|
||||
output_kv_nblk_idx = tl.arange(0, kv_hidden_size)
|
||||
output_kv_nmask = output_kv_nblk_idx < kv_hidden_size
|
||||
|
||||
for iter in tl.range(tl.cdiv(end - start, batch_size_per_iter_per_vec)):
|
||||
base_batch = start + iter * batch_size_per_iter_per_vec
|
||||
batch_indices = base_batch + tl.arange(0, batch_size_per_iter_per_vec)
|
||||
mmask = batch_indices < end
|
||||
|
||||
qk_cols = tl.arange(0, q_hidden_size + kv_hidden_size)
|
||||
mask = mmask[:, None] & (qk_cols[None, :] < total_hidden_size)
|
||||
|
||||
idx = batch_indices[:, None] * total_hidden_size + qk_cols[None, :]
|
||||
values_tmp1 = (
|
||||
tl.load(input_gm_ptr + idx, mask=mask).reshape(qk_head_nums_per_iter_per_vec, HEAD_DIM).to(tl.float32)
|
||||
)
|
||||
|
||||
if BIAS:
|
||||
q_bias_values = tl.load(q_bias_ptr + tl.arange(0, HEAD_DIM)).to(tl.float32)
|
||||
k_bias_values = tl.load(k_bias_ptr + tl.arange(0, HEAD_DIM)).to(tl.float32)
|
||||
|
||||
cos_sin_offset = base_batch * ROPE_DIM + tl.arange(0, batch_size_per_iter_per_vec * ROPE_DIM)
|
||||
cos_sin_value = tl.load(
|
||||
cos_sin_precomputed_ptr + cos_sin_offset, mask=cos_sin_offset < (end * ROPE_DIM)
|
||||
).reshape(batch_size_per_iter_per_vec, 1, ROPE_DIM)
|
||||
|
||||
cos = extract_slice(
|
||||
cos_sin_value, offsets=(0, 0, 0), sizes=(batch_size_per_iter_per_vec, 1, HALF_ROPE_DIM), strides=(1, 1, 1)
|
||||
)
|
||||
sin = extract_slice(
|
||||
cos_sin_value,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, 1, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
normalized_values = values_tmp1
|
||||
normalized_values = normalized_values * normalized_values
|
||||
normalized_values = tl.sum(normalized_values, axis=1) / HEAD_DIM
|
||||
normalized_values = 1 / tl.sqrt(normalized_values + eps).reshape(qk_head_nums_per_iter_per_vec, 1)
|
||||
normalized_values = values_tmp1 * normalized_values
|
||||
|
||||
normalized_values_tmp = extract_slice(
|
||||
normalized_values.reshape(batch_size_per_iter_per_vec, qk_head_num_sum, HEAD_DIM),
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HEAD_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
if BIAS:
|
||||
normalized_values_tmp = normalized_values_tmp * q_weight_values + q_bias_values
|
||||
else:
|
||||
normalized_values_tmp = normalized_values_tmp * q_weight_values
|
||||
|
||||
values_tmp = tl.zeros((batch_size_per_iter_per_vec, q_head_num, ROPE_DIM), dtype=tl.float32)
|
||||
x1 = extract_slice(
|
||||
normalized_values_tmp,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
x2 = extract_slice(
|
||||
normalized_values_tmp,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp = insert_slice(
|
||||
values_tmp,
|
||||
x1 * cos - x2 * sin,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp = insert_slice(
|
||||
values_tmp,
|
||||
x2 * cos + x1 * sin,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
q_output_idx = output_q_nblk_idx[None, :] + batch_indices[:, None] * q_hidden_size
|
||||
out_mask = mmask[:, None] & output_q_nmask[None, :]
|
||||
if IS_PARTIAL_ROPE:
|
||||
normalized_values_tmp = insert_slice(
|
||||
normalized_values_tmp,
|
||||
values_tmp,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, q_head_num, ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(
|
||||
q_gm_ptr + q_output_idx,
|
||||
normalized_values_tmp.reshape(batch_size_per_iter_per_vec, q_hidden_size),
|
||||
mask=out_mask,
|
||||
)
|
||||
else:
|
||||
tl.store(
|
||||
q_gm_ptr + q_output_idx, values_tmp.reshape(batch_size_per_iter_per_vec, q_hidden_size), mask=out_mask
|
||||
)
|
||||
|
||||
normalized_values_tmp1 = extract_slice(
|
||||
normalized_values.reshape(batch_size_per_iter_per_vec, qk_head_num_sum, HEAD_DIM),
|
||||
offsets=(0, q_head_num, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HEAD_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
if BIAS:
|
||||
normalized_values_tmp1 = normalized_values_tmp1 * k_weight_values + k_bias_values
|
||||
else:
|
||||
normalized_values_tmp1 = normalized_values_tmp1 * k_weight_values
|
||||
|
||||
values_tmp2 = tl.zeros((batch_size_per_iter_per_vec, kv_head_num, ROPE_DIM), dtype=tl.float32)
|
||||
x1 = extract_slice(
|
||||
normalized_values_tmp1,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
x2 = extract_slice(
|
||||
normalized_values_tmp1,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp2 = insert_slice(
|
||||
values_tmp2,
|
||||
x1 * cos - x2 * sin,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
values_tmp2 = insert_slice(
|
||||
values_tmp2,
|
||||
x2 * cos + x1 * sin,
|
||||
offsets=(0, 0, HALF_ROPE_DIM),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, HALF_ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
|
||||
kv_output_idx = output_kv_nblk_idx[None, :] + batch_indices[:, None] * kv_hidden_size
|
||||
out_mask = mmask[:, None] & output_kv_nmask[None, :]
|
||||
if IS_PARTIAL_ROPE:
|
||||
normalized_values_tmp1 = insert_slice(
|
||||
normalized_values_tmp1,
|
||||
values_tmp2,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(batch_size_per_iter_per_vec, kv_head_num, ROPE_DIM),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(
|
||||
k_gm_ptr + kv_output_idx,
|
||||
normalized_values_tmp1.reshape(batch_size_per_iter_per_vec, kv_hidden_size),
|
||||
mask=out_mask,
|
||||
)
|
||||
else:
|
||||
tl.store(
|
||||
k_gm_ptr + kv_output_idx,
|
||||
values_tmp2.reshape(batch_size_per_iter_per_vec, kv_hidden_size),
|
||||
mask=out_mask,
|
||||
)
|
||||
|
||||
for iter in tl.range(tl.cdiv(end - start, v_batch_size_per_iter_per_vec)):
|
||||
base_batch = start + iter * v_batch_size_per_iter_per_vec
|
||||
batch_indices = base_batch + tl.arange(0, v_batch_size_per_iter_per_vec)
|
||||
mmask = batch_indices < end
|
||||
|
||||
v_cols = tl.arange(q_hidden_size + kv_hidden_size, total_hidden_size)
|
||||
nmask = v_cols < total_hidden_size
|
||||
mask = mmask[:, None] & nmask[None, :]
|
||||
|
||||
idx = batch_indices[:, None] * total_hidden_size + v_cols[None, :]
|
||||
values = tl.load(input_gm_ptr + idx, mask=mask)
|
||||
|
||||
out_nblk_idx = tl.arange(0, kv_hidden_size)
|
||||
out_nmask = out_nblk_idx < kv_hidden_size
|
||||
out_idx = batch_indices[:, None] * kv_hidden_size + out_nblk_idx[None, :]
|
||||
out_mask = mmask[:, None] & out_nmask[None, :]
|
||||
tl.store(v_gm_ptr + out_idx, values, mask=out_mask)
|
||||
|
||||
|
||||
def split_qkv_rmsnorm_rope_simt_impl(
|
||||
input: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
eps: float,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
num_vectorcore = get_vectorcore_num()
|
||||
rope_dim = cos_sin_cache.shape[-1]
|
||||
batch_size = input.shape[0]
|
||||
BIAS = q_bias is not None
|
||||
IS_PARTIAL_ROPE = rope_dim != head_dim
|
||||
total_hidden_size = q_hidden_size + kv_hidden_size * 2
|
||||
|
||||
q_output = torch.empty(batch_size, q_hidden_size, device=input.device, dtype=input.dtype)
|
||||
k_output = torch.empty(batch_size, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
v_output = torch.empty(batch_size, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
|
||||
q_head_num = q_hidden_size // head_dim
|
||||
kv_head_num = kv_hidden_size // head_dim
|
||||
UB_SIZE = 87040
|
||||
|
||||
if IS_PARTIAL_ROPE:
|
||||
factor = 5 * q_hidden_size + 3 * kv_hidden_size + rope_dim * 4 + q_head_num * rope_dim
|
||||
batch_size_per_iter_per_vec = int(UB_SIZE / input.element_size()) // factor
|
||||
else:
|
||||
factor = 5 * q_hidden_size + 3 * kv_hidden_size + rope_dim * 2 + q_head_num * rope_dim // 2
|
||||
batch_size_per_iter_per_vec = int(UB_SIZE / input.element_size()) // factor
|
||||
batch_size_per_iter_per_vec = max(1, batch_size_per_iter_per_vec)
|
||||
qk_head_num_sum = int(q_head_num + kv_head_num)
|
||||
qk_head_nums_per_iter_per_vec = batch_size_per_iter_per_vec * qk_head_num_sum
|
||||
|
||||
v_batch_size_per_iter_per_vec = int(UB_SIZE / torch.float32.itemsize // (kv_hidden_size + 1))
|
||||
|
||||
cos_sin_precomputed = torch.empty(batch_size, rope_dim, dtype=torch.float32, device=input.device)
|
||||
grid = (num_vectorcore, 1, 1)
|
||||
precompute_rope_cos_sin_kernel[grid](
|
||||
positions,
|
||||
cos_sin_cache,
|
||||
cos_sin_precomputed,
|
||||
batch_size,
|
||||
rope_dim,
|
||||
num_vectorcore,
|
||||
force_simt_only=True,
|
||||
)
|
||||
|
||||
grid = (num_vectorcore, 1, 1)
|
||||
split_qkv_rmsnorm_rope_simt_kernel[grid](
|
||||
input,
|
||||
q_output,
|
||||
k_output,
|
||||
v_output,
|
||||
q_weight,
|
||||
q_bias,
|
||||
k_weight,
|
||||
k_bias,
|
||||
cos_sin_precomputed,
|
||||
batch_size,
|
||||
q_hidden_size,
|
||||
kv_hidden_size,
|
||||
total_hidden_size,
|
||||
eps,
|
||||
BIAS,
|
||||
head_dim,
|
||||
rope_dim,
|
||||
rope_dim // 2,
|
||||
IS_PARTIAL_ROPE,
|
||||
num_vectorcore,
|
||||
batch_size_per_iter_per_vec,
|
||||
v_batch_size_per_iter_per_vec,
|
||||
qk_head_nums_per_iter_per_vec,
|
||||
q_head_num,
|
||||
kv_head_num,
|
||||
q_head_num + kv_head_num,
|
||||
)
|
||||
return q_output, k_output, v_output
|
||||
|
||||
|
||||
def split_qkv_rmsnorm_rope_simt_impl_fake(
|
||||
input: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
eps: float,
|
||||
q_bias: torch.Tensor | None = None,
|
||||
k_bias: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# Fake implementation for shape inference during Dynamo/AOT tracing.
|
||||
# Note: sin and cos are not used in shape computation, but must be present in signature.
|
||||
batch_size = input.shape[0]
|
||||
q_output = torch.empty(
|
||||
batch_size,
|
||||
int(q_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
k_output = torch.empty(
|
||||
batch_size,
|
||||
int(kv_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
v_output = torch.empty(
|
||||
batch_size,
|
||||
int(kv_hidden_size),
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
return q_output, k_output, v_output
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="qkv_rmsnorm_rope_simt",
|
||||
op_func=split_qkv_rmsnorm_rope_simt_impl,
|
||||
fake_impl=split_qkv_rmsnorm_rope_simt_impl_fake,
|
||||
mutates_args=[],
|
||||
dispatch_key="PrivateUse1",
|
||||
)
|
||||
376
vllm_ascend/ops/triton/linearnorm/split_qkv_tp_rmsnorm_rope.py
Normal file
376
vllm_ascend/ops/triton/linearnorm/split_qkv_tp_rmsnorm_rope.py
Normal file
@@ -0,0 +1,376 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from vllm.distributed.communication_op import tensor_model_parallel_all_reduce
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import extract_slice, get_vectorcore_num, insert_slice
|
||||
|
||||
|
||||
# TODO: UB size differs across chips; consider whether BLOCK_SIZE can
|
||||
# be dynamically computed with a formula instead of autotuning {1,2,4}.
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({"BLOCK_SIZE": 1}),
|
||||
triton.Config({"BLOCK_SIZE": 2}),
|
||||
triton.Config({"BLOCK_SIZE": 4}),
|
||||
],
|
||||
key=["q_cols", "k_cols"],
|
||||
)
|
||||
@triton.jit
|
||||
def _split_qkv_and_compute_local_qk_var_kernel(
|
||||
input_ptr,
|
||||
q_out_ptr,
|
||||
k_out_ptr,
|
||||
v_out_ptr,
|
||||
qk_var_ptr,
|
||||
num_tokens,
|
||||
q_cols: tl.constexpr,
|
||||
k_cols: tl.constexpr,
|
||||
q_cols_pow2: tl.constexpr,
|
||||
k_cols_pow2: tl.constexpr,
|
||||
qkv_stride: tl.constexpr,
|
||||
q_inv_size: tl.constexpr,
|
||||
k_inv_size: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Grid Stride Loop + batch loading + precomputed reciprocal.
|
||||
(BLOCK_SIZE is limited to 1-4 to prevent UB overflow for large hidden_size)
|
||||
"""
|
||||
pid = tl.program_id(0)
|
||||
num_pids = tl.num_programs(0)
|
||||
block_range = tl.arange(0, BLOCK_SIZE)
|
||||
|
||||
# Grid Stride Loop: each program processes BLOCK_SIZE tokens at a time
|
||||
stride = num_pids * BLOCK_SIZE
|
||||
start_token_idx = pid * BLOCK_SIZE
|
||||
|
||||
for block_start in tl.range(start_token_idx, num_tokens, stride):
|
||||
token_indices = block_start + block_range
|
||||
token_mask = (token_indices < num_tokens)[:, None]
|
||||
|
||||
# === Batch load QKV data ===
|
||||
# Q: [BLOCK_SIZE, q_cols]
|
||||
q_offset = tl.arange(0, q_cols_pow2)[None, :]
|
||||
q_mask = token_mask & (q_offset < q_cols)
|
||||
q_batch = tl.load(
|
||||
input_ptr + token_indices[:, None] * qkv_stride + q_offset,
|
||||
mask=q_mask,
|
||||
other=0.0,
|
||||
)
|
||||
q_batch_f32 = q_batch.to(tl.float32)
|
||||
|
||||
# K: [BLOCK_SIZE, k_cols], K follows immediately after Q
|
||||
k_offset = tl.arange(0, k_cols_pow2)[None, :]
|
||||
k_mask = token_mask & (k_offset < k_cols)
|
||||
k_batch = tl.load(
|
||||
input_ptr + token_indices[:, None] * qkv_stride + q_cols + k_offset,
|
||||
mask=k_mask,
|
||||
other=0.0,
|
||||
)
|
||||
k_batch_f32 = k_batch.to(tl.float32)
|
||||
|
||||
# V: [BLOCK_SIZE, k_cols], V is at offset Q + 2*K
|
||||
v_offset = tl.arange(0, k_cols_pow2)[None, :]
|
||||
v_mask = token_mask & (v_offset < k_cols)
|
||||
v_batch = tl.load(
|
||||
input_ptr + token_indices[:, None] * qkv_stride + q_cols + k_cols + v_offset,
|
||||
mask=v_mask,
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
# === Batch compute sum of squares ===
|
||||
q_squaresum = tl.sum(q_batch_f32 * q_batch_f32, axis=-1) * q_inv_size
|
||||
k_squaresum = tl.sum(k_batch_f32 * k_batch_f32, axis=-1) * k_inv_size
|
||||
|
||||
# === Batch store QKV output ===
|
||||
# Store Q
|
||||
q_out_offset = token_indices[:, None] * q_cols + q_offset
|
||||
q_out_mask = token_mask & (q_offset < q_cols)
|
||||
tl.store(q_out_ptr + q_out_offset, q_batch, mask=q_out_mask)
|
||||
|
||||
# Store K
|
||||
k_out_offset = token_indices[:, None] * k_cols + k_offset
|
||||
k_out_mask = token_mask & (k_offset < k_cols)
|
||||
tl.store(k_out_ptr + k_out_offset, k_batch, mask=k_out_mask)
|
||||
|
||||
# Store V
|
||||
v_out_offset = token_indices[:, None] * k_cols + v_offset
|
||||
v_out_mask = token_mask & (v_offset < k_cols)
|
||||
tl.store(v_out_ptr + v_out_offset, v_batch, mask=v_out_mask)
|
||||
|
||||
# === Store variance ===
|
||||
var_offset = token_indices * 2
|
||||
var_mask = token_indices < num_tokens
|
||||
tl.store(qk_var_ptr + var_offset, q_squaresum, mask=var_mask)
|
||||
tl.store(qk_var_ptr + var_offset + 1, k_squaresum, mask=var_mask)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _apply_global_rmsnorm_kernel(
|
||||
q_ptr,
|
||||
k_ptr,
|
||||
cos_ptr,
|
||||
sin_ptr,
|
||||
cs_row_stride,
|
||||
q_weight_ptr,
|
||||
k_weight_ptr,
|
||||
qk_global_var_ptr,
|
||||
eps: tl.constexpr,
|
||||
inv_tp_world: tl.constexpr,
|
||||
num_tokens,
|
||||
q_cols: tl.constexpr,
|
||||
k_cols: tl.constexpr,
|
||||
q_num_heads: tl.constexpr,
|
||||
k_num_heads: tl.constexpr,
|
||||
head_dim: tl.constexpr,
|
||||
rotary_dim: tl.constexpr,
|
||||
HALF: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0).to(tl.int64)
|
||||
num_programs = tl.num_programs(0)
|
||||
tokens_per_program = tl.cdiv(num_tokens, num_programs)
|
||||
iter_num_per_program = tokens_per_program
|
||||
program_token_offset = pid * tokens_per_program
|
||||
program_token_end = min(program_token_offset + tokens_per_program, num_tokens)
|
||||
|
||||
token_tile_offsets = tl.arange(0, 1)
|
||||
q_head_offsets = tl.arange(0, q_num_heads)[:, None]
|
||||
k_head_offsets = tl.arange(0, k_num_heads)[:, None]
|
||||
hd_offsets = tl.arange(0, head_dim)[None, :]
|
||||
|
||||
q_row_offsets = q_head_offsets * head_dim + hd_offsets
|
||||
k_row_offsets = k_head_offsets * head_dim + hd_offsets
|
||||
|
||||
q_weight = tl.load(q_weight_ptr + q_row_offsets).to(tl.float32)
|
||||
k_weight = tl.load(k_weight_ptr + k_row_offsets).to(tl.float32)
|
||||
|
||||
half_offsets = tl.arange(0, HALF)
|
||||
base_token_offsets = program_token_offset + token_tile_offsets
|
||||
|
||||
for iter in tl.range(iter_num_per_program):
|
||||
token_offsets = base_token_offsets + iter
|
||||
token_mask = token_offsets < program_token_end
|
||||
|
||||
q_gv = tl.load(qk_global_var_ptr + token_offsets * 2, mask=token_mask, other=0.0).to(tl.float32)
|
||||
q_gv = q_gv * inv_tp_world
|
||||
k_gv = tl.load(qk_global_var_ptr + token_offsets * 2 + 1, mask=token_mask, other=0.0).to(tl.float32)
|
||||
k_gv = k_gv * inv_tp_world
|
||||
q_scale = 1.0 / tl.sqrt(q_gv + eps)
|
||||
k_scale = 1.0 / tl.sqrt(k_gv + eps)
|
||||
|
||||
q_offsets = token_offsets[:, None, None] * q_cols + q_row_offsets[None, :, :]
|
||||
q_mask = token_mask[:, None, None]
|
||||
q_vals_raw = tl.load(q_ptr + q_offsets, mask=q_mask, other=0.0)
|
||||
q_vals = q_vals_raw.to(tl.float32) * q_scale[:, None, None] * q_weight[None, :, :]
|
||||
|
||||
k_offsets = token_offsets[:, None, None] * k_cols + k_row_offsets[None, :, :]
|
||||
k_mask = token_mask[:, None, None]
|
||||
k_vals_raw = tl.load(k_ptr + k_offsets, mask=k_mask, other=0.0)
|
||||
k_vals = k_vals_raw.to(tl.float32) * k_scale[:, None, None] * k_weight[None, :, :]
|
||||
|
||||
# Neox-style RoPE on the first rotary_dim dimensions of each head
|
||||
cs_offsets = token_offsets[:, None] * cs_row_stride + half_offsets[None, :]
|
||||
cs_mask = token_mask[:, None]
|
||||
cos_row = tl.load(cos_ptr + cs_offsets, mask=cs_mask, other=0.0).to(tl.float32)
|
||||
sin_row = tl.load(sin_ptr + cs_offsets, mask=cs_mask, other=0.0).to(tl.float32)
|
||||
|
||||
q1 = extract_slice(
|
||||
q_vals,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(1, q_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
q2 = extract_slice(
|
||||
q_vals,
|
||||
offsets=(0, 0, HALF),
|
||||
sizes=(1, q_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
q_vals = insert_slice(
|
||||
q_vals,
|
||||
q1 * cos_row[:, None, :] - q2 * sin_row[:, None, :],
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(1, q_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
q_vals = insert_slice(
|
||||
q_vals,
|
||||
q2 * cos_row[:, None, :] + q1 * sin_row[:, None, :],
|
||||
offsets=(0, 0, HALF),
|
||||
sizes=(1, q_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(q_ptr + q_offsets, q_vals.to(q_vals_raw.dtype), mask=q_mask)
|
||||
|
||||
k1 = extract_slice(
|
||||
k_vals,
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(1, k_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
k2 = extract_slice(
|
||||
k_vals,
|
||||
offsets=(0, 0, HALF),
|
||||
sizes=(1, k_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
k_vals = insert_slice(
|
||||
k_vals,
|
||||
k1 * cos_row[:, None, :] - k2 * sin_row[:, None, :],
|
||||
offsets=(0, 0, 0),
|
||||
sizes=(1, k_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
k_vals = insert_slice(
|
||||
k_vals,
|
||||
k2 * cos_row[:, None, :] + k1 * sin_row[:, None, :],
|
||||
offsets=(0, 0, HALF),
|
||||
sizes=(1, k_num_heads, HALF),
|
||||
strides=(1, 1, 1),
|
||||
)
|
||||
tl.store(k_ptr + k_offsets, k_vals.to(k_vals_raw.dtype), mask=k_mask)
|
||||
|
||||
|
||||
def split_qkv_tp_rmsnorm_rope_impl(
|
||||
input: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
rotary_dim: int,
|
||||
eps: float,
|
||||
tp_world: int,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
num_tokens = input.shape[0]
|
||||
input_2d = input.view(num_tokens, -1)
|
||||
q = torch.empty(num_tokens, q_hidden_size, device=input.device, dtype=input.dtype)
|
||||
k = torch.empty(num_tokens, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
v = torch.empty(num_tokens, kv_hidden_size, device=input.device, dtype=input.dtype)
|
||||
if num_tokens == 0:
|
||||
return q, k, v
|
||||
|
||||
num_vectorcore = get_vectorcore_num()
|
||||
grid = (min(num_tokens, num_vectorcore),)
|
||||
q_cols = q_hidden_size
|
||||
k_cols = kv_hidden_size
|
||||
q_num_heads = q_hidden_size // head_dim
|
||||
k_num_heads = kv_hidden_size // head_dim
|
||||
|
||||
qk_var = torch.empty(num_tokens, 2, dtype=torch.float32, device=q.device)
|
||||
# Precompute reciprocal to avoid division inside kernel
|
||||
q_inv_size = 1.0 / q_cols
|
||||
k_inv_size = 1.0 / k_cols
|
||||
# Pad to power-of-2 for tl.arange (required by Ascend NPU Triton backend)
|
||||
q_cols_pow2 = 1 << (q_cols - 1).bit_length()
|
||||
k_cols_pow2 = 1 << (k_cols - 1).bit_length()
|
||||
_split_qkv_and_compute_local_qk_var_kernel[grid](
|
||||
input_2d,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
qk_var,
|
||||
num_tokens,
|
||||
q_cols,
|
||||
k_cols,
|
||||
q_cols_pow2,
|
||||
k_cols_pow2,
|
||||
q_cols + 2 * k_cols,
|
||||
q_inv_size,
|
||||
k_inv_size,
|
||||
)
|
||||
if tp_world > 1:
|
||||
qk_var = tensor_model_parallel_all_reduce(qk_var)
|
||||
|
||||
cos_2d = cos.view(num_tokens, -1)
|
||||
sin_2d = sin.view(num_tokens, -1)
|
||||
q_2d = q.view(num_tokens, -1)
|
||||
k_2d = k.view(num_tokens, -1)
|
||||
_apply_global_rmsnorm_kernel[grid](
|
||||
q_2d,
|
||||
k_2d,
|
||||
cos_2d,
|
||||
sin_2d,
|
||||
cos_2d.stride(0),
|
||||
q_weight,
|
||||
k_weight,
|
||||
qk_var,
|
||||
eps,
|
||||
1.0 / tp_world,
|
||||
num_tokens,
|
||||
q_cols,
|
||||
k_cols,
|
||||
q_num_heads,
|
||||
k_num_heads,
|
||||
head_dim,
|
||||
rotary_dim,
|
||||
rotary_dim // 2,
|
||||
)
|
||||
|
||||
return q, k, v
|
||||
|
||||
|
||||
def split_qkv_tp_rmsnorm_rope_impl_fake(
|
||||
input: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_hidden_size: int,
|
||||
kv_hidden_size: int,
|
||||
head_dim: int,
|
||||
rotary_dim: int,
|
||||
eps: float,
|
||||
tp_world: int,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
num_tokens = input.shape[0]
|
||||
q_out = torch.empty(
|
||||
num_tokens,
|
||||
q_hidden_size,
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
k_out = torch.empty(
|
||||
num_tokens,
|
||||
kv_hidden_size,
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
v_out = torch.empty(
|
||||
num_tokens,
|
||||
kv_hidden_size,
|
||||
device=input.device,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
return q_out, k_out, v_out
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="split_qkv_tp_rmsnorm_rope",
|
||||
op_func=split_qkv_tp_rmsnorm_rope_impl,
|
||||
fake_impl=split_qkv_tp_rmsnorm_rope_impl_fake,
|
||||
mutates_args=[],
|
||||
dispatch_key="PrivateUse1",
|
||||
)
|
||||
Reference in New Issue
Block a user