SegEarth-R2 项目改造
·
SegEarth-R2 项目改造使用手册
目标:帮助理解项目结构、数据流、训练/推理流程,达到能够基于本项目改造并参加其他比赛的能力。
目录
1. 项目概述
SegEarth-R2 是一个面向**遥感图像的语言引导分割(Reasoning Segmentation)**模型。
核心能力:输入一张遥感图像 + 一段自然语言描述(如"找出图中所有飞机"),模型输出:
- 一段推理文本
- 一个或多个分割 mask(对应文本中的
[SEG]token 位置)
技术栈:
- 多模态大模型(Mipha-3B,基于 Phi-2 语言模型)
- 视觉编码器:SigLIP / CLIP / DINOv2
- Mask2Former 风格的 mask decoder(Swin Transformer 骨干 + MSDeformAttnPixelDecoder + Transformer Decoder)
- LoRA 微调
- DeepSpeed 训练
2. 仓库结构速查
SegEarth-R2/
├── README.md # 项目说明
├── requirements.txt # Python 依赖
├── docs/ # 文档
│ ├── Installation.md # 环境安装
│ ├── Preparation.md # 数据准备
│ ├── Training.md # 训练说明
│ └── Evaluation.md # 评估说明
├── scripts/ # 启动脚本
│ ├── train.sh # 训练启动脚本
│ ├── test.sh # 评估启动脚本
│ ├── merge_lora_weights.sh # LoRA 权重合并
│ ├── zero1.json # DeepSpeed ZeRO stage1
│ ├── zero2.json # DeepSpeed ZeRO stage2
│ └── zero3.json # DeepSpeed ZeRO stage3
└── segearth_r2/
├── datasets/
│ └── dataset.py # LaSeRSDataset、DataCollator、预处理
├── eval/
│ └── eval.py # 推理评估脚本
├── model/
│ ├── language_model/
│ │ └── llava_phi.py # SegEarthR2 主模型
│ ├── mipha/model/
│ │ ├── mipha_arch.py # MiphaMetaModel / MiphaMetaForCausalLM
│ │ ├── language_model/ # Phi-2 语言模型相关
│ │ ├── multimodal_encoder/ # CLIP/SigLIP/DINOv2
│ │ └── multimodal_projector/ # 视觉-语言投影器
│ ├── mask_encoder/
│ │ └── swin_trans.py # Swin Transformer 骨干
│ ├── mask_decoder/
│ │ ├── Mask2Former_Simplify/ # Mask2Former 实现
│ │ ├── mask_config/ # YAML 配置文件
│ │ └── mask_criterion/ # 分割损失函数
│ └── datasets_mapper/
│ └── IVS_mapper.py # Detectron2 风格数据映射器
├── train/
│ ├── train.py # 训练主脚本
│ ├── llava_trainer.py # 自定义 Trainer
│ └── merge_lora_weights_and_save_hf_model.py
└── utils/
├── builder.py # load_pretrained_model
├── conversation.py # 对话模板
├── constants.py # 特殊 token 索引
└── mm_utils.py # 图像/token 工具
文件重要性分级
| 优先级 | 文件 | 说明 |
|---|---|---|
| ⭐⭐⭐⭐⭐ | segearth_r2/model/language_model/llava_phi.py |
主模型,控制前向/推理/loss |
| ⭐⭐⭐⭐⭐ | segearth_r2/datasets/dataset.py |
数据加载与预处理 |
| ⭐⭐⭐⭐⭐ | segearth_r2/train/train.py |
训练入口 |
| ⭐⭐⭐⭐ | segearth_r2/eval/eval.py |
评估入口 |
| ⭐⭐⭐⭐ | segearth_r2/utils/builder.py |
模型加载 |
| ⭐⭐⭐ | segearth_r2/model/mipha/model/mipha_arch.py |
多模态架构基类 |
| ⭐⭐⭐ | segearth_r2/model/mask_encoder/swin_trans.py |
Swin 骨干 |
| ⭐⭐⭐ | segearth_r2/model/mask_decoder/Mask2Former_Simplify/ |
Mask Decoder |
| ⭐⭐ | segearth_r2/train/llava_trainer.py |
自定义 Trainer |
| ⭐⭐ | segearth_r2/model/datasets_mapper/IVS_mapper.py |
视频/RefVOS 数据映射 |
3. 核心概念与术语
3.1 特殊 Token
| Token | 索引 | 含义 |
|---|---|---|
<image> |
-200 | 图像占位符,会被替换为图像特征 |
[SEG] |
新增 | 分割 token,LLM 预测到此处时生成对应 mask |
<refer> |
-204 | 引用 token,指示分割目标 |
[PAD] |
新增 | padding token |
3.2 两条视觉分支
| 分支 | 输入 | 用途 | 骨干 |
|---|---|---|---|
images_clip |
原始 RGB 图像 | 给 LLM 提供视觉语义 | SigLIP / CLIP |
images |
1024×1024 归一化图像 | 给 mask decoder 提供多尺度特征 | Swin Transformer |
3.3 关键超参数
| 参数 | 默认值 | 含义 |
|---|---|---|
base_data_path |
/data1/xzp/data |
数据集根目录 |
model_name_or_path |
pretrained_model/mllm/Mipha-3B |
预训练 MLLM |
vision_tower |
pretrained_model/CLIP/siglip-so400m-patch14-384 |
视觉编码器 |
vision_tower_mask |
pretrained_model/mask2former/...pkl |
Mask2Former 预训练权重 |
mask_config |
maskformer2_swin_base_384_bs16_50ep.yaml |
Mask Decoder 配置 |
lora_r |
8 | LoRA rank |
lora_alpha |
16 | LoRA alpha |
lora_target_modules |
q_proj, v_proj | LoRA 目标层 |
4. 数据格式与准备
4.1 数据集目录结构
base_data_path/
├── train/
│ ├── annotations/
│ │ └── train_data.json # 训练标注
│ └── images/ # 训练图像
└── test/
├── annotations/
│ └── test_data.json # 测试标注(可能无 mask)
└── images/ # 测试图像
4.2 标注 JSON 格式
[
{
"id": 123,
"image_name": "sample_001.jpg",
"description": "Locate all commercial aircraft in the airport.",
"answer": "The commercial aircraft are located at [SEG].",
"mask": [
{"size": [512, 512], "counts": "PQRd06jT0..."},
{"size": [512, 512], "counts": "abc123..."}
]
}
]
4.3 mask 编码
- 使用 pycocotools RLE 格式
- 在
dataset.py中通过pycocotools.mask.decode(rle)解码为 numpy 二值 mask - 比赛数据如果不是 RLE,需要改写
LaSeRSDataset.__getitem__的 mask 读取部分
4.4 参加比赛时改造数据加载
场景 1:比赛数据是 PNG mask 文件
修改位置:segearth_r2/datasets/dataset.py 中 LaSeRSDataset.__getitem__
# 原代码:从 JSON 中解码 RLE
rle_list = data_info['mask']
masks = []
for rle in rle_list:
mask = M.decode(rle)
masks.append(mask)
masks = np.stack(masks, axis=0)
# 改造后:从 PNG 文件读取
import cv2
mask_path = os.path.join(self.mask_path, data_info['mask_name'])
mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
mask = (mask > 0).astype(np.uint8)
masks = np.stack([mask], axis=0)
场景 2:比赛数据一个样本有多个 mask 类别
需要把每个类别作为一个独立 mask,并在 answer 中为每个类别加入 [SEG] token。
5. 模型架构详解
5.1 类继承关系
torch.utils.data.Dataset
└── RS_Base_Dataset
└── LaSeRSDataset
└── UnifyDatasetSingleDatasetForBatch
transformers.Trainer
└── LLaVATrainer
MiphaMetaModel
└── MiphaPhiModel
└── SegEarthR2Model
MiphaMetaForCausalLM (ABC)
└── MiphaPhiForCausalLM
└── SegEarthR2
CausalLMOutputWithPast
└── CausalOutputWithMask
5.2 SegEarthR2 主模型组成
5.3 关键前向流程
prepare_inputs_labels_for_multimodal:把<image>、<refer>替换为视觉 embedding- LLM forward:生成 hidden_states 和 logits
get_SEG_embedding:提取[SEG]token 位置的 hidden state,经SEG_token_projector得到 query embeddingget_vision_tower_feature:Swin 提取多尺度特征(res2/res3/res4/res5)pixel_decoder:把多尺度特征融合为 per-pixel embeddingspredictor:用 SEG embedding 作为 query,与 pixel features 做 cross-attention,输出 maskCriterion:计算 mask BCE loss、dice lossAttentionLoss:监督 LLM attention 关注目标区域- 最终 loss =
loss_llm + loss_mask + loss_dice + loss_attention
5.4 推理流程
- 构造 prompt:
This is an image <image> , please doing Reasoning Segmentation according to the following instruction: {description} - 生成文本序列,遇到
[SEG]时记录位置 - 从
[SEG]位置的 hidden state 得到 mask query - 运行 pixel_decoder + predictor 得到 pred_masks
- 输出文本和 mask
6. 训练流程详解
6.1 启动命令
# 从 scripts/train.sh
CUDA_VISIBLE_DEVICES=0 deepspeed --master_port=29500 segearth_r2/train/train.py \
--model_name_or_path pretrained_model/mllm/Mipha-3B \
--vision_tower pretrained_model/CLIP/siglip-so400m-patch14-384 \
--vision_tower_mask pretrained_model/mask2former/model_final_54b88a.pkl \
--base_data_path /path/to/data \
--output_dir ./output \
--max_steps 5000 \
--per_device_train_batch_size 1 \
--save_strategy steps \
--save_steps 1000 \
--bf16 True \
--lora_r 4 \
--deepspeed scripts/zero3.json \
--mask_config segearth_r2/model/mask_decoder/mask_config/maskformer2_swin_base_384_bs16_50ep.yaml \
--data_ratio 1 \
--switch_bs 4
6.2 train.py 执行步骤
渲染错误: Mermaid 渲染失败: Parse error on line 2: ... Start([开始 train()]) --> Parse["1. 解 ----------------------^ Expecting 'SQE', 'DOUBLECIRCLEEND', 'PE', '-)', 'STADIUMEND', 'SUBROUTINEEND', 'PIPE', 'CYLINDEREND', 'DIAMOND_STOP', 'TAGEND', 'TRAPEND', 'INVTRAPEND', 'UNICODE_TEXT', 'TEXT', 'TAGSTART', got 'PS'
6.3 可训练参数控制
在 train.py 中:
train_module_list = ["lm_head", "pixel_decoder", "predictor", "SEG_token_projector"]
if model_args.train_swin_backbone:
train_module_list.append('vision_tower_mask')
改造建议:
- 如果比赛数据域与遥感差异大,可以解冻
vision_tower_mask(Swin 骨干) - 如果希望学习更多语言-分割对齐,可以添加
mm_projector - 如果显存不足,可以只保留
SEG_token_projector和predictor
6.4 LoRA 配置
lora_config = LoraConfig(
r=4, # rank
lora_alpha=16, # scaling
target_modules=['q_proj', 'v_proj'],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
改造建议:
- 显存大 / 数据多:增大
lora_r到 8 或 16 - 需要更多可训练参数:添加
k_proj,o_proj,gate_proj,up_proj,down_proj
7. 评估/推理流程详解
7.1 启动命令
deepspeed --include localhost:0 --master_port=29500 segearth_r2/eval/eval.py \
--base_data_path data_path \
--model_path model_path \
--vision_tower_mask pretrained_model/mask2former/model_final_54b88a.pkl \
--mask_config segearth_r2/model/mask_decoder/mask_config/maskformer2_swin_base_384_bs16_50ep.yaml \
--output_dir output/res
7.2 eval.py 执行步骤
load_pretrained_model()加载 tokenizer、model、image_processor- 设置 conversation template 为
llava_phi - 对每个测试 split:
- 创建
LaSeRSDataset - 创建
DataLoader(batch_size=1)
- 创建
do_eval():- 对每个 sample 调用
preprocess_input() - 调用
model.inference()得到 output_ids 和 masks_pred - 计算与 GT mask 的 IoU
- 对每个 sample 调用
- 输出 gIoU 指标
7.3 保存推理结果的改造
当前 eval.py 的 output_dir 没有真正保存文件。比赛通常需要提交 mask 或可视化结果。
改造位置:segearth_r2/eval/eval.py 的 do_eval() 函数
# 在得到 masks_pred 后保存
import os
import cv2
import numpy as np
output_mask_dir = os.path.join(data_args.output_dir, split, 'masks')
os.makedirs(output_mask_dir, exist_ok=True)
if masks_pred is not None:
for i in range(masks_pred.shape[0]):
mask = masks_pred[i, 0] * 255
mask_path = os.path.join(output_mask_dir, f"{idx}_{i}.png")
cv2.imwrite(mask_path, mask.astype(np.uint8))
# 保存 per-sample 结果 JSON
import json
result = {
'image_path': image_path,
'instruction': text,
'iou': iou.item(),
'mask_files': [f"{idx}_{i}.png" for i in range(masks_pred.shape[0])]
}
with open(os.path.join(data_args.output_dir, split, 'results.jsonl'), 'a') as f:
f.write(json.dumps(result) + '\n')
8. 关键文件改造指南
8.1 数据集改造:segearth_r2/datasets/dataset.py
| 改造点 | 位置 | 说明 |
|---|---|---|
| 读取不同格式标注 | LaSeRSDataset.__init__ / __getitem__ |
JSON 格式、PNG mask、COCO format |
| 图像预处理尺寸 | preprocess_image |
默认 1024×1024,可按比赛调整 |
| mask 解码方式 | LaSeRSDataset.__getitem__ |
RLE / PNG / polygon |
| 数据增强 | preprocess_image 或新增 transform |
随机翻转、颜色抖动等 |
| 批处理 | DataCollatorForCOCODatasetV2.__call__ |
增加新的字段 |
8.2 模型改造:segearth_r2/model/language_model/llava_phi.py
| 改造点 | 位置 | 说明 |
|---|---|---|
| 修改 loss 权重 | forward() |
loss = loss_llm * w1 + loss_mask * w2 + loss_dice * w3 + loss_attention * w4 |
| 添加新任务头 | SegEarthR2.__init__ / forward() |
如分类头、边界框头 |
| 改变 mask query 来源 | get_SEG_embedding |
可加入文本 embedding |
| 替换 mask decoder | initial_mask_module |
改用 SAM、UNet 等 |
| 多尺度特征融合 | get_vision_tower_feature |
添加 FPN、PPM 等 |
8.3 训练脚本改造:segearth_r2/train/train.py
| 改造点 | 位置 | 说明 |
|---|---|---|
| 添加新数据集 | make_unify_datamodule |
支持多数据集混合训练 |
| 修改学习率 | TrainingArguments / 脚本 |
不同层使用不同 lr |
| 解冻更多参数 | train_module_list |
控制可训练模块 |
| 添加回调 | LLaVATrainer |
自定义日志、保存策略 |
| 梯度累积 | TrainingArguments.gradient_accumulation_steps |
小显存模拟大 batch |
8.4 评估脚本改造:segearth_r2/eval/eval.py
| 改造点 | 位置 | 说明 |
|---|---|---|
| 保存预测 mask | do_eval() |
生成提交文件 |
| 保存可视化 | do_eval() |
叠加 mask 到原图 |
| 计算更多指标 | do_eval() |
mIoU、F1、Dice 等 |
| 支持单图推理 | 新增函数 | 上传图片+文本直接推理 |
| 支持模型 ensemble | do_eval() |
多模型投票 |
9. 参加比赛常见改造场景
场景 1:比赛任务也是语言引导分割
步骤:
- 把比赛数据转换为 LaSeRS JSON 格式
- 修改
preprocess_image和preprocess_mask适应图像尺寸 - 调整
lora_r和训练步数 - 运行
train.sh - 运行
test.sh,并改造eval.py保存结果
场景 2:比赛只有分割 mask,没有文本描述
需要把类别名/标签扩展为自然语言描述:
# 在 dataset.py 中
CLASS_NAME = {0: 'background', 1: 'building', 2: 'road', 3: 'vehicle'}
data_info['description'] = f"Segment all {CLASS_NAME[cls_id]}s in the image."
data_info['answer'] = f"The {CLASS_NAME[cls_id]}s are located at [SEG]."
场景 3:比赛需要输出多个类别 mask
改造 LaSeRSDataset.__getitem__:
masks = [] # (N, H, W)
descriptions = []
answers = []
for cls_id in class_ids:
masks.append((gt_mask == cls_id).astype(np.uint8))
descriptions.append(f"Segment {CLASS_NAME[cls_id]}.")
answers.append(f"[SEG]")
# 或者使用统一的 instruction
instruction = "Segment buildings, roads, and vehicles."
answer = "Buildings [SEG], roads [SEG], vehicles [SEG]."
场景 4:比赛数据分辨率差异大
遥感图像通常很大,可以考虑:
- 滑动窗口切图:在
dataset.py中对大图切 patch - 多尺度训练:在
preprocess_image中随机缩放 - 大分辨率输入:修改
mask_configYAML 中的INPUT.IMAGE_SIZE
场景 5:显存不足
策略:
- 使用
zero3.json替代zero1.json - 减小
per_device_train_batch_size到 1 - 增大
gradient_accumulation_steps - 减小
lora_r到 4 或 2 - 冻结
vision_tower_mask - 使用
gradient_checkpointing=True
场景 6:需要更好的 mask 质量
策略:
- 训练更久(增加
max_steps) - 解冻
vision_tower_mask训练 Swin 骨干 - 替换 Mask2Former 为 SAM / SAM2 作为 decoder
- 增加 mask 数据增强(随机裁剪、旋转)
- 调整 loss 权重,增大 dice loss 权重
10. 调试与问题排查
10.1 常见错误
| 错误 | 原因 | 解决 |
|---|---|---|
No module named 'deepspeed' |
未安装 deepspeed | pip install deepspeed==0.10.0 |
No module named 'detectron2' |
未安装 detectron2 | 按官方说明安装 |
CUDA out of memory |
显存不足 | 减小 batch size / 使用 zero3 / 开启 gradient checkpointing |
KeyError: 'mask' |
标注 JSON 缺少 mask 字段 | 检查数据格式 |
IMAGE_TOKEN_INDEX 未替换 |
tokenizer 处理错误 | 检查 tokenizer_special_tokens |
| 预测 mask 全黑 | 训练不充分或 loss 不平衡 | 增加训练步数 / 调整 loss 权重 |
10.2 调试技巧
- 先跑通 eval.py:用官方权重验证环境是否正确
- 小数据集过拟合:取 10 张图训练,看能否 overfit
- 打印可训练参数:确认 LoRA 和 train_module_list 生效
- 可视化中间结果:在
forward()中保存pred_masks和 GT - 检查特殊 token:确认
[SEG]token 已加入 tokenizer
10.3 推荐检查清单
- 环境安装完成(torch、transformers、deepspeed、detectron2、fvcore)
- 预训练权重下载完成
- 数据路径正确
- 标注 JSON 格式正确
- mask 解码后形状正确
(N, H, W) - 图像归一化参数正确(ImageNet mean/std)
-
lora_r、lora_alpha、max_steps合理 - DeepSpeed 配置文件路径正确
- output_dir 有写入权限
11. 附录:命令速查表
11.1 环境安装
conda create -n segearthr2 python=3.10
conda activate segearthr2
pip install torch==2.1.0 torchvision==0.16.0 torchaudio==2.1.0 --index-url https://download.pytorch.org/whl/cu121
pip install -r requirements.txt
# 单独安装 detectron2
python -m pip install 'git+https://github.com/facebookresearch/detectron2.git'
# 编译 MSDeformAttn
sh segearth_r2/model/mask_decoder/Mask2Former_Simplify/modeling/pixel_decoder/ops/make.sh
11.2 训练
sh scripts/train.sh
11.3 评估
sh scripts/test.sh
11.4 合并 LoRA 权重
sh scripts/merge_lora_weights.sh
11.5 单图推理示例
import torch
from segearth_r2.utils.builder import load_pretrained_model
from segearth_r2.eval.eval import preprocess_input
from transformers import SiglipImageProcessor
model_path = "path/to/checkpoint"
mask_config = "segearth_r2/model/mask_decoder/mask_config/maskformer2_swin_base_384_bs16_50ep.yaml"
class Args:
vision_tower = "pretrained_model/CLIP"
vision_tower_mask = "pretrained_model/mask2former/model_final_54b88a.pkl"
seg_task = "instance"
tokenizer, model, image_processor, _ = load_pretrained_model(
model_path, model_args=Args(), mask_config=mask_config, device="cuda"
)
model.to(dtype=torch.float32, device="cuda")
model.eval()
clip_image_processor = SiglipImageProcessor.from_pretrained("pretrained_model/CLIP")
SEG_token_id = tokenizer.encode('[SEG]', add_special_tokens=False)[0]
text = "Segment all buildings in this image."
image_path = "path/to/image.jpg"
input_ids, images, images_clip = preprocess_input(text, image_path, tokenizer, clip_image_processor)
with torch.no_grad():
output_ids, masks_pred = model.inference(
input_ids=input_ids.cuda(),
images=images.cuda(),
images_clip=images_clip.cuda(),
do_sample=True,
temperature=0.2,
num_beams=1,
max_new_tokens=128,
eos_token_id=tokenizer.eos_token_id,
SEG_token_id=SEG_token_id
)
print(tokenizer.decode(output_ids[0], skip_special_tokens=True))
print("mask shape:", masks_pred.shape if masks_pred is not None else None)
12. 学习路径建议
| 阶段 | 目标 | 操作 |
|---|---|---|
| 第 1 周 | 跑通训练和评估 | 准备数据,运行 train.sh 和 test.sh |
| 第 2 周 | 理解数据流 | 精读 dataset.py 和 llava_phi.py 的 forward |
| 第 3 周 | 理解模型架构 | 精读 mipha_arch.py、swin_trans.py、Mask2Former 组件 |
| 第 4 周 | 简单改造 | 替换自己的数据集,保存预测结果 |
| 第 5-6 周 | 比赛适配 | 调整 loss、数据增强、模型结构 |
| 第 7-8 周 | 调参优化 | 学习率、训练步数、LoRA rank、多尺度等 |
本手册基于 SegEarth-R2 仓库代码整理,建议配合源码阅读,边改边学。
更多推荐




所有评论(0)