sheepss/autogluon-mitra-classifier-1.1-NPU
模型介绍
文件和版本
Pull Requests
讨论
分析

autogluon/mitra-classifier-1.1 昇腾 NPU 适配(Ascend 910B)

Mitra classifier 是 AutoGluon 的表格基础模型:12 层 Transformer、75,670,026 个 fp32 参数,在纯合成表格数据上预训练,推理时以训练行作为支持集(in-context learning)完成分类。本仓库将原始权重 autogluon/mitra-classifier-1.1(revision 10eead743ab2d810e0738fe2234c566bc8cea0d6)在单卡昇腾 NPU 上完成真实表格分类推理验证(平台标签 #NPU,Ascend 910B)。

任务与数据契约

  • 任务类型:表格分类(tabular-classification),pipeline_tag 与官方模型卡一致。
  • 数据集:sklearn load_wine(AutoGluon Mitra 官方 README 示例数据),13 个数值特征列,顺序固定为 alcohol, malic_acid, ash, alcalinity_of_ash, magnesium, total_phenols, flavanoids, nonflavanoid_phenols, proanthocyanins, color_intensity, hue, od280/od315_of_diluted_wines, proline。
  • 划分:80/20 训练/测试(train_test_split(random_state=42, stratify=y)),142 行训练 / 36 行测试;类别 {0, 1, 2},标签编码与 sklearn 原始值一致。
  • 输入 dtype:特征在送入 estimator 前转为 float32;缺失值策略:wine 数据无缺失,Mitra 预处理器对数值列做标准化变换,列顺序不得改变。
  • 输出:(N, 3) 概率矩阵,类别顺序 [0, 1, 2],预测标签为 argmax。

环境与依赖

  • 硬件:Atlas Ascend 910B(Ascend910_9362,单卡 npu:0,CANN 8.5.1)。
  • 软件:Python 3.11、torch 2.9.0、torch-npu 2.9.0.post1、autogluon.tabular[mitra]==1.5.0、huggingface_hub、safetensors、scikit-learn、pandas。
  • 安装:pip install -r requirements.txt(torch/torch_npu 请按昇腾官方渠道安装配套版本)。
  • 权重获取:默认由 inference.py 通过 snapshot_download(repo_id="autogluon/mitra-classifier-1.1", revision="10eead743ab2d810e0738fe2234c566bc8cea0d6") 显式下载;离线环境可设 MITRA11_LOCAL_DIR 指向本地 checkpoint 目录。

NPU 推理

python inference.py

脚本自包含:检查 NPU 可用性 → 加载本地缓存权重 → wine 数据固定划分 → MitraClassifier(fine_tune=False, device="npu:0") fit → 测试集 predict_proba → 打印预测、概率、accuracy/macro-F1 与同步耗时。任何异常返回非零退出码。

适配要点:AutoGluon 1.5.0 的 Tab2D.from_pretrained 仅支持 Hub repo id,本仓库在其外层做了本地目录感知加载补丁(config.json + safetensors 直接读入后再迁移到 npu:0),未修改任何模型结构与权重。fit 阶段的预处理(标准化)保留在 CPU/pandas 层,Transformer 前向的权重与全部参与计算的 Tensor 位于 npu:0(日志可见 first_parameter_device=npu:0)。

真实结果(本次运行)

  • 数据:wine 官方示例,142 行训练 / 36 行测试,13 特征。
  • 设备证据:模型参数 npu:0 fp32,75,670,026 参数;前向输入/输出 Tensor device 为 npu:0 -> npu:0。
  • 默认命令真实输出:accuracy = 1.0000,macro-F1 = 1.0000(CPU 与 NPU 相同),同步耗时 0.159 s(首轮含编译)。

一致性(CPU vs NPU)

同一权重、同一数据划分、同一种子(seed=0)、同一 dtype(fp32)下比较 36 行 × 3 类概率:

  • argmax 标签一致率:100%(36/36)
  • 概率最大绝对误差:1.03e-2,平均绝对误差:7.25e-4
  • 校验阈值:atol=0.01, rtol=0.05 下 elementwise 全部通过(compare_outputs.py passed=true)。阈值依据:12 层 ICL Transformer 在 fp32 下注意力累积误差的实测量级,且两端 accuracy/macro-F1 完全一致。该结果属于 smoke consistency(单数据集冒烟一致性),不是完整基准精度评测。

性能(npu:0)

  • 预热 3 次、正式计时 10 次,每次前后 torch.npu.synchronize();predict_proba 计时包含 CPU 侧预处理(已说明边界)。
  • batch shape [36, 13],dtype float32:
    • 首轮(含 NPU 图编译):约 0.16 s;稳定态 latency avg/min/max = 42.54 / 42.32 / 43.30 ms,p50/p90/p95 = 42.37 / 43.20 / 43.25 ms;
    • 吞吐 ≈ 846 rows/s;峰值显存 ≈ 502 MB。
  • 动态 shape 复跑:首次编译后分别以 10/24/36 行查询集复跑均正常。

证据图

以下三张 PNG 由 xterm.js 根据本次适配的真实 stdout/stderr 日志渲染(非手工绘制、非桌面截图):

agent workflow

npu device call

model result

限制

  • 仅验证了分类任务 + wine 官方示例数据集上的 smoke consistency 与性能;未覆盖全量公开表格基准。
  • fine_tune=True 微调路径未在 NPU 上验证;fit 阶段预处理仍在 CPU,仅 Transformer 前向位于 NPU。
  • 大数据集下 Mitra 会按 max_samples_support/max_samples_query 截断支持/查询集,行为与上游一致。