请写一个precompute_freqs_cis函数的测试方法def precompute_freqs_cisdim int end int theta float = 100000 freqs = 10 theta torcharange0 dim 2 dim 2float dim t = torcharangeend device=freqsdevice # type i
下面是一个测试precompute_freqs_cis函数的例子:
import torch
from typing import Tuple
def test_precompute_freqs_cis():
dim = 4
end = 5
theta = 10000.0
# 调用被测试的函数
freqs_cis = precompute_freqs_cis(dim, end, theta)
# 验证输出的类型和形状是否正确
assert isinstance(freqs_cis, torch.Tensor)
assert freqs_cis.dtype == torch.complex64
assert freqs_cis.shape == (end, dim // 2)
# 验证计算是否正确
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
t = torch.arange(end, device=freqs.device)
expected_freqs_cis = torch.polar(torch.ones(end, dim // 2), torch.outer(t, freqs))
assert torch.allclose(freqs_cis, expected_freqs_cis)
你可以运行test_precompute_freqs_cis()来执行测试
原文地址: http://www.cveoy.top/t/topic/iriK 著作权归作者所有。请勿转载和采集!