Files
paddlepaddle--paddle/test/legacy_test/test_fused_rotary_position_embedding.py
2026-07-13 12:40:42 +08:00

968 lines
31 KiB
Python

# Copyright (c) 2023 PaddlePaddle Authors. 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.
import unittest
import numpy as np
import parameterized as param
from op_test import is_custom_device
import paddle
from paddle.base import core
from paddle.incubate.nn.functional import fused_rotary_position_embedding
position_ids_list = [[7, 5, 4, 6, 3, 1, 2, 0], [3, 1, 4, 0, 7, 6, 5, 2]]
def deal_qkv(init_value):
if init_value is None:
return None
perm = [0, 2, 1, 3]
return paddle.transpose(x=init_value, perm=perm)
def mult_qkv(value, cos_tensor, sin_tensor):
if value is None:
return None
rot_dim = cos_tensor.shape[-1]
value, value_pass = value[..., :rot_dim], value[..., rot_dim:]
rotate_half_q = paddle.reshape(
paddle.stack([-value[:, :, :, 1::2], value[:, :, :, 0::2]], axis=-1),
paddle.shape(value),
)
query = paddle.add(
paddle.multiply(value, cos_tensor),
paddle.multiply(rotate_half_q, sin_tensor),
)
return paddle.cat([query, value_pass], axis=-1)
def mult_qkv_rotate_half(value, cos_tensor, sin_tensor):
if value is None:
return None
rot_dim = cos_tensor.shape[-1]
value, value_pass = value[..., :rot_dim], value[..., rot_dim:]
rotate_half_q = paddle.reshape(
paddle.concat(
[
-value[..., value.shape[-1] // 2 :],
value[..., : value.shape[-1] // 2],
],
axis=-1,
),
paddle.shape(value),
)
query = paddle.add(
paddle.multiply(value, cos_tensor),
paddle.multiply(rotate_half_q, sin_tensor),
)
return paddle.cat([query, value_pass], axis=-1)
def get_sin_cos_tensor(seq_len, head_dim, sign=1, rotate_half=False):
pos_seq = paddle.arange(0, seq_len, 1, dtype="float32")
indices = paddle.arange(0, head_dim, 2, dtype="float32")
indices = 1 / 10000 ** (indices / head_dim)
sinusoid_inp = pos_seq.unsqueeze(1) * indices.unsqueeze(0)
sin_sin = np.empty((seq_len * head_dim), dtype=np.float32)
cos_cos = np.empty((seq_len * head_dim), dtype=np.float32)
numpy_array = sinusoid_inp.numpy()
iter_array = np.nditer(numpy_array)
i = 0
if rotate_half:
stride = head_dim // 2
for value in iter_array:
sin_sin[i] = sign * np.sin(value)
cos_cos[i] = np.cos(value)
sin_sin[i + stride] = np.sin(
value * 0.1
) # Verify the accuracy of the reverse computation logic for rotate_half by setting the front and back sin values inconsistently.
cos_cos[i + stride] = np.cos(value)
i += 1
if i % head_dim == stride:
i += stride
else:
for value in iter_array:
sin_sin[i * 2] = sign * np.sin(value)
cos_cos[i * 2 + 0] = np.cos(value)
sin_sin[i * 2 + 1] = np.sin(value)
cos_cos[i * 2 + 1] = np.cos(value)
i += 1
tensor_sin = paddle.reshape(
paddle.to_tensor(sin_sin),
[1, seq_len, 1, head_dim],
)
tensor_cos = paddle.reshape(
paddle.to_tensor(cos_cos),
[1, seq_len, 1, head_dim],
)
return tensor_sin, tensor_cos
def paddle_fused_rotary_position_embedding(
init_q,
init_k,
init_v,
sin_tensor=None,
cos_tensor=None,
position_ids=None,
use_neox_rotary_style=True,
**kwargs,
):
# permute q, k, v from [batch_size, seq_len, num_heads, head_dim]
# to [batch_size, num_heads, seq_len, head_dim]
q = deal_qkv(init_q)
k = deal_qkv(init_k)
v = deal_qkv(init_v)
if position_ids is not None:
sin_tensor = sin_tensor.squeeze(axis=[0, 2]) # [seq_len, dim]
cos_tensor = cos_tensor.squeeze(axis=[0, 2]) # [seq_len, dim]
sin_tensor = sin_tensor[position_ids].unsqueeze(
2
) # [bs, seq_len, 1, dim]
cos_tensor = cos_tensor[position_ids].unsqueeze(
2
) # [bs, seq_len, 1, dim]
perm = [0, 2, 1, 3]
sin_tensor = paddle.transpose(x=sin_tensor, perm=perm)
cos_tensor = paddle.transpose(x=cos_tensor, perm=perm)
if use_neox_rotary_style:
query = mult_qkv(q, cos_tensor, sin_tensor)
value = mult_qkv(v, cos_tensor, sin_tensor)
key = mult_qkv(k, cos_tensor, sin_tensor)
else:
query = mult_qkv_rotate_half(q, cos_tensor, sin_tensor)
value = mult_qkv_rotate_half(v, cos_tensor, sin_tensor)
key = mult_qkv_rotate_half(k, cos_tensor, sin_tensor)
# permute the result back to [batch_size, seq_len, num_heads, head_dim]
r_query = deal_qkv(query)
r_key = deal_qkv(key)
r_value = deal_qkv(value)
return r_query, r_key, r_value
@unittest.skipIf(
not (core.is_compiled_with_cuda() or is_custom_device())
and not paddle.is_compiled_with_rocm(),
"core is not compiled with CUDA or ROCM ",
)
@param.parameterized_class(
("name", "shape_q", "shape_k", "shape_v", "position_ids_list"),
[
(
"qkv_input",
[2, 8, 2, 16], # bs, seq_len, num_heads, head_dim
[2, 8, 2, 16], # bs, seq_len, num_heads, head_dim
[2, 8, 2, 16], # bs, seq_len, num_heads, head_dim
position_ids_list,
),
("qk_input", [2, 8, 2, 16], [2, 8, 2, 16], None, position_ids_list),
("qv_input", [2, 8, 2, 16], None, [2, 8, 2, 16], position_ids_list),
("q_input", [2, 8, 2, 16], None, None, position_ids_list),
(
"qkv_input_mqa",
[2, 8, 4, 8],
[2, 8, 1, 8],
[2, 8, 1, 8],
position_ids_list,
),
("qk_input_mqa", [2, 8, 4, 8], [2, 8, 1, 8], None, position_ids_list),
("qv_input_mqa", [2, 8, 4, 8], None, [2, 8, 1, 8], position_ids_list),
(
"qkv_input_gqa",
[1, 8, 4, 8],
[1, 8, 2, 8],
[1, 8, 2, 8],
position_ids_list[:1],
),
(
"qk_input_gqa",
[1, 8, 4, 8],
[1, 8, 2, 8],
None,
position_ids_list[:1],
),
(
"qv_input_gqa",
[1, 8, 4, 8],
None,
[1, 8, 2, 8],
position_ids_list[:1],
),
],
)
class TestFusedRotaryPositionEmbedding(unittest.TestCase):
def setUp(self):
self.dtype = "float32"
self.training = True
self.seed = 1203
self.rtol = 1e-5
self.atol = 1e-6
def get_paddle_tensor(self, shape):
if shape is None:
return None
tmp = paddle.randn(shape, self.dtype)
tmp.stop_gradient = False
return tmp
def get_inputs(
self,
seed,
with_sin_cos,
rotary_percent=1.0,
with_grads=False,
rotate_half=False,
):
paddle.disable_static()
paddle.seed(seed)
# tensor_q shape: [batch_size, seq_len, num_heads, head_dim]
tensor_q = self.get_paddle_tensor(self.shape_q)
tensor_k = self.get_paddle_tensor(self.shape_k)
tensor_v = self.get_paddle_tensor(self.shape_v)
tensor_sin, tensor_cos = (
get_sin_cos_tensor(
tensor_q.shape[1],
int(tensor_q.shape[3] * rotary_percent),
1,
rotate_half=rotate_half,
)
if with_sin_cos
else (None, None)
)
if not with_grads:
return (tensor_q, tensor_k, tensor_v, tensor_sin, tensor_cos)
tensor_grad_outq = self.get_paddle_tensor(self.shape_q)
tensor_grad_outk = self.get_paddle_tensor(self.shape_k)
tensor_grad_outv = self.get_paddle_tensor(self.shape_v)
return (
tensor_q,
tensor_k,
tensor_v,
tensor_sin,
tensor_cos,
tensor_grad_outq,
tensor_grad_outk,
tensor_grad_outv,
)
def get_forward_backward(
self,
rope_function,
seed,
with_sin_cos=True,
rotary_percent=1.0,
use_neox_rotary_style=True,
position_ids=None,
test_time_major=False,
):
paddle.disable_static()
fw = []
bw = []
(
tensor_q,
tensor_k,
tensor_v,
tensor_sin,
tensor_cos,
tensor_grad_outq,
tensor_grad_outk,
tensor_grad_outv,
) = self.get_inputs(
seed,
with_sin_cos,
rotary_percent,
with_grads=True,
rotate_half=not use_neox_rotary_style,
)
if test_time_major:
# [batch_size, seq_len, num_heads, head_dim] -> [seq_len, batch_size, num_heads, head_dim]
if tensor_q is not None:
tensor_q = paddle.transpose(tensor_q, perm=[1, 0])
if tensor_k is not None:
tensor_k = paddle.transpose(tensor_k, perm=[1, 0])
if tensor_v is not None:
tensor_v = paddle.transpose(tensor_v, perm=[1, 0])
if tensor_grad_outq is not None:
tensor_grad_outq = paddle.transpose(
tensor_grad_outq, perm=[1, 0]
)
if tensor_grad_outk is not None:
tensor_grad_outk = paddle.transpose(
tensor_grad_outk, perm=[1, 0]
)
if tensor_grad_outv is not None:
tensor_grad_outv = paddle.transpose(
tensor_grad_outv, perm=[1, 0]
)
tensor_q = tensor_q.detach().clone()
tensor_q.stop_gradient = False
if tensor_k is not None:
tensor_k = tensor_k.detach().clone()
tensor_k.stop_gradient = False
if tensor_v is not None:
tensor_v = tensor_v.detach().clone()
tensor_v.stop_gradient = False
out_q, out_k, out_v = rope_function(
tensor_q,
tensor_k,
tensor_v,
tensor_sin,
tensor_cos,
position_ids=position_ids,
use_neox_rotary_style=use_neox_rotary_style,
time_major=test_time_major,
)
out_init_grad = []
for out_value in [out_q, out_k, out_v]:
if out_value is None or not out_value._is_initialized():
continue
fw.append(out_value)
for grad_value in [
tensor_grad_outq,
tensor_grad_outk,
tensor_grad_outv,
]:
if grad_value is None or not grad_value._is_initialized():
continue
out_init_grad.append(grad_value)
paddle.autograd.backward(fw, out_init_grad, True)
bw = list(
filter(lambda x: x is not None, [tensor_q, tensor_k, tensor_v])
)
bw = [x.grad for x in bw]
if test_time_major:
# transpose back
# [seq_len, batch_size, num_heads, head_dim] -> [batch_size, seq_len, num_heads, head_dim]
fw = [paddle.transpose(x, perm=[1, 0]) for x in fw]
bw = [paddle.transpose(x, perm=[1, 0]) for x in bw]
return fw, bw
def check_results(self, p_results, f_results):
for i in range(len(p_results)):
np.testing.assert_allclose(
p_results[i].numpy(),
f_results[i].numpy(),
rtol=self.rtol,
atol=self.atol,
err_msg=f"Tensor {i} not match",
)
def test_fused_rope(self):
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding, seed=self.seed
)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
test_time_major=False,
)
f_fw_time_major, f_bw_time_major = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
test_time_major=True,
)
self.check_results(p_fw, f_fw)
self.check_results(p_bw, f_bw)
self.check_results(p_fw, f_fw_time_major)
self.check_results(p_bw, f_bw_time_major)
def test_fused_rope_with_sin_cos(self):
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
test_time_major=False,
)
f_fw_time_major, f_bw_time_major = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
test_time_major=True,
)
self.check_results(p_fw, f_fw)
self.check_results(p_bw, f_bw)
self.check_results(p_fw, f_fw_time_major)
self.check_results(p_bw, f_bw_time_major)
def test_fused_rope_with_sin_cos_with_rotary_percent(self):
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
rotary_percent=0.5,
)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
rotary_percent=0.5,
test_time_major=False,
)
f_fw_time_major, f_bw_time_major = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
with_sin_cos=True,
rotary_percent=0.5,
test_time_major=True,
)
self.check_results(p_fw, f_fw)
self.check_results(p_bw, f_bw)
self.check_results(p_fw, f_fw_time_major)
self.check_results(p_bw, f_bw_time_major)
def test_fused_rope_rotate_half(self):
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
use_neox_rotary_style=False,
)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
use_neox_rotary_style=False,
test_time_major=False,
)
f_fw_time_major, f_bw_time_major = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
use_neox_rotary_style=False,
test_time_major=True,
)
self.check_results(p_fw, f_fw)
self.check_results(p_bw, f_bw)
self.check_results(p_fw, f_fw_time_major)
self.check_results(p_bw, f_bw_time_major)
def test_fused_rope_position_ids(self):
position_ids = paddle.to_tensor(self.position_ids_list)
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
position_ids=position_ids,
)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
position_ids=position_ids,
test_time_major=False,
)
f_fw_time_major, f_bw_time_major = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
position_ids=position_ids,
test_time_major=True,
)
self.check_results(p_fw, f_fw)
self.check_results(p_bw, f_bw)
self.check_results(p_fw, f_fw_time_major)
self.check_results(p_bw, f_bw_time_major)
def test_static(self):
paddle.disable_static()
tensor_q, tensor_k, tensor_v, tensor_sin, tensor_cos = self.get_inputs(
self.seed, True, rotate_half=True
)
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
use_neox_rotary_style=False,
)
paddle.enable_static()
main = paddle.static.Program()
startup = paddle.static.Program()
with paddle.static.program_guard(main, startup):
q = (
None
if self.shape_q is None
else paddle.static.data(
name="q", shape=self.shape_q, dtype=self.dtype
)
)
k = (
None
if self.shape_k is None
else paddle.static.data(
name="k", shape=self.shape_k, dtype=self.dtype
)
)
v = (
None
if self.shape_v is None
else paddle.static.data(
name="v", shape=self.shape_v, dtype=self.dtype
)
)
sin = paddle.static.data(
name="sin",
shape=(1, tensor_q.shape[1], 1, tensor_q.shape[3]),
dtype=self.dtype,
)
cos = paddle.static.data(
name="cos",
shape=(1, tensor_q.shape[1], 1, tensor_q.shape[3]),
dtype=self.dtype,
)
out_q, out_k, out_v = fused_rotary_position_embedding(
q,
k,
v,
sin,
cos,
position_ids=None,
use_neox_rotary_style=False,
)
exe = paddle.static.Executor()
feed = {
"sin": tensor_sin.numpy(),
"cos": tensor_cos.numpy(),
}
for var_name, input_tensor in zip(
["q", "k", "v"], [tensor_q, tensor_k, tensor_v]
):
if input_tensor is not None:
feed[var_name] = input_tensor.numpy()
fetch_list = []
for x, out in zip([q, k, v], [out_q, out_k, out_v]):
# The reason why fetch `out` based on `x` is that
# if input is None, the output of static function might be not NoneType
# but pir.Value with type builtin.tensor<0xf32> in pir mode.
if x is not None:
fetch_list.append(out)
outs = exe.run(
main,
feed=feed,
fetch_list=fetch_list,
)
for i in range(len(p_fw)):
np.testing.assert_allclose(
p_fw[i].numpy(), outs[i], rtol=self.rtol, atol=self.atol
)
paddle.disable_static()
def test_static_time_major(self):
paddle.disable_static()
tensor_q, tensor_k, tensor_v, tensor_sin, tensor_cos = self.get_inputs(
self.seed, True, rotate_half=True
)
p_fw, p_bw = self.get_forward_backward(
paddle_fused_rotary_position_embedding,
seed=self.seed,
use_neox_rotary_style=False,
test_time_major=False,
)
paddle.enable_static()
shape_q = (
[self.shape_q[1], self.shape_q[0], self.shape_q[2], self.shape_q[3]]
if self.shape_q
else None
)
shape_k = (
[self.shape_k[1], self.shape_k[0], self.shape_k[2], self.shape_k[3]]
if self.shape_k
else None
)
shape_v = (
[self.shape_v[1], self.shape_v[0], self.shape_v[2], self.shape_v[3]]
if self.shape_v
else None
)
main = paddle.static.Program()
startup = paddle.static.Program()
with paddle.static.program_guard(main, startup):
q = (
None
if shape_q is None
else paddle.static.data(
name="q", shape=shape_q, dtype=self.dtype
)
)
k = (
None
if shape_k is None
else paddle.static.data(
name="k", shape=shape_k, dtype=self.dtype
)
)
v = (
None
if shape_v is None
else paddle.static.data(
name="v", shape=shape_v, dtype=self.dtype
)
)
sin = paddle.static.data(
name="sin",
shape=(1, shape_q[0], 1, shape_q[3]),
dtype=self.dtype,
)
cos = paddle.static.data(
name="cos",
shape=(1, shape_q[0], 1, shape_q[3]),
dtype=self.dtype,
)
q.stop_gradient = False
if v is not None:
v.stop_gradient = False
if k is not None:
k.stop_gradient = False
out_q, out_k, out_v = fused_rotary_position_embedding(
q,
k,
v,
sin,
cos,
position_ids=None,
use_neox_rotary_style=False,
time_major=True,
)
dout = paddle.static.gradients(out_q, q)
exe = paddle.static.Executor()
feed = {
"sin": tensor_sin.numpy(),
"cos": tensor_cos.numpy(),
}
for var_name, input_tensor in zip(
["q", "k", "v"], [tensor_q, tensor_k, tensor_v]
):
if input_tensor is not None:
feed[var_name] = input_tensor.numpy().transpose((1, 0, 2, 3))
fetch_list = []
for x, out in zip([q, k, v], [out_q, out_k, out_v]):
# The reason why fetch `out` based on `x` is that
# if input is None, the output of static function might be not NoneType
# but pir.Value with type builtin.tensor<0xf32> in pir mode.
if x is not None:
fetch_list.append(out)
outs = exe.run(
main,
feed=feed,
fetch_list=fetch_list,
)
for i in range(len(p_fw)):
np.testing.assert_allclose(
p_fw[i].numpy(),
outs[i].transpose((1, 0, 2, 3)),
rtol=self.rtol,
atol=self.atol,
)
paddle.disable_static()
def test_errors(self):
def test_error1():
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
test_time_major=False,
with_sin_cos=False,
use_neox_rotary_style=False,
)
self.assertRaises(AssertionError, test_error1)
def test_error2():
position_ids = paddle.to_tensor(self.position_ids_list)
f_fw, f_bw = self.get_forward_backward(
fused_rotary_position_embedding,
seed=self.seed,
test_time_major=False,
with_sin_cos=False,
position_ids=position_ids,
)
self.assertRaises(AssertionError, test_error2)
@unittest.skipIf(
not (core.is_compiled_with_cuda() or is_custom_device())
and not paddle.is_compiled_with_rocm(),
"core is not compiled with CUDA or ROCM ",
)
class TestFusedRotaryPositionEmbeddingZeroSize(unittest.TestCase):
def setUp(self):
self.dtype = "float32"
self.qkv_shape = [0, 1, 8, 8]
self.sin_cos_shape = [1, 1, 1, 8]
def init_data(self):
self.q = paddle.randn(self.qkv_shape, dtype=self.dtype)
self.k = paddle.randn(self.qkv_shape, dtype=self.dtype)
self.v = paddle.randn(self.qkv_shape, dtype=self.dtype)
self.q.stop_gradient = False
self.k.stop_gradient = False
self.v.stop_gradient = False
self.sin = paddle.sin(
paddle.randn(self.sin_cos_shape, dtype=self.dtype)
)
self.cos = paddle.cos(
paddle.randn(self.sin_cos_shape, dtype=self.dtype)
)
def _test_forward_backward(self):
out_q, out_k, out_v = fused_rotary_position_embedding(
self.q,
self.k,
self.v,
sin=self.sin,
cos=self.cos,
use_neox_rotary_style=False,
)
out = out_q + out_k + out_v
out.backward()
np.testing.assert_allclose(
self.q.shape, self.q.grad.shape, rtol=1e-05, atol=1e-06
)
np.testing.assert_allclose(
self.k.shape, self.k.grad.shape, rtol=1e-05, atol=1e-06
)
np.testing.assert_allclose(
self.v.shape, self.v.grad.shape, rtol=1e-05, atol=1e-06
)
def test_zero_size(self):
self.init_data()
self._test_forward_backward()
@unittest.skipIf(
not (core.is_compiled_with_cuda() or is_custom_device())
and not paddle.is_compiled_with_rocm(),
"core is not compiled with CUDA or ROCM ",
)
class TestFusedRotaryPositionEmbeddingZeroNumHeads(unittest.TestCase):
"""Test fused_rotary_position_embedding with k or v tensors that have
zero num_heads (e.g. shape [batch, seq, 0, head_dim]).
Regression test for a bug where:
1. The MQA/GQA validation `num_heads % v_num_heads == 0` caused a
SIGFPE (integer division by zero) when v_num_heads == 0.
2. FusedRopeKernelLauncher launched CUDA kernels for zero-element
tensors even when numel == 0.
"""
def setUp(self):
self.dtype = "float32"
self.batch_size = 1
self.seq_len = 8
self.num_heads_q = 4
self.head_dim = 8
self.sin_cos_shape = [1, self.seq_len, 1, self.head_dim]
def _make_tensor(self, shape, requires_grad=True):
t = paddle.randn(shape, dtype=self.dtype)
t.stop_gradient = not requires_grad
return t
def _make_sin_cos(self):
sin = paddle.sin(paddle.randn(self.sin_cos_shape, dtype=self.dtype))
cos = paddle.cos(paddle.randn(self.sin_cos_shape, dtype=self.dtype))
return sin, cos
def _run_forward_backward(self, q, k, v, sin, cos, **kwargs):
"""Run forward + backward; return outputs and check no crash."""
out_q, out_k, out_v = fused_rotary_position_embedding(
q, k, v, sin=sin, cos=cos, **kwargs
)
# Build loss from initialized, non-empty outputs
loss_terms = []
for out in [out_q, out_k, out_v]:
if out is not None and out._is_initialized() and out.numel() > 0:
loss_terms.append(out.sum())
if loss_terms:
sum(loss_terms).backward()
return out_q, out_k, out_v
def test_v_zero_num_heads(self):
"""v with 0 num_heads should not crash (original bug scenario)."""
q_shape = [
self.batch_size,
self.seq_len,
self.num_heads_q,
self.head_dim,
]
kv_shape = [self.batch_size, self.seq_len, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(kv_shape)
v = self._make_tensor(kv_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = self._run_forward_backward(
q, k, v, sin, cos, use_neox_rotary_style=False
)
self.assertEqual(list(out_q.shape), q_shape)
self.assertEqual(list(out_k.shape), kv_shape)
self.assertEqual(list(out_v.shape), kv_shape)
def test_k_zero_num_heads(self):
"""k with 0 num_heads should not crash."""
q_shape = [
self.batch_size,
self.seq_len,
self.num_heads_q,
self.head_dim,
]
k_shape = [self.batch_size, self.seq_len, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(k_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = self._run_forward_backward(
q, k, None, sin, cos, use_neox_rotary_style=False
)
self.assertEqual(list(out_q.shape), q_shape)
self.assertEqual(list(out_k.shape), k_shape)
def test_kv_zero_num_heads(self):
"""Both k and v with 0 num_heads should not crash."""
q_shape = [
self.batch_size,
self.seq_len,
self.num_heads_q,
self.head_dim,
]
kv_shape = [self.batch_size, self.seq_len, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(kv_shape)
v = self._make_tensor(kv_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = self._run_forward_backward(
q, k, v, sin, cos, use_neox_rotary_style=False
)
self.assertEqual(list(out_q.shape), q_shape)
self.assertEqual(list(out_k.shape), kv_shape)
self.assertEqual(list(out_v.shape), kv_shape)
def test_v_zero_num_heads_neox_style(self):
"""v with 0 num_heads, neox rotary style, should not crash."""
q_shape = [
self.batch_size,
self.seq_len,
self.num_heads_q,
self.head_dim,
]
kv_shape = [self.batch_size, self.seq_len, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(kv_shape)
v = self._make_tensor(kv_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = self._run_forward_backward(
q, k, v, sin, cos, use_neox_rotary_style=True
)
self.assertEqual(list(out_q.shape), q_shape)
self.assertEqual(list(out_k.shape), kv_shape)
self.assertEqual(list(out_v.shape), kv_shape)
def test_v_zero_num_heads_time_major(self):
"""v with 0 num_heads, time_major=True, should not crash."""
# time_major: [seq_len, batch_size, num_heads, head_dim]
q_shape = [
self.seq_len,
self.batch_size,
self.num_heads_q,
self.head_dim,
]
kv_shape = [self.seq_len, self.batch_size, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(kv_shape)
v = self._make_tensor(kv_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = self._run_forward_backward(
q, k, v, sin, cos, use_neox_rotary_style=False, time_major=True
)
self.assertEqual(list(out_q.shape), q_shape)
self.assertEqual(list(out_k.shape), kv_shape)
self.assertEqual(list(out_v.shape), kv_shape)
def test_q_grad_shape_with_zero_kv(self):
"""Backward pass gradient shape for q should be correct when k/v have 0 heads."""
q_shape = [
self.batch_size,
self.seq_len,
self.num_heads_q,
self.head_dim,
]
kv_shape = [self.batch_size, self.seq_len, 0, self.head_dim]
q = self._make_tensor(q_shape)
k = self._make_tensor(kv_shape)
v = self._make_tensor(kv_shape)
sin, cos = self._make_sin_cos()
out_q, out_k, out_v = fused_rotary_position_embedding(
q, k, v, sin=sin, cos=cos, use_neox_rotary_style=False
)
out_q.sum().backward()
self.assertEqual(list(q.grad.shape), q_shape)
if __name__ == "__main__":
unittest.main()