基于 TinyChart-3B 的层次化结构建模适配器,通过 DETR 元素检测 + GCN 关系建模 + Cross-Attention 特征融合,增强图表理解能力。
本项目在 TinyChart-3B (SigLIP-SO400M + Phi-2) 的视觉编码器与投影器之间插入一个 ~20.3M 参数的层次化结构适配器,使模型能够显式感知图表的几何结构(标题、图例、数据区域等)和语义关系(包含、相邻等),从而提升 ChartQA 等下游任务的性能。
输入 [B, 3, 768, 768]
→ [Frozen] SigLIP+ToMe (27层) → [B, N, 1152]
→ [Trainable] VisualGeometryPerception (DETR, ~16.7M)
→ [Trainable] SemanticRelationModeling (GCN, ~0.67M)
→ [Trainable] MultiFeatureFusion (CrossAttn, ~2.9M)
→ 输出 [B, N, 1152]
→ [Frozen] Resampler Projector → [B, 64, 2560]
→ [Frozen] Phi-2 (2.7B)
适配器总参数: ~20.3M
# 安装依赖
uv sync
# 激活虚拟环境
source .venv/bin/activate
# 阶段1: 结构预训练 (DePlot 数据集)
uv run python train.py --stage 1 --config configs/stage1_config.yaml
# 阶段1 调试模式
uv run python train.py --stage 1 --debug
# 阶段2: 端到端微调 (ChartQA 数据集)
uv run python train.py --stage 2 --config configs/stage2_config.yaml
# 从检查点恢复训练
uv run python train.py --stage 1 --resume outputs/stage1/checkpoint_epoch3.pt
# Chart-to-Table 评估 (RMSF1)
uv run python evaluate.py --task table --checkpoint outputs/stage1/best_model.pt
# 元素检测评估 (mAP@IoU)
uv run python evaluate.py --task detection --checkpoint outputs/stage1/best_model.pt
# ChartQA 评估 (Relaxed Accuracy)
uv run python evaluate.py --task chartqa --checkpoint outputs/stage2/best_model.pt
# 运行全部测试 (61个)
uv run python -m pytest tests/ -v -k "not slow"
# 运行单个模块测试
uv run python -m pytest tests/test_hierarchical_adapter.py -v
vis-struct-exp/
├── models/
│ ├── tinychart_base.py # TinyChart 模型加载/冻结封装
│ ├── visual_geometry.py # DETR 元素检测器 (~16.7M params)
│ ├── semantic_relation.py # GCN 图关系建模 (~0.67M params)
│ ├── feature_fusion.py # Cross-attention 特征融合 (~2.9M params)
│ ├── hierarchical_adapter.py # 主适配器组装 (三个子模块串联)
│ └── tinychart/ # TinyChart 源码 (fork from mPLUG-DocOwl)
├── training/
│ ├── losses.py # 元素检测/关系预测/表格重建损失
│ ├── train_stage1_table.py # Stage 1 结构预训练
│ └── train_stage2_qa.py # Stage 2 端到端微调
├── evaluation/
│ ├── eval_table.py # RMSF1 指标
│ ├── eval_detection.py # mAP@IoU 指标
│ └── eval_chartqa.py # Relaxed Accuracy 指标
├── data/
│ └── data_loader.py # ChartQA/DePlot/PlotQA 数据集加载
├── utils/
│ └── visualization.py # 结构图叠加可视化
├── configs/
│ ├── stage1_config.yaml # lr=1e-4, batch=32, 冻结全部基座
│ └── stage2_config.yaml # lr=5e-5, batch=16, 解冻 projector
├── tests/ # 61 个测试 (全部通过)
├── train.py # 训练入口
└── evaluate.py # 评估入口
| 组件 | Stage 1 | Stage 2 |
|---|---|---|
| SigLIP 视觉编码器 | ❄️ 冻结 | ❄️ 冻结 |
| 适配器 (~20M) | 🔥 训练 | 🔥 训练 |
| Resampler 投影器 | ❄️ 冻结 | 🔥 训练 |
| Phi-2 语言模型 | ❄️ 冻结 | ❄️ 冻结 |
| 显卡 | Stage 1 (batch=32) | Stage 2 (batch=16) |
|---|---|---|
| 24 GB (3090/4090) | ✅ ~14-16 GB | ⚠️ ~18-22 GB,可能需降 batch |
| 48 GB (A6000) | ✅ 宽裕 | ✅ 可行 |
| 80 GB (A100) | ✅ 非常宽裕 | ✅ 宽裕 |
Stage 2 在 24GB 显卡上如遇 OOM,可降低 batch_size 至 8 并增大 gradient_accumulation_steps。
encode_images() 调用适配器本项目采用 MIT 许可证。详见 LICENSE 文件。
7 commits
Python
100.0%
基于 TinyChart-3B 的层次化结构建模适配器,通过 DETR 元素检测 + GCN 关系建模 + Cross-Attention 特征融合,增强图表理解能力。
本项目在 TinyChart-3B (SigLIP-SO400M + Phi-2) 的视觉编码器与投影器之间插入一个 ~20.3M 参数的层次化结构适配器,使模型能够显式感知图表的几何结构(标题、图例、数据区域等)和语义关系(包含、相邻等),从而提升 ChartQA 等下游任务的性能。
输入 [B, 3, 768, 768]
→ [Frozen] SigLIP+ToMe (27层) → [B, N, 1152]
→ [Trainable] VisualGeometryPerception (DETR, ~16.7M)
→ [Trainable] SemanticRelationModeling (GCN, ~0.67M)
→ [Trainable] MultiFeatureFusion (CrossAttn, ~2.9M)
→ 输出 [B, N, 1152]
→ [Frozen] Resampler Projector → [B, 64, 2560]
→ [Frozen] Phi-2 (2.7B)
适配器总参数: ~20.3M
# 安装依赖
uv sync
# 激活虚拟环境
source .venv/bin/activate
# 阶段1: 结构预训练 (DePlot 数据集)
uv run python train.py --stage 1 --config configs/stage1_config.yaml
# 阶段1 调试模式
uv run python train.py --stage 1 --debug
# 阶段2: 端到端微调 (ChartQA 数据集)
uv run python train.py --stage 2 --config configs/stage2_config.yaml
# 从检查点恢复训练
uv run python train.py --stage 1 --resume outputs/stage1/checkpoint_epoch3.pt
# Chart-to-Table 评估 (RMSF1)
uv run python evaluate.py --task table --checkpoint outputs/stage1/best_model.pt
# 元素检测评估 (mAP@IoU)
uv run python evaluate.py --task detection --checkpoint outputs/stage1/best_model.pt
# ChartQA 评估 (Relaxed Accuracy)
uv run python evaluate.py --task chartqa --checkpoint outputs/stage2/best_model.pt
# 运行全部测试 (61个)
uv run python -m pytest tests/ -v -k "not slow"
# 运行单个模块测试
uv run python -m pytest tests/test_hierarchical_adapter.py -v
vis-struct-exp/
├── models/
│ ├── tinychart_base.py # TinyChart 模型加载/冻结封装
│ ├── visual_geometry.py # DETR 元素检测器 (~16.7M params)
│ ├── semantic_relation.py # GCN 图关系建模 (~0.67M params)
│ ├── feature_fusion.py # Cross-attention 特征融合 (~2.9M params)
│ ├── hierarchical_adapter.py # 主适配器组装 (三个子模块串联)
│ └── tinychart/ # TinyChart 源码 (fork from mPLUG-DocOwl)
├── training/
│ ├── losses.py # 元素检测/关系预测/表格重建损失
│ ├── train_stage1_table.py # Stage 1 结构预训练
│ └── train_stage2_qa.py # Stage 2 端到端微调
├── evaluation/
│ ├── eval_table.py # RMSF1 指标
│ ├── eval_detection.py # mAP@IoU 指标
│ └── eval_chartqa.py # Relaxed Accuracy 指标
├── data/
│ └── data_loader.py # ChartQA/DePlot/PlotQA 数据集加载
├── utils/
│ └── visualization.py # 结构图叠加可视化
├── configs/
│ ├── stage1_config.yaml # lr=1e-4, batch=32, 冻结全部基座
│ └── stage2_config.yaml # lr=5e-5, batch=16, 解冻 projector
├── tests/ # 61 个测试 (全部通过)
├── train.py # 训练入口
└── evaluate.py # 评估入口
| 组件 | Stage 1 | Stage 2 |
|---|---|---|
| SigLIP 视觉编码器 | ❄️ 冻结 | ❄️ 冻结 |
| 适配器 (~20M) | 🔥 训练 | 🔥 训练 |
| Resampler 投影器 | ❄️ 冻结 | 🔥 训练 |
| Phi-2 语言模型 | ❄️ 冻结 | ❄️ 冻结 |
| 显卡 | Stage 1 (batch=32) | Stage 2 (batch=16) |
|---|---|---|
| 24 GB (3090/4090) | ✅ ~14-16 GB | ⚠️ ~18-22 GB,可能需降 batch |
| 48 GB (A6000) | ✅ 宽裕 | ✅ 可行 |
| 80 GB (A100) | ✅ 非常宽裕 | ✅ 宽裕 |
Stage 2 在 24GB 显卡上如遇 OOM,可降低 batch_size 至 8 并增大 gradient_accumulation_steps。
encode_images() 调用适配器本项目采用 MIT 许可证。详见 LICENSE 文件。
7 commits
Python
100.0%