You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用graph_tool更快构建3D seam carving的大型图结构?

问题描述

我正在Python中复现3D seam carving算法(出自论文《Video Retargeting with Optimal Seam Carving》),使用mask_data这个3D数组标记需保留的体素区域(测试时用全0数组,20%概率设为1)。算法核心是构建图并通过最小割确定待移除的seam,但当前实现构建128³规模的图在M2 Max MacBook上耗时13分钟。尝试并行化后耗时反而更长,希望继续使用graph_tool库(其C++实现的最大流/最小割计算速度快),求加快图结构创建的思路或实现方案。

注:论文提到的多分辨率分带可提速,但目前仅图构建环节缓慢;当前版本仅支持'left'方向;曾尝试拆分mask_data构建子图再合并,但重组开销大导致速度更慢。

当前实现代码:

import graph_tool.all as gt

def create_directed_energy_graph_from_mask(mask_data, direction='left', large_weight=1e8):
    z,y,x = mask_data.shape  # Dimensions of the 3D mask array
    g = gt.Graph(directed=True)
    weight_prop = g.new_edge_property("int")  # Edge property for weights

    # Function to get linear index from 3D coordinates, assuming C-style row-major order
    index = lambda i, j, k: i * y * x + j * x + k

    # Add vertices
    num_vertices = z * y * x
    g.add_vertex(num_vertices)
    vertices = list(g.vertices())  # Get the list of vertices

    # Define neighbor offsets based on directionality
    # Source to sink propagation direction - positive axes directions
    directions = {
        'left': [(0, 0, 1)],  # propagate right
        'right': [(0, 0, -1)],  # propagate left
        'top': [(0, 1, 0)],  # propagate downwards
        'bottom': [(0, -1, 0)],  # propagate upwards
        'front': [(1, 0, 0)],  # propagate back
        'back': [(-1, 0, 0)]  # propagate front
    }

    neighbors = directions[direction]
    print(x,y,z)
    
    for i in range(z):
        for j in range(y):
            for k in range(x):
                current_index = index(i, j, k)
                current_vertex = vertices[current_index]

                # Check each neighbor direction for valid connections
                for di, dj, dk in neighbors:
                    ni, nj, nk = i + di, j + dj, k + dk
                    if 0 <= ni < z and 0 <= nj < y and 0 <= nk < x:
                        neighbor_index = index(ni, nj, nk)
                        neighbor_vertex = vertices[neighbor_index]

                        # Determine edge weight
                        weight = 10 if mask_data[ni, nj, nk] != 0 or mask_data[i, j, k] != 0 else 1

                        # Add edge and assign weight
                        e = g.add_edge(current_vertex, neighbor_vertex) #forward edge with energy value
                        e2 = g.add_edge(neighbor_vertex, current_vertex) #backward edge with large energy value
                        weight_prop[e] = weight
                        weight_prop[e2] = large_weight

                # Add each diagonal backwards neighbor inf edge, ie x-1, y-1 and x-1, y+1 for YX plane
                if k > 0 and j > 0:
                    neighbor_index = index(i, j-1, k-1)
                    neighbor_vertex = vertices[neighbor_index]
                    e = g.add_edge(current_vertex, neighbor_vertex)
                    weight_prop[e] = large_weight+1
                if k > 0 and j < y-1:
                    neighbor_index = index(i, j+1, k-1)
                    neighbor_vertex = vertices[neighbor_index]
                    e = g.add_edge(current_vertex, neighbor_vertex)
                    weight_prop[e] = large_weight+1
                    
                # Add each diagonal backwards neighbor inf edge for ik plane
                if k > 0 and i > 0:
                    neighbor_index = index(i-1, j, k-1)
                    neighbor_vertex = vertices[neighbor_index]
                    e = g.add_edge(current_vertex, neighbor_vertex)
                    weight_prop[e] = large_weight+2
                if k > 0 and i < z-1:
                    neighbor_index = index(i+1, j, k-1)
                    neighbor_vertex = vertices[neighbor_index]
                    e = g.add_edge(current_vertex, neighbor_vertex)
                    weight_prop[e] = large_weight+2

    g.edge_properties["weight"] = weight_prop
    return g, weight_prop
优化思路与实现方案

核心优化方向:减少Python与graph_tool C++层的交互开销

原代码的主要瓶颈是Python三层嵌套循环中频繁调用g.add_edge和设置边属性,每次调用都有跨语言交互的开销。以下是针对性优化:


1. 预生成所有边的源/目标索引与权重,批量添加

graph_tool支持通过数组批量添加边,避免循环内逐个调用API。具体步骤:

  • 预先收集所有需要添加的边的源顶点索引、目标顶点索引、对应的权重
  • 使用g.add_edge_list一次性添加所有边
  • 批量设置边属性,替代逐个赋值

2. 用numpy预处理mask数据,减少循环内条件判断

将mask的判断逻辑转为numpy数组运算,提前生成所有体素的权重映射,避免在循环中反复判断mask_data的值。

3. 取消vertices列表,直接通过索引访问顶点

原代码中vertices = list(g.vertices())会生成大量顶点对象,占用内存且访问耗时。直接使用整数索引(graph_tool支持用整数表示顶点)进行边的定义,无需提前生成顶点对象列表。

4. 合并重复逻辑,减少冗余计算

将不同类型的边(正向、反向、对角线)统一处理,避免分散的条件判断和重复的索引计算。


优化后的代码实现

import graph_tool.all as gt
import numpy as np

def create_directed_energy_graph_from_mask(mask_data, direction='left', large_weight=10**8):
    z, y, x = mask_data.shape
    num_vertices = z * y * x
    g = gt.Graph(directed=True)
    g.add_vertex(num_vertices)
    
    # 预计算索引转换系数,避免lambda在循环中的开销
    coeff_z = y * x
    coeff_y = x
    
    # 预处理mask:生成每个体素的"保留标记",用于快速计算权重
    mask_keep = (mask_data != 0).astype(np.int32)
    
    # 初始化边列表和权重列表
    edges = []
    weights = []
    
    # 处理方向邻居的正向/反向边(以left方向为例)
    if direction == 'left':
        # 遍历所有体素,除了最后一列(k=x-1)
        for i in range(z):
            for j in range(y):
                for k in range(x - 1):
                    current_idx = i * coeff_z + j * coeff_y + k
                    neighbor_idx = current_idx + 1  # (0,0,1)方向,索引+1
                    
                    # 计算正向边权重
                    weight = 10 if (mask_keep[i,j,k] or mask_keep[i,j,k+1]) else 1
                    edges.append((current_idx, neighbor_idx))
                    weights.append(weight)
                    
                    # 添加反向边
                    edges.append((neighbor_idx, current_idx))
                    weights.append(large_weight)
    
    # 处理YX平面的对角线反向边(k>0时)
    for i in range(z):
        for j in range(y):
            for k in range(1, x):
                current_idx = i * coeff_z + j * coeff_y + k
                # j-1, k-1
                if j > 0:
                    neighbor_idx = i * coeff_z + (j-1)*coeff_y + (k-1)
                    edges.append((current_idx, neighbor_idx))
                    weights.append(large_weight + 1)
                # j+1, k-1
                if j < y - 1:
                    neighbor_idx = i * coeff_z + (j+1)*coeff_y + (k-1)
                    edges.append((current_idx, neighbor_idx))
                    weights.append(large_weight + 1)
    
    # 处理IZ平面的对角线反向边(k>0时)
    for i in range(z):
        for j in range(y):
            for k in range(1, x):
                current_idx = i * coeff_z + j * coeff_y + k
                # i-1, k-1
                if i > 0:
                    neighbor_idx = (i-1)*coeff_z + j*coeff_y + (k-1)
                    edges.append((current_idx, neighbor_idx))
                    weights.append(large_weight + 2)
                # i+1, k-1
                if i < z - 1:
                    neighbor_idx = (i+1)*coeff_z + j*coeff_y + (k-1)
                    edges.append((current_idx, neighbor_idx))
                    weights.append(large_weight + 2)
    
    # 批量添加边和设置权重
    g.add_edge_list(edges)
    weight_prop = g.new_edge_property("int")
    weight_prop.set_2d_array(np.array(weights, dtype=np.int32))
    g.edge_properties["weight"] = weight_prop
    
    return g, weight_prop

额外优化建议

  • 减少循环层数:可以尝试用numpy的广播机制生成所有边的索引,完全替代Python循环(比如用meshgrid生成i,j,k的网格,然后计算所有符合条件的邻居索引),进一步降低Python循环的开销。
  • 数据类型优化:确保large_weight使用整数类型(比如10**8而非1e8),避免浮点数转整数的开销。
  • 禁用不必要的输出:原代码中的print(x,y,z)可以注释掉,减少IO开销。

内容的提问来源于stack exchange,提问作者James

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.23 05:28:12