如何计算efficient_kan构建的KAN_Linear的Params与MACs?
解决KAN_Linear无法被ptflops/thop计算参数量与MAC的问题
核心原因
ptflops、thop这类性能分析工具默认只识别PyTorch内置的标准模块(如nn.Linear、nn.Conv2d),对于自定义的KAN_Linear模块,它们无法自动解析内部参数和计算逻辑,因此会返回0值的统计结果。
可行解决方案
方法1:为KAN_Linear注册自定义统计钩子
针对ptflops和thop分别编写统计逻辑,让工具能正确识别模块的参数和计算量:
适配ptflops
from ptflops import get_model_complexity_info import torch from efficient_kan import KAN_Linear def count_kan_linear_flops(module, input, output): input_dim = input[0].size(1) output_dim = output.size(1) batch_size = input[0].size(0) # 统计参数总数:基函数参数 + 权重参数 + 偏置(可选) total_params = module.bases.numel() + module.weights.numel() if module.bias is not None: total_params += module.bias.numel() # 统计MAC:基函数映射计算 + 线性变换计算 total_macs = batch_size * (input_dim * module.num_bases + module.num_bases * output_dim) # 将结果注入ptflops的计数器 module.__flops__ += total_macs module.__params__ += total_params # 为KAN_Linear注册统计钩子 KAN_Linear.register_forward_hook(count_kan_linear_flops) # 测试统计 model = KAN_Linear(in_features=3, out_features=5) flops, params = get_model_complexity_info(model, (3,), as_strings=True, print_per_layer_stat=True) print(f"Flops: {flops}, Params: {params}")
适配thop
from thop import profile, clever_format import torch from efficient_kan import KAN_Linear def count_kan_linear_ops(module, input, output): input_dim = input[0].size(1) output_dim = output.size(1) batch_size = input[0].size(0) # 计算MAC和参数 macs = batch_size * (input_dim * module.num_bases + module.num_bases * output_dim) params = module.bases.numel() + module.weights.numel() if module.bias is not None: params += module.bias.numel() return macs, params # 测试统计 model = KAN_Linear(in_features=3, out_features=5) input_tensor = torch.randn(1, 3) macs, params = profile(model, inputs=(input_tensor,), custom_ops={KAN_Linear: count_kan_linear_ops}) macs, params = clever_format([macs, params], "%.3f") print(f"MACs: {macs}, Params: {params}")
方法2:调整KAN_Linear实现以兼容工具
检查KAN_Linear的源码,确保所有可训练参数都用nn.Parameter封装(原项目已满足),同时在forward方法中尽量使用PyTorch内置的矩阵乘法(如torch.matmul)代替自定义封装逻辑,让工具能自动捕获计算操作。
方法3:手动计算参数量与MAC
如果不想修改代码或注册钩子,可直接手动计算:
- 参数量:
in_features * num_bases + num_bases * out_features + (out_features if use_bias else 0) - MAC(单样本):
in_features * num_bases + num_bases * out_features
内容的提问来源于stack exchange,提问作者Viento
相关产品推荐
相关产品推荐

