Shardzzi/vis-struct-exp

视觉-结构学习实验 - 轻量级视觉编码层

0

stars

7

commits

Python

primary language

Apr 7, 2026

updated

README

视觉-结构学习实验 (Visual-Structure Learning Experiment)

Python 3.11 PyTorch License

基于 TinyChart-3B 的层次化结构建模适配器,通过 DETR 元素检测 + GCN 关系建模 + Cross-Attention 特征融合,增强图表理解能力。

项目概述

本项目在 TinyChart-3B (SigLIP-SO400M + Phi-2) 的视觉编码器与投影器之间插入一个 ~20.3M 参数的层次化结构适配器,使模型能够显式感知图表的几何结构(标题、图例、数据区域等)和语义关系(包含、相邻等),从而提升 ChartQA 等下游任务的性能。

主要特性

  • 轻量级适配器: 仅 ~20.3M 可训练参数,远低于 100M 上限
  • 三阶段流水线: DETR 元素检测 → GCN 关系建模 → Cross-Attention 特征融合
  • 两阶段训练: Stage 1 结构预训练 (DePlot) + Stage 2 端到端微调 (ChartQA)
  • 最小侵入: 适配器插入冻结基座之间,不修改 TinyChart 原始权重
  • 无额外依赖: GCN 使用自定义 GraphConv,无需 torch_geometric

模型架构

输入 [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

  • VisualGeometryPerception (DETR): ~16.7M
  • SemanticRelationModeling (GCN): ~0.67M
  • MultiFeatureFusion (CrossAttn): ~2.9M

快速开始

环境配置

# 安装依赖
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 1Stage 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。

性能目标

  • Chart-to-Table RMSF1: > 94.0
  • ChartQA 准确率提升: > 2% (相比 TinyChart baseline)
  • 推理速度下降: < 20% (相比 TinyChart baseline)
  • 可训练参数: < 100M (实际 ~20.3M)

技术栈

  • Python: 3.11
  • PyTorch: ≥ 2.9.1 (CUDA 支持)
  • Transformers: ≥ 4.57.6
  • timm: ≥ 1.0.24
  • Datasets: ≥ 4.5.0
  • einops: ≥ 0.7.0
  • 包管理器: uv

后续工作

  1. 加载真实 TinyChart 权重: 下载 ~6GB 模型权重,验证端到端前向传播
  2. 接入适配器到 TinyChart: 修改 encode_images() 调用适配器
  3. Stage 1 训练: 在 DePlot 数据集上进行结构预训练
  4. Stage 2 微调: 在 ChartQA 数据集上端到端微调
  5. 性能基准测试: 验证推理速度下降 < 20%

许可证

本项目采用 MIT 许可证。详见 LICENSE 文件。

致谢

Contributors

Shardzzi

7 commits

Shardzzi/vis-struct-exp

视觉-结构学习实验 - 轻量级视觉编码层

0

stars

7

commits

Python

primary language

Apr 7, 2026

updated

README

视觉-结构学习实验 (Visual-Structure Learning Experiment)

Python 3.11 PyTorch License

基于 TinyChart-3B 的层次化结构建模适配器,通过 DETR 元素检测 + GCN 关系建模 + Cross-Attention 特征融合,增强图表理解能力。

项目概述

本项目在 TinyChart-3B (SigLIP-SO400M + Phi-2) 的视觉编码器与投影器之间插入一个 ~20.3M 参数的层次化结构适配器,使模型能够显式感知图表的几何结构(标题、图例、数据区域等)和语义关系(包含、相邻等),从而提升 ChartQA 等下游任务的性能。

主要特性

  • 轻量级适配器: 仅 ~20.3M 可训练参数,远低于 100M 上限
  • 三阶段流水线: DETR 元素检测 → GCN 关系建模 → Cross-Attention 特征融合
  • 两阶段训练: Stage 1 结构预训练 (DePlot) + Stage 2 端到端微调 (ChartQA)
  • 最小侵入: 适配器插入冻结基座之间,不修改 TinyChart 原始权重
  • 无额外依赖: GCN 使用自定义 GraphConv,无需 torch_geometric

模型架构

输入 [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

  • VisualGeometryPerception (DETR): ~16.7M
  • SemanticRelationModeling (GCN): ~0.67M
  • MultiFeatureFusion (CrossAttn): ~2.9M

快速开始

环境配置

# 安装依赖
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 1Stage 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。

性能目标

  • Chart-to-Table RMSF1: > 94.0
  • ChartQA 准确率提升: > 2% (相比 TinyChart baseline)
  • 推理速度下降: < 20% (相比 TinyChart baseline)
  • 可训练参数: < 100M (实际 ~20.3M)

技术栈

  • Python: 3.11
  • PyTorch: ≥ 2.9.1 (CUDA 支持)
  • Transformers: ≥ 4.57.6
  • timm: ≥ 1.0.24
  • Datasets: ≥ 4.5.0
  • einops: ≥ 0.7.0
  • 包管理器: uv

后续工作

  1. 加载真实 TinyChart 权重: 下载 ~6GB 模型权重,验证端到端前向传播
  2. 接入适配器到 TinyChart: 修改 encode_images() 调用适配器
  3. Stage 1 训练: 在 DePlot 数据集上进行结构预训练
  4. Stage 2 微调: 在 ChartQA 数据集上端到端微调
  5. 性能基准测试: 验证推理速度下降 < 20%

许可证

本项目采用 MIT 许可证。详见 LICENSE 文件。

致谢

Contributors

Shardzzi

7 commits

Languages

Python

100.0%