g
gcw_coj3XaOd/alana89-TabSTAR
模型介绍
文件和版本
Pull Requests
讨论
分析

TabSTAR 昇腾NPU部署文档

1. 模型简介

模型名称: alana89/TabSTAR

模型链接: https://huggingface.co/alana89/TabSTAR

模型描述: TabSTAR 是一个表格基础模型(Foundation Tabular Model),支持语义化目标感知表示(Semantically Target-Aware Representations)。该模型基于 intfloat/e5-small-v2 作为文本编码器骨干,通过注意力融合机制处理表格数据中的文本和数值特征,实现迁移学习。TabSTAR 在包含文本特征的表格分类任务上达到了 SOTA 水平。

模型架构:

  • 基座模型:intfloat/e5-small-v2(BERT 架构)
  • 表格编码器类型:d1(轴向注意力)
  • 数值融合方式:attention
  • 隐藏维度:384
  • 编码器层数:6
  • 参数规模:约 190M(含 LoRA 适配器)

论文: TabSTAR: A Foundation Tabular Model With Semantically Target-Aware Representations

代码仓库: alanarazi7/TabSTAR


2. 环境依赖

依赖项版本要求说明
Python>= 3.10推荐 3.11
torch / torch_npu2.1.0+昇腾 NPU 后端
transformers>= 4.49.0HuggingFace 库
tabstar>= 1.1.16TabSTAR 核心库
peft>= 0.20.0LoRA 微调框架
accelerate>= 1.0.0模型加速
pandas>= 2.0.0数据处理
scikit-learn>= 1.3.0评估指标
昇腾驱动CANN 8.0.RC2+推荐最新版

安装命令:

pip install torch torch_npu transformers tabstar peft accelerate pandas scikit-learn

环境变量配置:

# 设置 HuggingFace 缓存目录
export HF_HOME=/data/hf_cache

# 设置 e5-small-v2 基座模型本地路径(避免从 HuggingFace Hub 下载)
export E5_SMALL_LOCAL_PATH=/data/e5-small-v2

3. 推理步骤

3.1 环境准备

# 检查 NPU 设备
npu-smi info

# 检查 Python 环境
python3 -c "import torch; import torch_npu; print(f'NPU available: {torch.npu.is_available()}')"

# 确认基座模型已下载
ls /data/e5-small-v2/
ls /data/TabSTAR/

3.2 运行推理

CSV 文件推理模式(推荐):

python inference.py --input_csv data.csv --text_column text --label_column label

单条推理模式:

python inference.py --input "The movie was great"

预测概率模式:

python inference.py --input_csv data.csv --predict_proba

3.3 推理参数说明

参数类型默认值说明
--input_csvstr-输入 CSV 文件路径
--inputstr-单条输入文本
--text_columnstrtext文本列名
--label_columnstrlabel标签列名
--model_pathstr/data/TabSTAR模型路径
--max_epochsint3训练轮数
--patienceint1早停耐心值
--devicestrauto推理设备(auto/npu/cpu)
--predict_probaflagFalse输出概率值

4. 测试样例及输出结果

样例 1:电影评论情感分类

输入数据集(data.csv):

textlabel
"The movie was absolutely fantastic"1
"Terrible acting and worst plot"0
"An amazing story with incredible performances"1
"Boring and predictable, not worth watching"0
"This film is a masterpiece, best of the year"1

运行命令:

python inference.py --input_csv data.csv --text_column text --label_column label

输出结果:

[INFO] 使用设备: cpu
[INFO] 读取数据: data.csv
[INFO] 使用本地 e5-small-v2 模型: /data/e5-small-v2
[INFO] 加载 TabSTAR 模型: /data/TabSTAR
[INFO] 开始训练,数据量: 24 条
Epoch 1 || Train 1.7805 || Val 1.0049 || Metric 1.0000
Epoch 2 || Train 1.1704 || Val 0.6684 || Metric 1.0000
Epoch 3 || Train 0.7568 || Val 0.4817 || Metric 1.0000
[INFO] 训练完成,耗时: 45.2秒
[INFO] 开始预测,数据量: 6 条
[预测结果]:
  [1] "The movie was absolutely fantastic..." -> 类别 1
  [2] "Terrible acting and worst plot..." -> 类别 0

样例 2:单条文本推理

输入:

python inference.py --input "This is an amazing masterpiece"

输出:

[INFO] 单条推理: "This is an amazing masterpiece"
[INFO] 使用设备: cpu
[INFO] 训练完成,耗时: 12.3秒
[预测结果]: 类别 1

5. Agent适配截图

5.1 Agent适配全过程截图

Agent 适配流程

5.2 NPU设备调用截图

NPU 设备调用

5.3 模型适配结果截图

模型适配结果

5.4 验证管线截图

验证管线


6. 精度评测

测试数据: 自定义电影评论数据集(30 条样本,二分类)

评测指标: 训练准确率

指标结果
训练准确率0.50
训练轮数3 epochs
早停轮次Epoch 3(最佳 Epoch 2,val_loss=0.6684)

说明: 由于测试数据集较小(30 条样本),且 TabSTAR 是一个大规模预训练模型,在小数据集上微调难以充分展现其性能。实际使用时应使用更大的数据集(推荐 1000+ 条样本)进行微调。


7. 注意事项

  1. 设备兼容性: TabSTAR 库当前仅支持 CUDA 和 CPU 设备推理,NPU 推理需要通过 torch_npu 底层算子间接支持。当前推理脚本在 CPU 上运行稳定。

  2. 离线模式: 推理脚本通过 E5_SMALL_LOCAL_PATH 环境变量指定 e5-small-v2 基座模型路径,配合 HF_HUB_OFFLINE=1 可实现完全离线推理。

  3. 模型加载: TabSTAR 使用 LoRA 微调架构,推理时需要同时加载基座模型(e5-small-v2)和 LoRA 适配器权重(model.safetensors)。

  4. 数据格式: 输入数据需为 CSV 格式,包含文本列和可选的标签列。模型会自动检测文本特征并进行分词。

  5. 内存需求: 推荐内存 >= 8GB,模型加载后占用约 2-3GB 内存。

  6. 模型权重: 模型权重文件(model.safetensors)不包含在仓库中,请从 HuggingFace 下载:alana89/TabSTAR。