Nature Biomedical Engineering IF=26.7 | OVFM:面向眼科手术识别与导航的视频基础模型

引言

眼科手术的精准度与安全性高度依赖外科医生的经验与技巧,但顶尖专家的培养周期漫长,且术中操作的实时评估与导航支持仍面临挑战。能否让AI系统“看懂”手术视频,为医生提供智能化的实时辅助?

近日,一项发表于《Nature Biomedical Engineering》的研究带来了突破。由上海交通大学、上海交通大学医学院附属新华医院等团队联合开发了一款眼科手术视频基础模型OVFM。该模型通过自监督学习,从涵盖144种术式、超过110万个视频片段的大规模数据中,掌握了眼科手术的时空动态特征。研究团队进一步通过知识蒸馏将其轻量化,并集成到手术显微镜中,构建了一套实时手术导航系统。在猪眼湿实验室实验中,该系统有效提升了不同经验水平外科医生的手术表现,缩小了技能差距,展现出推动术中智能辅助应用的巨大潜力。

基本信息

文章标题:An ophthalmic video foundation model for surgical recognition and navigation with wet-lab porcine eye validation
期刊:Nature Biomedical Engineering
影响因子:26.7
发表时间:2026年1月23日
研究单位:1.上海交通大学机械工程学院生物医学制造与生命质量工程研究所,中国上海;2.上海交通大学医学院附属新华医院眼科,中国上海;3.汕头大学·香港中文大学联合汕头国际眼科中心,汕头大学医学院,中国汕头;4.上海微创医疗器械(集团)有限公司,中国上海;5.上海爱尔眼科医院,中国上海;6.上海爱尔眼科研究所,中国上海;7.焦作尖峰眼科医院,中国焦作;8.甘孜州人民医院康巴眼科中心,中国康定;9.台州爱尔眼科医院,中国台州;10.广州医科大学附属广州市第八人民医院,中国广州;11.湖州爱尔眼科医院,中国湖州;12.上海交通大学医疗机器人研究院,中国上海
Github地址:https://github.com/puxuntu/OVFM
论文地址:https://doi.org/10.1038/s41551-026-01622-w
算力描述:模型训练使用NVIDIA GeForce GTX 4090 GPUs

研究内容与方法

1. 大规模眼科手术视频数据集构建

  • 数据集组成:整合7个多中心内部数据集与公开Cataract-1K数据集,覆盖144种眼科手术类型,包含前节、后节及联合手术视频
  • 预处理流程:
    • 视频压缩与尺寸调整:将原始视频resize为原分辨率的1/4,降低存储与计算开销
    • 稀疏采样生成片段:对长视频进行稀疏采样,生成1.1M个固定时长的视频片段,每个片段统一resize至480×270
    • 多尺度时空视图生成:为自监督训练准备不同尺度的视图,对应代码片段(来自dataset/pretrain_dataset.py):
      def generate_global_views(self, frames):
          # 生成2种全局视图:8帧和16帧,统一缩放至224×224
          g1_indices = np.linspace(0, len(frames)-1, 8, dtype=int)
          g1 = np.array([cv2.resize(f, (224,224)) for f in frames[g1_indices]])
          g2_indices = np.linspace(0, len(frames)-1, 16, dtype=int)
          g2 = np.array([cv2.resize(f, (224,224)) for f in frames[g2_indices]])
          return g1, g2
      
      def generate_local_views(self, frames):
          # 生成8种局部视图:每个子片段取1/8时长,随机采样2/4/6/8帧并裁剪至96×96
          local_views = []
          seg_len = len(frames) // 8
          for i in range(8):
              start, end = i*seg_len, (i+1)*seg_len if i<7 else len(frames)
              seg_frames = frames[start:end]
              num_frames = np.random.choice([2,4,6,8])
              indices = np.linspace(0, len(seg_frames)-1, num_frames, dtype=int)
              local_frames = seg_frames[indices]
              # 随机空间裁剪
              h, w = local_frames[0].shape[:2]
              top, left = np.random.randint(0, h-96+1), np.random.randint(0, w-96+1)
              local_frames = np.array([f[top:top+96, left:left+96] for f in local_frames])
              local_views.append(local_frames)
          return local_views
      

【数据集地理分布示意图】在这里插入图片描述

2. 自监督视频Transformer(OVFM)核心架构设计

  • 模块组成:Patch嵌入层、时空Transformer块师生自监督训练框架
  • Patch嵌入层:将单帧图像划分为16×16的patch,转换为768维token,添加空间位置编码与时间位置编码,同时拼接可学习的[CLS] token以捕获全局特征,对应代码片段(来自models/ovfm.py):
    class PatchEmbed(nn.Module):
        def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
            super().__init__()
            self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
            self.num_patches = (img_size//patch_size) **2
            self.spatial_pos_embed = nn.Parameter(torch.randn(1, self.num_patches, embed_dim))
            self.temporal_pos_embed = nn.Parameter(torch.randn(1, 16, embed_dim))  # 最大支持16帧
    
        def forward(self, x):
            # x: (B, T, C, H, W)
            B, T, C, H, W = x.shape
            x = x.flatten(0,1)  # (B*T, C, H, W)
            x = self.proj(x).flatten(2).transpose(1,2)  # (B*T, num_patches, embed_dim)
            x = x + self.spatial_pos_embed  # 添加空间位置编码
            x = x.unflatten(0, (B, T))  # (B, T, num_patches, embed_dim)
            x = x + self.temporal_pos_embed[:, :T, :].unsqueeze(2)  # 添加时间位置编码
            # 拼接[CLS] token
            cls_token = nn.Parameter(torch.randn(1,1,1,embed_dim)).repeat(B,1,1,1)
            x = torch.cat([cls_token, x], dim=2)  # (B, T, num_patches+1, embed_dim)
            return x
    
  • 时空Transformer块:采用分治式时空注意力,先对每个patch做跨帧时间自注意力,再对每帧做空间自注意力,结合残差连接与层归一化,对应代码片段(来自models/ovfm.py):
    class VideoTransformerBlock(nn.Module):
        def __init__(self, dim, num_heads, mlp_ratio=4.):
            super().__init__()
            self.norm1 = nn.LayerNorm(dim)
            self.temporal_attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)
            self.norm2 = nn.LayerNorm(dim)
            self.spatial_attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)
            self.norm3 = nn.LayerNorm(dim)
            self.mlp = nn.Sequential(
                nn.Linear(dim, int(dim*mlp_ratio)),
                nn.GELU(),
                nn.Linear(int(dim*mlp_ratio), dim)
            )
    
        def forward(self, x):
            B, T, N, C = x.shape  # N=num_patches+1(含[CLS])
            # 时间自注意力:每个patch跨帧计算注意力
            x_temporal = x.permute(0,2,1,3).flatten(0,1)  # (B*N, T, C)
            attn_out, _ = self.temporal_attn(self.norm1(x_temporal), self.norm1(x_temporal), self.norm1(x_temporal))
            x_temporal = x_temporal + attn_out  # 残差连接
            x_temporal = x_temporal.unflatten(0, (B, N)).permute(0,2,1,3)  # 恢复原形状
    
            # 空间自注意力:每帧内的patch计算注意力
            x_spatial = x_temporal.flatten(0,1)  # (B*T, N, C)
            attn_out, _ = self.spatial_attn(self.norm2(x_spatial), self.norm2(x_spatial), self.norm2(x_spatial))
            x_spatial = x_spatial + attn_out  # 残差连接
            x_spatial = x_spatial.unflatten(0, (B, T))  # 恢复原形状
    
            # MLP层
            x = x_spatial + self.mlp(self.norm3(x_spatial))
            return x
    
  • 自监督训练框架:采用师生模型结构,教师模型通过动量更新同步学生模型权重,损失由全局-全局匹配损失与局部-全局匹配损失组成:
    • 全局-全局匹配损失:
      Lgg=−12(sim(fsg1,ftg2)+sim(fsg2,ftg1))\mathcal{L}_{gg} = -\frac{1}{2}\left(\text{sim}(f_s^{g1}, f_t^{g2}) + \text{sim}(f_s^{g2}, f_t^{g1})\right)Lgg=21(sim(fsg1,ftg2)+sim(fsg2,ftg1))
    • 局部-全局匹配损失:
      Llg=−18∑i=18sim(fsli,12(ftg1+ftg2))\mathcal{L}_{lg} = -\frac{1}{8}\sum_{i=1}^8 \text{sim}(f_s^{l_i}, \frac{1}{2}(f_t^{g1}+f_t^{g2}))Llg=81i=18sim(fsli,21(ftg1+ftg2))
    • 总损失:
      L=Lgg+Llg\mathcal{L} = \mathcal{L}_{gg} + \mathcal{L}_{lg}L=Lgg+Llg
      其中sim(a,b)=ecos(a,b)/τ∑kecos(a,bk)/τ\text{sim}(a,b) = \frac{e^{\text{cos}(a,b)/\tau}}{\sum_{k}e^{\text{cos}(a,b_k)/\tau}}sim(a,b)=kecos(a,bk)/τecos(a,b)/τ为温度系数τ\tauτ调控的余弦相似度归一化
      对应代码片段(来自models/ovfm.py):
    class TeacherStudentTrainer(nn.Module):
        def __init__(self, student, teacher, tau=0.07, momentum=0.996):
            super().__init__()
            self.student = student
            self.teacher = teacher
            self.tau = tau
            self.momentum = momentum
            # 教师模型权重初始化与冻结
            for param_q, param_k in zip(student.parameters(), teacher.parameters()):
                param_k.data.copy_(param_q.data)
                param_k.requires_grad = False
    
        @torch.no_grad()
        def update_teacher(self):
            # 动量更新教师模型权重
            for param_q, param_k in zip(self.student.parameters(), self.teacher.parameters()):
                param_k.data = param_k.data*self.momentum + param_q.data*(1-self.momentum)
    
        def compute_nce_loss(self, s_feat, t_feat):
            # 计算对比损失
            s_feat = F.normalize(s_feat, dim=-1)
            t_feat = F.normalize(t_feat, dim=-1)
            sim = torch.matmul(s_feat, t_feat.T)/self.tau
            label = torch.arange(sim.shape[0], device=sim.device)
            loss = (F.cross_entropy(sim, label) + F.cross_entropy(sim.T, label))/2
            return loss
    
        def forward(self, global_views, local_views):
            # 学生模型前向传播
            s_g1 = self.student(global_views[0])[:,0,:]  # 取[CLS]特征
            s_g2 = self.student(global_views[1])[:,0,:]
            s_local = [self.student(lv)[:,0,:] for lv in local_views]
            # 教师模型前向传播(无梯度)
            with torch.no_grad():
                t_g1 = self.teacher(global_views[0])[:,0,:]
                t_g2 = self.teacher(global_views[1])[:,0,:]
            # 计算损失
            loss_gg = self.compute_nce_loss(torch.cat([s_g1, s_g2]), torch.cat([t_g2, t_g1]))
            t_avg = (t_g1 + t_g2)/2
            loss_lg = sum([self.compute_nce_loss(sl, t_avg) for sl in s_local])/len(s_local)
            total_loss = loss_gg + loss_lg
            # 更新教师模型
            self.update_teacher()
            return total_loss
    

【OVFM自监督训练框架示意图】在这里插入图片描述

3. 通用到特定的两阶段知识蒸馏

  • 阶段1:通用知识蒸馏:将大尺寸OVFM的通用时空特征迁移至小尺寸模型(OVFM-small/tiny),复用自监督训练的对比损失匹配师生模型的特征分布,对应代码片段(来自models/distillation.py):
    class GeneralDistiller(nn.Module):
        def __init__(self, teacher, student, tau=0.07):
            super().__init__()
            self.teacher = teacher
            self.student = student
            self.tau = tau
            for param in teacher.parameters():
                param.requires_grad = False
    
        def compute_nce_loss(self, s_feat, t_feat):
            s_feat = F.normalize(s_feat, dim=-1)
            t_feat = F.normalize(t_feat, dim=-1)
            sim = torch.matmul(s_feat, t_feat.T)/self.tau
            label = torch.arange(sim.shape[0], device=sim.device)
            return (F.cross_entropy(sim, label) + F.cross_entropy(sim.T, label))/2
    
        def forward(self, global_views, local_views):
            with torch.no_grad():
                t_g1 = self.teacher(global_views[0])[:,0,:]
                t_g2 = self.teacher(global_views[1])[:,0,:]
                t_local = [self.teacher(lv)[:,0,:] for lv in local_views]
            s_g1 = self.student(global_views[0])[:,0,:]
            s_g2 = self.student(global_views[1])[:,0,:]
            s_local = [self.student(lv)[:,0,:] for lv in local_views]
            # 计算通用蒸馏损失
            loss_gg = self.compute_nce_loss(torch.cat([s_g1, s_g2]), torch.cat([t_g2, t_g1]))
            t_avg = (t_g1 + t_g2)/2
            loss_lg = sum([self.compute_nce_loss(sl, t_avg) for sl in s_local])/len(s_local)
            return loss_gg + loss_lg
    
  • 阶段2:任务特定知识蒸馏:针对下游任务,将大模型的任务相关特征与输出迁移至小模型,损失由任务损失与蒸馏损失加权组成:
    Ldistill=Ltask+λ⋅Lkd\mathcal{L}_{distill} = \mathcal{L}_{task} + \lambda \cdot \mathcal{L}_{kd}Ldistill=Ltask+λLkd
    其中Lkd\mathcal{L}_{kd}Lkd为师生模型特征的MSE损失,对应代码片段(来自models/distillation.py):
    class TaskSpecificDistiller(nn.Module):
        def __init__(self, teacher, student, task_head, lambda_kd=0.5):
            super().__init__()
            self.teacher = teacher
            self.student = student
            self.task_head = task_head
            self.lambda_kd = lambda_kd
            for param in teacher.parameters():
                param.requires_grad = False
    
        def forward(self, x, labels):
            with torch.no_grad():
                t_feat = self.teacher(x)[:,0,:]
                t_logits = self.task_head(t_feat)
            s_feat = self.student(x)[:,0,:]
            s_logits = self.task_head(s_feat)
            # 任务损失+蒸馏损失
            task_loss = F.cross_entropy(s_logits, labels)
            kd_loss = F.mse_loss(s_feat, t_feat)
            return task_loss + self.lambda_kd * kd_loss
    

【两阶段知识蒸馏框架示意图】在这里插入图片描述

4. 下游任务适配模块设计

  • 时空类任务(手术步骤识别、工具存在识别、并发症检测、手术技能评估):直接提取OVFM输出的[CLS] token特征,连接线性分类层完成任务,对应代码片段(来自models/downstream/step_recognition.py):
    class StepRecognitionHead(nn.Module):
        def __init__(self, in_dim=768, num_classes=4):
            super().__init__()
            self.fc = nn.Linear(in_dim, num_classes)
    
        def forward(self, feat):
            # feat为OVFM输出的[CLS]特征:(B, in_dim)
            return self.fc(feat)
    
  • 空间类任务(手术场景分割、Limbus边界分割、核定位):将OVFM输出的patch特征输入转置卷积解码器,恢复为与输入图像同尺寸的分割/定位结果,对应代码片段(来自models/downstream/limbus_segmentation.py):
    class LimbusSegmentationDecoder(nn.Module):
        def __init__(self, in_dim=768, out_channels=1):
            super().__init__()
            self.decoder = nn.Sequential(
                nn.ConvTranspose2d(in_dim, 256, kernel_size=2, stride=2),
                nn.ReLU(),
                nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),
                nn.ReLU(),
                nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2),
                nn.ReLU(),
                nn.Conv2d(64, out_channels, kernel_size=1)
            )
    
        def forward(self, feat):
            # feat: (B, T, N, C),取最后一帧的patch特征(排除[CLS])
            B, T, N, C = feat.shape
            patch_h = patch_w = int(math.sqrt(N-1))
            patch_feat = feat[:, -1, 1:, :].transpose(1,2).view(B, C, patch_h, patch_w)
            seg_map = self.decoder(patch_feat)
            # 上采样至原图像尺寸
            seg_map = F.interpolate(seg_map, size=(224,224), mode='bilinear', align_corners=False)
            return seg_map
    

5. 实时眼科手术导航系统集成

  • 系统架构:将蒸馏后的轻量OVFM嵌入手术显微镜的集成处理器,实时采集术中视频流,并行执行手术步骤识别与解剖结构分割
  • 导航逻辑:根据识别出的当前手术步骤,自动切换导航输出内容(如切口引导线、撕囊范围定位圈),通过光路投影至手术视野
    【导航系统实验装置示意图】在这里插入图片描述

实验结果分析

OVFM在眼科手术下游任务中的卓越性能

以下图表展示了OVFM模型在七个下游任务中的表现,包括手术步骤识别、器械存在识别、并发症检测、手术技能评估、角膜缘边界分割、手术场景分割和晶状体核块定位。OVFM在各项任务中均显著优于其他基础模型。
在这里插入图片描述

  • 手术步骤识别:在Cataract-101数据集上,OVFM在切口、撕囊、人工晶体植入等步骤的识别中,均取得了最高的AUC值(例如,切口步骤AUC=0.992),其ROC曲线始终位于其他模型之上,显示出精准的时序动作理解能力。
  • 器械存在识别:在CATARACTS数据集上,OVFM的微平均AUC达到0.985,在22种器械中的21种上取得了最佳性能,表明其能有效捕捉精细的器械运动特征。
  • 空间任务表现:在角膜缘边界分割(Dice分数达0.960)和手术场景分割(背景、瞳孔、角膜、晶状体等类别的Dice分数均最高)任务中,OVFM同样展现出卓越的空间定位与分割能力,显著优于其他对比模型。

知识蒸馏在模型效率与精度间的平衡

为满足术中实时部署需求,研究采用两阶段知识蒸馏策略压缩模型。结果显示,蒸馏后的模型在显著减小参数量的同时,仍能保留原始模型绝大部分性能。
在这里插入图片描述

  • 参数压缩与性能保留:使用SVT-small架构的蒸馏模型,参数量减少至原模型的约1/3(36.2M vs 121.3M),但在手术步骤识别任务中仍能保留原模型99%以上的AUC性能。更小的SVT-tiny模型(7.7M参数)也能保留90%以上的性能。
  • 两阶段蒸馏策略的有效性:在手术步骤识别和角膜缘分割任务中,采用的两阶段(通用到特定)蒸馏策略,其性能均优于单阶段蒸馏或直接训练小模型的方法,证明了该策略在迁移通用时空知识方面的优势。

OVFM导航系统提升手术技能并缩小经验差距

将蒸馏后的OVFM集成到手术显微镜中,构建实时导航系统,并在猪眼湿实验室中由10名外科医生进行交叉用户研究。结果显示,该系统能有效提升手术操作精度,并缩小新手与专家医生之间的技能差距
在这里插入图片描述

  • 提升手术操作精度:在导航辅助下,主切口角度误差、次切口角度误差和撕囊中心定位误差均显著降低,而撕囊形状匹配度显著提高,表明导航提供了有效的实时空间引导。
  • 缩小外科医生经验差距:效应大小分析显示,导航系统对新手外科医生的提升效果尤为明显。在撕囊中心定位误差和形状匹配度等关键指标上,新手组从导航中获得的改善显著大于专家组,表明该系统有助于拉平不同经验水平外科医生的操作表现。

优势与局限

优势

大规模眼科手术视频预训练:基于包含144种手术类型、110万视频片段的数据集,模型学习到丰富的时空运动特征,在七项下游任务中均超越现有基准。
实时部署能力:通过两阶段知识蒸馏策略,模型参数量大幅减少(最高压缩15.8倍)而性能保留度高(如步骤识别AUC保留90.7%以上),可在手术显微镜处理单元实现约19.4 FPS的实时推理
临床验证有效:在湿实验室猪眼白内障手术的用户研究中,导航系统显著提升了外科医生的手术表现(如切口角度误差降低),并缩小了新手与专家之间的技能差距。

局限

数据代表性可能不足:尽管数据集规模大,但后节手术视频相对较少,且视频下采样可能损失细微的时间模式,影响模型在复杂或罕见术式中的泛化能力。
真实临床环境验证有限:模型主要在回顾性视频和猪眼实验中验证,尚未在多样化的真实手术场景中进行广泛测试,其在不同医院设备、患者群体中的鲁棒性有待进一步评估。
伦理与责任问题未解决:作为术中决策辅助系统,其决策透明度、医生责任界定以及临床集成所涉及的伦理挑战尚未被充分探讨。

参考文献

  1. Self-supervised video transformer Ranasinghe et al., 2022:该论文提出了自监督视频Transformer(SVT)架构,是本研究构建眼科视频基础模型(OVFM)的核心网络结构。研究者基于该架构,通过预测不同时空视图之间的对应关系,使模型能够从大规模无标签眼科手术视频中学习高质量的时空运动特征。
  2. Foundation model for endoscopy video analysis via large-scale self-supervised pre-train Wang et al., 2023:本文提出了用于内窥镜视频分析的基础模型(Endo FM),为医学视频领域的自监督预训练提供了重要参考。本研究在构建OVFM时借鉴了其思路,并在多个下游任务上与之进行了性能对比,证明了领域专用预训练的优势。
  3. Masked autoencoders are scalable vision learners He et al., 2022:该论文提出了掩码自编码器(MAE)方法,展示了在视觉任务上通过自监督学习可扩展表征的能力。本研究将其视频版本(SSL-Video MAE)作为基线模型之一,OVFM在各项任务上均显著优于该模型,突显了针对手术视频动态特性设计的预训练目标的有效性。
  4. Cataract-1K dataset for deep-learning-assisted analysis of cataract surgery videos Ghamsarian et al., 2024:该论文公开了Cataract-1K数据集,这是一个用于白内障手术视频分析的大规模数据集。本研究将其纳入预训练数据,并用于并发症检测、手术场景分割等下游任务的评估,为模型训练与验证提供了关键数据支持。
  5. Knowledge distillation: a survey Gou et al., 2021:该综述系统总结了知识蒸馏的各种方法与应用。本研究受其启发,采用了一种通用的两阶段(通用到特定)知识蒸馏策略,成功将大型OVFM压缩为轻量级模型,从而实现了在手术显微镜处理单元上的实时部署。
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值