MMProLong长文档LMM训练新范式深度解析:问答对数据为何比OCR转录高效百倍——字节Seed团队与港科大联合研究
一、引言:长文档多模态训练的"隐形成本"
2026年7月底,字节跳动Seed团队与香港科技大学联合发布了MMProLong——一个突破长文档多模态大语言模型(LMM)训练效率的新框架。这项研究揭示了一个被行业长期忽视的根本性问题:长文档LMM训练的瓶颈不在模型架构,而在训练数据的组织方式。
当前主流的长文档LMM训练思路是"OCR转录":将长文档扫描为图像,用OCR提取文本,然后喂给模型。但MMProLong的研究团队发现,这种方法不仅效率低下,在特定场景下甚至损害了模型的性能。
更关键的是,他们发现:精心构造的问答对(Q&A Pairs)训练数据,在仅使用128K token的极小训练预算下,就能超越使用百万级OCR数据的传统方法。
本文将从数据组织、训练策略、跨模态迁移三个维度,深度解析MMProLong的技术架构。
二、问题建模:长文档LMM训练的"数据效率困局"
2.1 长文档LMM的核心挑战
多模态大语言模型在处理长文档时面临三个相互关联的挑战:
- 视觉-文本对齐失效:OCR将视觉布局信息压缩为纯文本,丢失了段落位置、表格结构、图表纹理等关键视觉线索
- 注意力稀释:在长上下文中,关键信息被大量无关文本稀释,模型难以定位答案
- 训练-推理分布偏移:OCR训练数据与真实推理时看到的文档图像分布不一致
2.2 数据效率的形式化分析
设一个长文档 $D$ 包含 $N$ 页图像 ${I_1, …, I_N}$ 和对应的文本 ${T_1, …, T_N}$。传统OCR方法的训练目标是:
$$ \mathcal{L}{OCR} = \sum{i=1}^N \sum_{j=1}^{|T_i|} -\log P(t_{i,j} | I_i, t_{i,<j}) $$
即模型需要从图像中重建每一个文本token。但MMProLong发现,大多数文本token对下游任务的理解是无用的——模型并不需要逐字重建文档,而是需要学会在文档中定位并回答具体问题。
MMProLong的训练目标改为:
$$ \mathcal{L}{QA} = \sum{(q, a) \in \mathcal{Q}} -\log P(a | D, q) $$
其中 $\mathcal{Q}$ 是精心构造的问答对集合,每个问答对只关注文档中的一个特定信息片段。
import numpy as np
from typing import List, Dict, Tuple
import random
class DataEfficiencyAnalyzer:
"""
数据效率分析器 - 量化OCR vs QA对训练的效率差异
核心发现:QA对的"信息密度"远高于OCR转录
"""
def __init__(self):
self.ocr_token_count = 0
self.qa_token_count = 0
self.qa_accuracy = 0.0
self.ocr_accuracy = 0.0
def analyze_document(self,
document_pages: List[Dict],
qa_pairs: List[Dict]) -> Dict:
"""
分析单个文档的OCR和QA训练效率
Args:
document_pages: 文档页列表,每页包含图像和文本
qa_pairs: 问答对列表,每个包含问题和答案
Returns:
efficiency_report: 效率分析报告
"""
# OCR数据量
ocr_tokens = sum(
len(page['text'].split())
for page in document_pages
)
# QA数据量
qa_tokens = sum(
len(qa['question'].split()) + len(qa['answer'].split())
for qa in qa_pairs
)
# 信息密度计算
# 定义:每token的信息量 = 下游任务准确率提升 / 训练token数
ocr_information_density = 0.05 / ocr_tokens # OCR训练提升约5%
qa_information_density = 0.35 / qa_tokens # QA训练提升约35%
# 效率比
efficiency_ratio = qa_information_density / ocr_information_density
print(f"=== 数据效率分析 ===")
print(f"文档总页数: {len(document_pages)}")
print(f"OCR转录token数: {ocr_tokens:,}")
print(f"QA对token数: {qa_tokens:,}")
print(f"OCR信息密度: {ocr_information_density:.6f}")
print(f"QA信息密度: {qa_information_density:.6f}")
print(f"效率比 (QA/OCR): {efficiency_ratio:.1f}x")
print(f"\n核心结论: QA对的训练效率是OCR的{efficiency_ratio:.0f}倍")
print(f"原因是QA对聚焦于关键信息,剔除了文档中95%以上的冗余文本")
return {
'ocr_tokens': ocr_tokens,
'qa_tokens': qa_tokens,
'efficiency_ratio': efficiency_ratio,
'ocr_density': ocr_information_density,
'qa_density': qa_information_density
}
def simulate_training_curve(self,
num_epochs: int = 10,
qa_budget: int = 128000,
ocr_budget: int = 1000000) -> Dict:
"""
模拟训练曲线 - QA vs OCR
Args:
num_epochs: 训练轮数
qa_budget: QA训练预算(token数)
ocr_budget: OCR训练预算(token数)
"""
epochs = np.arange(1, num_epochs + 1)
# 模拟准确率曲线
# QA: 快速收敛,最终准确率高
qa_accuracy = 0.20 + 0.65 * (1 - np.exp(-epochs * 0.5))
# OCR: 慢速收敛,最终准确率低
ocr_accuracy = 0.15 + 0.25 * (1 - np.exp(-epochs * 0.2))
print(f"\n=== 训练曲线模拟 ===")
print(f"QA预算: {qa_budget:,} tokens")
print(f"OCR预算: {ocr_budget:,} tokens")
print(f"QA预算仅为OCR的{qa_budget/ocr_budget*100:.1f}%")
print(f"\nEpoch | QA准确率 | OCR准确率")
print("-" * 35)
for ep, qa, ocr in zip(epochs, qa_accuracy, ocr_accuracy):
print(f" {ep:2d} | {qa:.1%} | {ocr:.1%}")
print(f"\n最终提升: QA = {qa_accuracy[-1]:.1%}, OCR = {ocr_accuracy[-1]:.1%}")
print(f"QA优势: {qa_accuracy[-1] - ocr_accuracy[-1]:.1%}")
return {
'epochs': epochs.tolist(),
'qa_accuracy': qa_accuracy.tolist(),
'ocr_accuracy': ocr_accuracy.tolist()
}
analyzer = DataEfficiencyAnalyzer()
# 模拟10页文档
doc = [{'text': 'Page ' + str(i) + ' content ' * 200} for i in range(10)]
qa = [{'question': f'Q{i}', 'answer': f'A{i} with detail ' * 5} for i in range(10)]
analyzer.analyze_document(doc, qa)
analyzer.simulate_training_curve()
三、MMProLong核心技术架构
3.1 整体架构
MMProLong的核心是一个两阶段训练框架:
第一阶段:QA数据生成
┌──────────────┐ ┌──────────────┐
│ 长文档图像 │────▶│ Seed 2.0 │
│ (PDF/扫描件) │ │ (教师模型) │
└──────────────┘ └──────┬───────┘
│
▼
┌──────────────────────────────────────┐
│ 高质量QA对 │
│ - 定位型: "第3页表格第2行数据是多少" │
│ - 推理型: "根据文档判断XX趋势" │
│ - 比较型: "对比A方案和B方案" │
└──────────────────────────────────────┘
第二阶段:LMM训练
┌──────────────────────────────────────┐
│ MMProLong训练流程 │
│ │
│ [文档图像 + QA对] ──▶ Encoder │
│ │ │
│ ▼ │
│ [Vision Encoder] ──▶ [Projector] │
│ │ │
│ ▼ │
│ [LLM Backbone] ──▶ [Answer] │
│ │
│ 损失函数: QA交叉熵 + 对比学习 │
│ 训练数据: 128K tokens QA对 │
└──────────────────────────────────────┘
3.2 QA数据生成:Seed 2.0教师模型
MMProLong使用字节跳动的Seed 2.0模型作为教师,从长文档中自动生成高质量的问答对。生成策略分为三步:
class QAGenerator:
"""
QA对生成器 - 使用教师模型从长文档中提取高质量问答对
三步策略:
1. 文档分块与视觉特征提取
2. 锚点定位 - 识别文档中的关键信息区域
3. 多类型问答生成
"""
def __init__(self, teacher_model: Any = None):
self.teacher = teacher_model
self.question_templates = [
# 定位型问题
"根据文档{page}页的{section},{content}",
"{document}中关于{entity}的描述是什么",
# 推理型问题
"分析{context},{inference_question}",
"根据{evidence},{conclusion_question}",
# 比较型问题
"对比{option_a}和{option_b},{comparison_question}",
"{method_a}和{method_b}有什么区别"
]
def generate_qa_pairs(self,
document: Dict,
num_pairs: int = 50) -> List[Dict]:
"""
从文档中生成QA对
Args:
document: {pages: [image_paths], text: [page_texts],
structure: {headings, tables, figures}}
num_pairs: 生成的QA对数量
Returns:
qa_pairs: [{question, answer, page, type, difficulty}]
"""
# 步骤1: 文档结构解析
structure = self._parse_document_structure(document)
# 步骤2: 关键信息锚点定位
anchors = self._locate_information_anchors(document, structure)
# 步骤3: 多类型QA生成
qa_pairs = []
types = ['locating', 'reasoning', 'comparison']
# 分配各类型数量
counts = {
'locating': int(num_pairs * 0.4), # 40%定位型
'reasoning': int(num_pairs * 0.35), # 35%推理型
'comparison': int(num_pairs * 0.25) # 25%比较型
}
# 生成定位型问题
for anchor in anchors[:counts['locating']]:
qa = self._generate_locating_qa(document, anchor)
if qa:
qa_pairs.append(qa)
# 生成推理型问题
reasoning_contexts = self._extract_reasoning_contexts(document)
for ctx in reasoning_contexts[:counts['reasoning']]:
qa = self._generate_reasoning_qa(document, ctx)
if qa:
qa_pairs.append(qa)
# 生成比较型问题
comparison_pairs = self._find_comparable_sections(document)
for pair in comparison_pairs[:counts['comparison']]:
qa = self._generate_comparison_qa(document, pair)
if qa:
qa_pairs.append(qa)
return qa_pairs
def _parse_document_structure(self, document: Dict) -> Dict:
"""解析文档结构:标题、段落、表格、图表"""
# 使用视觉布局分析模型
structure = {
'headings': [], # 标题层级
'paragraphs': [], # 段落位置
'tables': [], # 表格位置和结构
'figures': [] # 图表位置和描述
}
for page_idx, page in enumerate(document.get('pages', [])):
# 检测页面元素
elements = self._detect_page_elements(page)
for elem in elements:
if elem['type'] == 'heading':
structure['headings'].append({
'page': page_idx,
'text': elem['text'],
'level': elem['level'],
'bbox': elem['bbox']
})
elif elem['type'] == 'table':
structure['tables'].append({
'page': page_idx,
'rows': elem['rows'],
'cols': elem['cols'],
'bbox': elem['bbox']
})
elif elem['type'] == 'figure':
structure['figures'].append({
'page': page_idx,
'caption': elem.get('caption', ''),
'type': elem.get('figure_type', 'chart'),
'bbox': elem['bbox']
})
else:
structure['paragraphs'].append({
'page': page_idx,
'text': elem['text'],
'bbox': elem['bbox']
})
return structure
def _detect_page_elements(self, page) -> List[Dict]:
"""检测页面元素(简化实现)"""
# 在实际实现中,使用视觉模型检测
# 这里返回模拟数据
return [
{'type': 'heading', 'text': 'Introduction', 'level': 1, 'bbox': (0, 0, 100, 20)},
{'type': 'paragraph', 'text': 'This is a sample paragraph...', 'bbox': (0, 25, 100, 50)},
{'type': 'table', 'rows': 5, 'cols': 3, 'bbox': (0, 55, 100, 80)},
]
def _locate_information_anchors(self,
document: Dict,
structure: Dict) -> List[Dict]:
"""定位文档中的关键信息锚点"""
anchors = []
# 基于结构的锚点
for table in structure.get('tables', []):
anchors.append({
'type': 'table',
'page': table['page'],
'content': f"table at page {table['page']}",
'importance': 0.8
})
for figure in structure.get('figures', []):
anchors.append({
'type': 'figure',
'page': figure['page'],
'content': figure.get('caption', 'figure'),
'importance': 0.7
})
# 基于文本密度的锚点(信息密集段落)
for para in structure.get('paragraphs', []):
text = para.get('text', '')
if len(text) > 50: # 长段落可能包含关键信息
# 使用关键词密度评分
keywords = ['result', 'conclusion', 'important', 'key',
'significant', 'finding', 'result', 'summary']
keyword_density = sum(
1 for kw in keywords if kw.lower() in text.lower()
) / len(text) * 1000 # 每千字关键词数
if keyword_density > 2.0:
anchors.append({
'type': 'paragraph',
'page': para['page'],
'content': text[:200],
'importance': min(1.0, keyword_density / 10.0)
})
# 按重要性排序
anchors.sort(key=lambda x: x['importance'], reverse=True)
return anchors
def _generate_locating_qa(self, document: Dict, anchor: Dict) -> Dict:
"""生成定位型问答对"""
page = anchor['page']
if anchor['type'] == 'table':
question = f"根据文档第{page + 1}页的表格,表格包含多少行多少列?"
answer = f"该表格包含{anchor.get('rows', 'N')}行{anchor.get('cols', 'N')}列"
elif anchor['type'] == 'figure':
question = f"文档第{page + 1}页的图表展示了什么内容?"
answer = anchor.get('content', '图表内容描述')
else:
question = f"文档第{page + 1}页的关键段落中提到了什么?"
answer = anchor.get('content', '')[:200]
return {
'question': question,
'answer': answer,
'page': page,
'type': 'locating',
'difficulty': 'easy',
'anchor': anchor
}
def _extract_reasoning_contexts(self, document: Dict) -> List[Dict]:
"""提取可推理的上下文片段"""
contexts = []
full_text = ' '.join(document.get('text', []))
# 找因果关系的句子
causal_patterns = ['because', 'therefore', 'leads to', 'results in',
'due to', 'as a consequence', '从而', '导致', '因此']
sentences = full_text.split('.')
for sent in sentences:
if any(p in sent.lower() for p in causal_patterns):
contexts.append({
'text': sent.strip(),
'type': 'causal',
'complexity': 'medium'
})
return contexts
def _generate_reasoning_qa(self, document: Dict, context: Dict) -> Dict:
"""生成推理型问答对"""
text = context['text']
if context['type'] == 'causal':
# 提取因果关系
parts = text.split('because')
if len(parts) == 2:
question = f"根据文档,{parts[0].strip()}的原因是什么?"
answer = parts[1].strip()
else:
question = f"根据文档,{text[:100]}...说明了什么?"
answer = text
return {
'question': question,
'answer': answer,
'type': 'reasoning',
'difficulty': 'medium',
'source_text': text
}
def _find_comparable_sections(self, document: Dict) -> List[Dict]:
"""查找文档中可比较的段落"""
pairs = []
headings = []
for page in document.get('pages', []):
text = page.get('text', '')
# 找对比性关键词
if any(kw in text.lower() for kw in ['comparison', 'vs', 'vs.',
'versus', '对比', '比较']):
headings.append({'page': text[:50], 'text': text[:200]})
# 配对
for i in range(0, len(headings) - 1, 2):
if i + 1 < len(headings):
pairs.append((headings[i], headings[i + 1]))
return pairs
class QADataSelector:
"""QA数据选择器 - 选择最有训练价值的QA对"""
def __init__(self, budget: int = 128000):
self.budget = budget # token预算
def select(self, qa_pairs: List[Dict]) -> List[Dict]:
"""
在token预算内选择最优QA对
选择策略:
1. 多样性优先(覆盖不同页面、类型)
2. 难度适中(太难或太简单都降低优先级)
3. 信息密度高(答案包含关键信息)
"""
scored_pairs = []
for qa in qa_pairs:
score = self._score_qa_pair(qa)
tokens = len(qa['question'].split()) + len(qa['answer'].split())
scored_pairs.append((score, tokens, qa))
# 按得分降序排列
scored_pairs.sort(reverse=True)
# 贪心选择
selected = []
total_tokens = 0
for score, tokens, qa in scored_pairs:
if total_tokens + tokens <= self.budget:
selected.append(qa)
total_tokens += tokens
print(f"从{len(qa_pairs)}个QA对中选出{len(selected)}个")
print(f"总token数: {total_tokens:,} / {self.budget:,}")
return selected
def _score_qa_pair(self, qa: Dict) -> float:
"""评分QA对"""
score = 0.0
# 1. 类型多样性加分
type_scores = {'locating': 0.7, 'reasoning': 1.0, 'comparison': 0.9}
score += type_scores.get(qa.get('type', 'locating'), 0.5)
# 2. 难度适中加分
difficulty = qa.get('difficulty', 'medium')
if difficulty == 'medium':
score += 0.3
elif difficulty == 'hard':
score += 0.2
else:
score += 0.1
# 3. 答案长度加分(适中长度)
answer_len = len(qa.get('answer', '').split())
if 10 <= answer_len <= 50:
score += 0.2
elif answer_len < 10:
score += 0.1
else:
score += 0.05
return score
3.3 对比学习增强
MMProLong引入了一种轻量级的对比学习机制,进一步强化模型在长文档中的定位能力:
import torch
import torch.nn as nn
import torch.nn.functional as F
class ContrastiveDocumentEncoder(nn.Module):
"""
对比文档编码器 - 增强模型在长文档中的信息定位能力
核心思想:同一文档中,答案相关的图像块应与问题embedding接近,
无关的图像块应远离
"""
def __init__(self,
vision_encoder: nn.Module,
text_encoder: nn.Module,
hidden_dim: int = 4096,
temperature: float = 0.07):
super().__init__()
self.vision_encoder = vision_encoder
self.text_encoder = text_encoder
self.temperature = temperature
# 投影头
self.vision_proj = nn.Linear(hidden_dim, hidden_dim)
self.text_proj = nn.Linear(hidden_dim, hidden_dim)
def contrastive_loss(self,
question_embeds: torch.Tensor,
document_patches: torch.Tensor,
positive_mask: torch.Tensor) -> torch.Tensor:
"""
对比学习损失
Args:
question_embeds: [batch, hidden_dim] - 问题embedding
document_patches: [batch, num_patches, hidden_dim] - 文档图像块
positive_mask: [batch, num_patches] - 正样本掩码(答案所在区域)
Returns:
loss: 对比损失
"""
batch_size, num_patches, _ = document_patches.shape
# 投影
q = self.text_proj(question_embeds) # [batch, hidden_dim]
q = F.normalize(q, dim=-1)
p = self.vision_proj(document_patches) # [batch, num_patches, hidden_dim]
p = F.normalize(p, dim=-1)
# 相似度矩阵
sim = torch.matmul(q.unsqueeze(1), p.transpose(-2, -1)) # [batch, 1, num_patches]
sim = sim.squeeze(1) / self.temperature # [batch, num_patches]
# 正样本:答案所在区域
# 负样本:其他区域
positive_mask = positive_mask.float()
num_positives = positive_mask.sum(dim=1, keepdim=True).clamp(min=1)
# InfoNCE损失
exp_sim = torch.exp(sim)
pos_exp = (exp_sim * positive_mask).sum(dim=1) / num_positives.squeeze()
neg_exp = (exp_sim * (1 - positive_mask)).sum(dim=1)
loss = -torch.log(pos_exp / (pos_exp + neg_exp + 1e-8)).mean()
return loss
def forward(self,
questions: List[str],
document_images: torch.Tensor,
answer_regions: torch.Tensor = None) -> Dict:
"""
Args:
questions: 问题文本列表
document_images: [batch, num_pages, 3, H, W] 文档图像
answer_regions: [batch, num_patches] 答案区域标注
Returns:
outputs: 包含loss和embeddings
"""
# 编码问题
question_embeds = self.text_encoder(questions)
# 编码文档图像
batch_size, num_pages, C, H, W = document_images.shape
doc_flat = document_images.view(-1, C, H, W)
patch_embeds = self.vision_encoder(doc_flat)
# [batch * num_pages, num_patches_per_page, hidden_dim]
# reshape
_, num_patches_per_page, hidden_dim = patch_embeds.shape
document_patches = patch_embeds.view(
batch_size, num_pages * num_patches_per_page, hidden_dim
)
outputs = {'question_embeds': question_embeds,
'document_patches': document_patches}
if answer_regions is not None:
loss = self.contrastive_loss(
question_embeds, document_patches, answer_regions
)
outputs['contrastive_loss'] = loss
return outputs
四、训练策略与数据配比
4.1 多阶段训练
MMProLong采用三阶段渐进式训练策略:
class MMProLongTrainer:
"""
MMProLong训练器 - 三阶段渐进式训练
阶段1: 短上下文适应 (≤32K tokens)
阶段2: 长上下文扩展 (≤128K tokens)
阶段3: 超长上下文强化 (≤512K tokens)
"""
def __init__(self,
model: nn.Module,
max_contexts: List[int] = [32768, 131072, 524288],
qa_ratios: List[float] = [0.3, 0.5, 0.7]):
self.model = model
self.max_contexts = max_contexts
self.qa_ratios = qa_ratios
self.optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)
def train_stage(self,
stage: int,
train_data: List[Dict],
num_epochs: int = 3) -> Dict:
"""
训练单个阶段
Args:
stage: 阶段编号 (0, 1, 2)
train_data: 训练数据,包含文档和QA对
num_epochs: 训练轮数
"""
max_ctx = self.max_contexts[stage]
qa_ratio = self.qa_ratios[stage]
print(f"=== 阶段{stage + 1}: 最大上下文={max_ctx:,}, QA比例={qa_ratio:.0%} ===")
metrics = {'loss': [], 'accuracy': []}
for epoch in range(num_epochs):
epoch_loss = 0.0
num_batches = 0
for batch in self._create_batches(train_data, max_ctx, qa_ratio):
# 前向传播
outputs = self.model(
questions=batch['questions'],
document_images=batch['images'],
answer_regions=batch.get('answer_regions')
)
# 计算损失
loss = outputs.get('contrastive_loss', 0)
if 'generation_loss' in outputs:
loss = loss + outputs['generation_loss']
# 反向传播
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
self.optimizer.step()
self.optimizer.zero_grad()
epoch_loss += loss.item()
num_batches += 1
avg_loss = epoch_loss / max(num_batches, 1)
metrics['loss'].append(avg_loss)
# 验证
acc = self._evaluate(stage)
metrics['accuracy'].append(acc)
print(f"Epoch {epoch + 1}: loss={avg_loss:.4f}, accuracy={acc:.2%}")
return metrics
def _create_batches(self,
data: List[Dict],
max_ctx: int,
qa_ratio: float) -> List[Dict]:
"""创建训练批次,按比例混合QA和OCR数据"""
batches = []
# 分离QA和OCR数据
qa_data = [d for d in data if d.get('type') == 'qa']
ocr_data = [d for d in data if d.get('type') == 'ocr']
# 按比例采样
batch_size = 8
num_qa = int(batch_size * qa_ratio)
num_ocr = batch_size - num_qa
for i in range(0, max(len(qa_data), 1), num_qa):
batch_qa = qa_data[i:i + num_qa]
batch_ocr = ocr_data[i:i + num_ocr] if ocr_data else []
if batch_qa:
batches.append({
'questions': [d['question'] for d in batch_qa],
'images': torch.stack([d['image'] for d in batch_qa]),
'answer_regions': torch.stack([d['region'] for d in batch_qa])
if 'region' in batch_qa[0] else None,
'type': 'mixed'
})
return batches
def _evaluate(self, stage: int) -> float:
"""验证当前阶段性能"""
# 使用验证集评估
# 这里返回模拟值
return 0.75 + stage * 0.08
4.2 关键发现
QA数据在低预算下优势更明显:仅使用128K token的QA数据训练,即可在512K token的长文档检索任务上超越使用1M+ token OCR数据的模型。
跨模型迁移性:在Qwen3-VL-8B上使用MMProLong方法训练的模型,同样观察到了明显提升,证明该方法具有架构无关性。
视频理解的正向迁移:QA训练带来的"聚焦关键信息"能力,自动迁移到了视频理解任务——模型在长视频问答中表现出更强的上下文定位能力,尽管从未使用视频进行训练。
五、实验结果与分析
5.1 长文档检索基准
| 模型 | 训练数据 | 预算 | 128K准确率 | 256K准确率 | 512K准确率 |
|---|---|---|---|---|---|
| InternVL3-38B | OCR | 2M tokens | 72.3% | 65.1% | 51.2% |
| Gemma3-27B | OCR | 2M tokens | 68.7% | 60.3% | 47.8% |
| MMProLong (ours) | QA | 128K tokens | 78.5% | 73.2% | 64.6% |
| MMProLong (ours) | QA+OCR | 256K tokens | 81.3% | 76.8% | 68.1% |
5.2 视频理解迁移效果
| 模型 | 视频问答准确率 | 时序定位F1 | 描述一致性 |
|---|---|---|---|
| 基线 (OCR训练) | 52.3% | 0.45 | 0.61 |
| MMProLong (QA训练) | 61.7% | 0.53 | 0.72 |
| 提升 | +9.4% | +0.08 | +0.11 |
六、工程实践指南
6.1 数据准备
def prepare_mmprolong_data(
document_dir: str,
output_dir: str,
teacher_model: str = "Seed2.0",
budget: int = 128000
):
"""
准备MMProLong训练数据
流程:
1. 解析文档
2. 生成QA对
3. 质量筛选
4. 格式化输出
"""
import json
import os
generator = QAGenerator()
selector = QADataSelector(budget=budget)
all_qa_pairs = []
for doc_file in os.listdir(document_dir):
if not doc_file.endswith(('.pdf', '.png', '.jpg')):
continue
doc_path = os.path.join(document_dir, doc_file)
doc = load_document(doc_path)
# 生成QA对
qa_pairs = generator.generate_qa_pairs(doc, num_pairs=100)
all_qa_pairs.extend(qa_pairs)
# 质量筛选
selected = selector.select(all_qa_pairs)
# 格式化输出
output = []
for qa in selected:
output.append({
'id': f"qa_{len(output)}",
'conversations': [
{'from': 'human', 'value': qa['question']},
{'from': 'gpt', 'value': qa['answer']}
],
'page': qa.get('page', 0),
'type': qa.get('type', 'locating')
})
# 保存
os.makedirs(output_dir, exist_ok=True)
with open(os.path.join(output_dir, 'mmprolong_data.json'), 'w') as f:
json.dump(output, f, ensure_ascii=False, indent=2)
print(f"生成{len(output)}个QA对,共{sum(len(q['conversations'][0]['value'].split()) + len(q['conversations'][1]['value'].split()) for q in output)} tokens")
def load_document(path: str) -> Dict:
"""加载文档(简化实现)"""
return {
'pages': [{'text': 'Sample document content for testing...'}],
'text': ['Sample document content for testing...']
}
七、总结
MMProLong的核心贡献在于揭示了长文档LMM训练的"数据效率"问题:
- 数据组织比数据规模更重要:128K QA token > 1M OCR token
- 问答对迫使模型学会"聚焦":在长上下文中定位关键信息,而非逐字重建
- 跨模态迁移是额外收益:长文档的定位能力自然迁移到视频理解
这项研究为长上下文LMM训练提供了一条更经济、更高效的路径——在算力和数据都受限的现实条件下,数据质量比数据规模更重要。
参考:ByteDance Seed Team & HKUST, “MMProLong: Long Document LMM Training with QA Pairs”, 2026.