下面是一个测试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()来执行测试

请写一个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

原文地址: http://www.cveoy.top/t/topic/iriK 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录