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的核心挑战

多模态大语言模型在处理长文档时面临三个相互关联的挑战:

  1. 视觉-文本对齐失效:OCR将视觉布局信息压缩为纯文本,丢失了段落位置、表格结构、图表纹理等关键视觉线索
  2. 注意力稀释:在长上下文中,关键信息被大量无关文本稀释,模型难以定位答案
  3. 训练-推理分布偏移: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 关键发现

  1. QA数据在低预算下优势更明显:仅使用128K token的QA数据训练,即可在512K token的长文档检索任务上超越使用1M+ token OCR数据的模型。

  2. 跨模型迁移性:在Qwen3-VL-8B上使用MMProLong方法训练的模型,同样观察到了明显提升,证明该方法具有架构无关性。

  3. 视频理解的正向迁移:QA训练带来的"聚焦关键信息"能力,自动迁移到了视频理解任务——模型在长视频问答中表现出更强的上下文定位能力,尽管从未使用视频进行训练。


五、实验结果与分析

5.1 长文档检索基准

模型训练数据预算128K准确率256K准确率512K准确率
InternVL3-38BOCR2M tokens72.3%65.1%51.2%
Gemma3-27BOCR2M tokens68.7%60.3%47.8%
MMProLong (ours)QA128K tokens78.5%73.2%64.6%
MMProLong (ours)QA+OCR256K tokens81.3%76.8%68.1%

5.2 视频理解迁移效果

模型视频问答准确率时序定位F1描述一致性
基线 (OCR训练)52.3%0.450.61
MMProLong (QA训练)61.7%0.530.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训练的"数据效率"问题

  1. 数据组织比数据规模更重要:128K QA token > 1M OCR token
  2. 问答对迫使模型学会"聚焦":在长上下文中定位关键信息,而非逐字重建
  3. 跨模态迁移是额外收益:长文档的定位能力自然迁移到视频理解

这项研究为长上下文LMM训练提供了一条更经济、更高效的路径——在算力和数据都受限的现实条件下,数据质量比数据规模更重要


参考:ByteDance Seed Team & HKUST, “MMProLong: Long Document LMM Training with QA Pairs”, 2026.