MAIFS:Model-Agnostic Ising Feature Selection
MAIFS 是一个面向 PyTorch 的模型无关特征选择插件。它把一个普通
torch.nn.Module 包起来,在输入进入基模型之前先做内部特征预处理,再乘上
一个二进制特征 mask,然后通过 QUBO/Ising 求解器决定哪些特征保留。
只要满足下面条件,就可以接入 MAIFS:
- 基模型是
torch.nn.Module; - loss 是 PyTorch 标量 loss;
- loss 可以对输入 mask 求导;
- 输入张量里有明确的特征维度。
当前公开求解器包括:
local_search:本地贪心翻转搜索;sa:普通本地模拟退火;kaiwu_cim:Kaiwu CIM 真机/云端求解。
算法流程
MAIFS 的核心思想是:先用 PyTorch 正常训练基模型权重,再把“选择哪些特征” 建模成一个二进制优化问题。
整体流程如下:
输入 x
│
▼
内部特征预处理:
连续特征标准化,0/1 二值特征保持不变
│
▼
乘上二进制 mask
│
▼
基模型 nn.Module
│
▼
计算标量 loss
│
▼
PyTorch autograd 计算 loss 对连续 mask 的 gradient / Hessian
│
▼
构造 QUBO:
0.5 * z @ q @ z + c @ z
z ∈ {0, 1}^n
│
▼
QuadraticLinearSolver:
q, c → QUBO 矩阵 → Ising 矩阵
│
▼
具体 Ising 求解函数:
local_search / sa / kaiwu_cim
│
▼
得到 Ising 自旋解
│
▼
恢复为 0/1 特征 mask
│
▼
用原始 q、c 重新计算目标函数,选择最优 mask
│
▼
写回 FeatureSelectionWrapper.mask
其中 QuadraticLinearSolver 是唯一 Adapter。它负责把二次项 q 和一次项
c 转成 Ising 矩阵。local_search、sa 和 kaiwu_cim 再分别接收
这个 Ising 矩阵并执行求解;其中 kaiwu_cim 直接调用 Kaiwu 提供的
kw.cim.CIMOptimizer。
PyTorch 版和传统 NumPy 版的主要区别
传统 NumPy 版本通常要给不同模型分别写梯度和 Hessian,例如线性回归一套、 逻辑回归一套。PyTorch 版不需要这样做。
在 PyTorch 版中,MAIFS 只要求基模型是 nn.Module,loss 是标量。然后用
PyTorch autograd 自动计算 loss 对 mask 的一阶导和二阶导。因此,同一套
mask 更新逻辑可以包住线性回归、逻辑回归、多层神经网络等不同模型。
安装
从源码目录安装:
pip install -r requirements.txt
pip install -e .
当前 requirements.txt 内容和 Kaiwu 1.3.1 对齐:
numpy==2.2.6
torch>=2.0
kaiwu==1.3.1
注意:kaiwu==1.3.1 依赖 numpy==2.2.6,不要再把 NumPy 限制成
numpy<2.0,否则 pip 会出现依赖冲突。
最简使用示例:线性回归 + local_search
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from maifs import FeatureSelectionWrapper
torch.manual_seed(0)
x = torch.randn(64, 5)
y = 2.0 * x[:, :1] - 3.0 * x[:, 1:2]
loader = DataLoader(TensorDataset(x, y), batch_size=64, shuffle=False)
model = nn.Linear(5, 1)
selector = FeatureSelectionWrapper(
model,
feature_dim=5,
cardinality_k=2,
solver="local_search",
solver_kwargs={"max_iter": 2000},
mask_update_epochs=10,
)
loss_fn = nn.MSELoss()
optimizer = torch.optim.Adam(selector.model.parameters(), lr=0.05)
selector.fit_weights(loader, loss_fn, optimizer, epochs=20)
mask = selector.update_mask(loader, loss_fn, hessian_mode="diagonal")
print(mask)
print(selector.selected_indices())
这里的 mask 是长度为 feature_dim 的 0/1 数组:
1表示保留该特征;0表示屏蔽该特征。
交替训练
如果希望插件内部按固定周期更新 mask,可以在初始化时传入
mask_update_epochs。然后用户只需要调用 fit_weights(),每训练到指定
epoch 周期时,内部会自动调用 update_mask()。
selector = FeatureSelectionWrapper(
model,
feature_dim=n_features,
cardinality_k=10,
solver="local_search",
solver_kwargs={"max_iter": 2000},
mask_update_epochs=10,
)
selector.fit_weights(
loader,
loss_fn,
optimizer,
epochs=100,
)
print(selector.selected_indices())
含义是:
- 每个 epoch 都训练一次基模型权重;
- 每隔
mask_update_epochs个 epoch 更新一次 mask; - 更新 mask 时,基模型权重固定,只优化特征选择。
求解器参数 solver_kwargs
FeatureSelectionWrapper 只提供一个统一参数 solver_kwargs。不同求解器自己的参数都放
进这个字典,不在主类里拆成很多独立参数。
1. local_search
from kaiwu.cim._optimizer_adapter import TaskMode
selector = FeatureSelectionWrapper(
model,
feature_dim=5,
cardinality_k=2,
solver="local_search",
solver_kwargs={"max_iter": 2000},
)
参数:
max_iter:最大单点翻转搜索次数,默认2000。
2. sa
selector = FeatureSelectionWrapper(
model,
feature_dim=5,
cardinality_k=2,
solver="sa",
solver_kwargs={
"max_iter": 5000,
"random_state": 0,
},
)
参数:
max_iter:最大随机翻转次数,默认2000;random_state:随机种子,默认0。
3. kaiwu_cim
selector = FeatureSelectionWrapper(
model,
feature_dim=5,
cardinality_k=2,
solver="kaiwu_cim",
solver_kwargs={
"target_precision": 14,
"max_bits": 1000,
"max_precision": 32,
"precision_step": 4,
"sample_number": 512,
"task_mode": TaskMode.SAMPLE,
"project_no": "你的 Kaiwu 项目 ID",
"sample_sort_mode": 1,
},
)
参数:
target_precision:提交到 CIM 前的目标矩阵精度,默认14;max_bits:精度拆分后允许的最大 bit 数,默认1000;max_precision:搜索允许的最大源精度,默认32;precision_step:精度搜索步长,默认4;sample_number:CIM 采样数量,默认512;task_mode:Kaiwu 任务模式,默认"sample";project_no:Kaiwu 项目 ID,填 CPQC-X 项目列表里的项目编号;sample_sort_mode:Kaiwu 采样结果排序方式,默认1;save_dir:CIM 临时记录目录,默认None;cleanup_records:是否清理临时记录,默认True。
这些都是求解策略参数,不是 license 参数。通常只改输入特征维度时,不一定要改
这些默认值;如果真机资源或问题规模变化较大,再调整 max_bits、
target_precision 或 sample_number。
Kaiwu CIM 真机运行
kaiwu_cim 会提交真实 CIM 任务,可能消耗额度,所以不会放在 tests/ 目录下
让 pipeline 自动运行。
真机测试文件放在:
manual_tests/test_maifs_cim_linear.py
手动运行:
cd D:\maifs
python manual_tests\test_maifs_cim_linear.py
运行前需要保证:
- 当前环境已安装
kaiwu==1.3.1; - Kaiwu license 已配置;
- 当前网络能访问 Kaiwu 真机服务;
- 账户额度足够。
常见真机错误包括:
- license 未配置;
- 网络不可达;
- 资源不足,例如
CPQC-1资源不足; - 提交矩阵规模或精度超过当前资源限制。
特征数量约束
MAIFS 会在求解器返回结果后做一次兜底修正,避免出现不合理的 mask。
selector = FeatureSelectionWrapper(
model,
feature_dim=100,
min_selected_features=3,
max_selected_features=20,
)
含义:
min_selected_features:最少保留多少个特征;max_selected_features:最多保留多少个特征;cardinality_k:希望选择的目标特征数,会作为 QUBO 惩罚项加入目标函数。
如果用户不设置 min_selected_features,MAIFS 默认至少保留一部分特征,避免
mask 全为 0 导致模型没有输入信息。
本地测试
普通测试不会提交 CIM 真机任务:
python -m pytest tests
手动 CIM 测试单独运行:
python manual_tests/test_maifs_cim_linear.py
错误处理
try:
selector.update_mask(loader, loss_fn, hessian_mode="diagonal")
except ImportError as exc:
print("缺少可选依赖:", exc)
except RuntimeError as exc:
print("mask 更新失败:", exc)
一般来说:
ImportError表示依赖缺失;ValueError表示输入参数或矩阵形状不合法;RuntimeError表示求解器运行失败,或者 CIM 真机提交失败。
Release files for maifs 0.1.2
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| maifs-0.1.2.tar.gz | 25.8 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| maifs-0.1.2-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 47.8 kB
Release files / maifs-0.1.2.tar.gz
| Download URL | maifs-0.1.2.tar.gz |
|---|---|
| Size | 25.8 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
366c9e3825147044b0fc3ed849d89e132078a091ead2de750632c08e2f599284
|
|
BLAKE2b-256 checksum How to use checksums |
807195ffdcde6f1226329d7efe6a0522303906d423ef4d6be55d87b55594b0ee
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.10.20
|
Release files / maifs-0.1.2-py3-none-any.whl
| Download URL | maifs-0.1.2-py3-none-any.whl |
|---|---|
| Size | 22.1 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
398bf74d93ed081f555f911d4da51b178c57c32a247bf58e6cc735ce7d927b60
|
|
BLAKE2b-256 checksum How to use checksums |
6188793aca7f9cd6110ca54ab05a80dc0069fe2b6977deef6390d96a3b1e8aa0
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.10.20
|