ai_compiler_tutorial_v2

第2章:计算图表示与抽象

计算图是AI编译器的核心抽象,它将复杂的神经网络计算表示为节点和边的有向图结构。本章将深入探讨计算图的表示方法、静态图与动态图的权衡、中间表示(IR)的设计原则,并通过具身智能的实际案例展示这些概念的应用。通过学习本章,你将理解AI编译器如何将高层的模型描述转换为可优化的计算表示,为后续的优化和代码生成奠定基础。

2.1 计算图的基本概念

2.1.1 什么是计算图

计算图是一种数据流图(Dataflow Graph),用于表示计算过程中数据和操作之间的依赖关系。在AI编译器中,计算图将神经网络的前向传播和反向传播过程表示为一个有向无环图(DAG)或更一般的有向图结构。这种抽象方式起源于数据流架构的思想,但在深度学习时代获得了新的生命力。

与传统程序的控制流图不同,计算图强调数据的流动而非控制的转移。在计算图中,数据像河流一样从上游节点流向下游节点,每个节点都是一个计算单元,对流经的数据进行变换。这种表示方式天然地暴露了并行性——没有数据依赖关系的节点可以同时执行。

计算图的核心组成要素:

     输入X           权重W
   [8,768]         [768,2048]
        |              |
        v              v
    [-------MatMul-------]
              |
         [8,2048]
              v
            偏置B
          [2048]
              |
              v
         [---Add---]
              |
         [8,2048]
              v
         [--ReLU--]
              |
         [8,2048]
              v
           输出Y

计算图的表示不仅仅是静态的结构,它还隐含了执行的语义:

2.1.2 计算图的层次结构

现代AI编译器通常采用多层次的图表示,这种分层架构借鉴了传统编译器的设计思想,但针对张量计算做了专门优化。层次化设计的核心理念是”逐步降低(Progressive Lowering)”——从高层的、抽象的、与硬件无关的表示,逐步转换为低层的、具体的、针对特定硬件优化的表示。

这种设计哲学解决了AI编译器面临的根本矛盾:一方面,我们希望为用户提供高层的、易用的编程接口;另一方面,我们需要生成高效的、充分利用硬件特性的机器码。通过多层次的IR,我们可以在每个层次上进行最适合的优化,同时保持层次间的清晰分离。

  1. 高层图(High-level Graph):
    • 接近用户API的表示,保留了用户的编程抽象。这一层的设计目标是让用户能够用熟悉的概念来表达计算,如”注意力机制”、”批归一化”等。
    • 包含复杂的复合操作,如BatchNorm、LSTM Cell、Attention Block。这些操作在概念上是原子的,但在实现上可能包含数十个基础运算。保留这些高层操作使得编译器可以进行语义级别的优化。
    • 携带高层语义信息,便于进行领域特定优化。例如,知道一个操作是BatchNorm而不仅仅是一系列数学运算,编译器可以应用专门的融合模式或数值稳定性技巧。
    • 支持动态控制流和符号化形状。在这一层,循环的边界可以是符号表达式,条件分支可以依赖运行时值。
  2. 中层图(Mid-level Graph):
    • 将复合操作分解为基础算子的组合。例如,BatchNorm被分解为均值计算、方差计算、归一化和仿射变换。这种分解暴露了更多的优化机会。
    • 标准化的算子集合,便于进行通用优化。中层图通常定义一个精心设计的算子集,既不太高层(失去优化机会)也不太低层(过于复杂)。
    • 开始引入硬件相关的考虑,如数据布局偏好。虽然还不绑定具体硬件,但会考虑目标硬件类别的特性。
    • 适合进行算子融合、内存优化等变换。这是大多数图级优化发生的地方,如将多个逐元素操作融合为一个kernel。
  3. 低层图(Low-level Graph):
    • 接近硬件的表示,包含具体的实现细节。在这一层,抽象的”矩阵乘法”变成了具体的循环嵌套和向量指令。
    • 内存分配、数据布局、并行策略已经确定。每个张量都有确定的内存地址,每个循环都有确定的并行化方案。
    • 循环嵌套、向量化、分块等底层优化。这些优化直接影响缓存利用率和指令级并行度。
    • 可以直接映射到目标硬件的指令。低层图基本上是硬件指令的一对一映射,只是保留了一些结构信息便于最后的调度。

这种层次化设计允许在不同抽象层次进行针对性优化:

高层:BatchNorm(x, gamma, beta, eps=1e-5)
        |
        | [语义保持的降级]
        v
中层:mean = ReduceMean(x, axis=[0,2,3])
      var = ReduceVar(x, axis=[0,2,3])  
      x_norm = (x - mean) / sqrt(var + eps)
      y = gamma * x_norm + beta
        |
        | [循环和内存优化]
        v
低层:for n in range(N):
        for c in range(C):
          sum = 0; sum_sq = 0
          for h in range(H):
            for w in range(W):
              val = load(x[n,c,h,w])
              sum += val
              sum_sq += val * val
          mean[c] = sum / (H*W)
          var[c] = sum_sq / (H*W) - mean[c]^2
          # ... normalization and scaling

每个层次的优化重点不同:

2.1.3 数据依赖与控制依赖

计算图中存在两种主要的依赖关系,它们共同决定了程序的执行语义和优化空间。理解和正确处理这些依赖关系是编译器正确性的基础,也是性能优化的关键。依赖分析不仅影响并行化决策,还影响内存管理、指令调度等多个方面。

数据依赖(Data Dependency):

数据依赖描述了操作之间通过数据建立的时序关系。在计算图中,如果操作B使用了操作A产生的数据,我们说B数据依赖于A。这种依赖关系形成了一个偏序——某些操作必须在其他操作之前执行,但不相关的操作可以以任意顺序执行。

数据依赖的细分类型:

控制依赖(Control Dependency):

控制依赖描述了操作的执行与否依赖于某个控制决策的情况。与数据依赖不同,控制依赖不是关于数据的流动,而是关于执行路径的选择。在包含条件分支和循环的程序中,控制依赖决定了哪些操作会被执行。

控制依赖的处理策略:

if condition:
    y = expensive_op(x)
else:
    y = cheap_op(x)
z = common_op(y)

优化考虑:
1. 投机执行:预先计算两个分支
   - 优点:隐藏分支延迟
   - 缺点:浪费计算资源
   - 适用:分支预测准确率高的场景

2. 延迟执行:推迟到条件确定
   - 优点:避免不必要的计算
   - 缺点:可能增加关键路径长度
   - 适用:分支开销差异大的场景

3. 部分执行:提取公共子计算
   - 优点:减少重复计算
   - 缺点:增加代码复杂度
   - 适用:分支有大量共同计算的场景

隐式依赖:

除了显式的数据流和控制流依赖,还存在一些不那么明显但同样重要的依赖关系。这些隐式依赖常常是bug的来源,也是性能瓶颈的所在。

2.1.4 计算图的属性推导

计算图构建过程中的一个关键任务是属性推导(Attribute Inference),这是一个自底向上和自顶向下相结合的过程。属性推导不仅是为了验证计算的合法性,更是为了收集优化所需的信息。一个强大的属性推导系统可以在编译时发现错误,避免运行时的意外,同时为后续的优化passes提供必要的元数据。

属性推导的复杂性在于它需要处理部分信息的情况。在构建计算图时,某些属性可能未知(如动态形状),某些属性相互依赖(如设备分配影响内存布局),某些属性有多个合法选择(如数据类型的自动提升)。一个好的推导系统需要优雅地处理这些情况。

  1. 形状推导(Shape Inference):

    形状推导是最基础也是最重要的属性推导。知道张量的形状对于内存分配、并行化策略、算子选择都至关重要。现代AI编译器的形状推导系统需要处理越来越复杂的场景。

    • 根据输入形状和操作语义推导输出形状。每个操作都有其形状变换规则,编译器需要实现这些规则。例如,卷积的输出形状取决于输入形状、卷积核大小、步长和填充。

    • 处理符号维度和动态形状。在许多应用中,某些维度(如batch size)在编译时未知。编译器需要用符号变量表示这些维度,并维护它们之间的约束关系。

    • 验证形状兼容性,早期发现错误。形状不匹配是深度学习中最常见的错误之一。通过静态形状推导,我们可以在运行前就发现这些错误。

    • 支持广播语义和形状约束。广播机制让我们可以在不同形状的张量间进行运算,但需要遵循特定规则。编译器需要理解和验证这些规则。

  2. 类型推导(Type Inference):

    类型推导确保数据类型的一致性和效率。在混合精度训练越来越普及的今天,智能的类型推导可以在保证精度的同时最大化性能。

    • 确定数据类型的传播和转换规则。不同操作对数据类型有不同要求,编译器需要确定何时保持类型、何时转换类型。

    • 处理混合精度计算的类型提升。当不同精度的数据相遇时,需要决定提升到哪种精度。这需要在精度损失和性能之间权衡。

    • 插入必要的类型转换节点。类型转换不是免费的,需要额外的计算和内存。编译器需要最小化转换次数。

    • 优化类型转换的位置以减少开销。同一个类型转换可能在多个位置进行,编译器需要找到最优位置。

  3. 设备推导(Device Inference):

    在异构计算环境中,决定每个操作在哪个设备上执行是一个关键决策。这不仅影响性能,还影响内存管理和数据传输。

    • 确定操作在哪个计算设备上执行。某些操作可能只能在特定设备上执行,而另一些操作在不同设备上有不同的性能特性。

    • 考虑设备亲和性和数据局部性。将相关操作放在同一设备可以减少数据传输。数据的当前位置也影响设备选择。

    • 插入必要的数据传输节点。当连续的操作在不同设备上执行时,需要显式的数据传输。这些传输可能成为瓶颈。

    • 优化跨设备通信模式。批量传输、异步传输、传输与计算重叠等技术可以隐藏通信开销。

  4. 内存推导(Memory Inference):

    内存是AI计算的关键瓶颈之一。准确的内存推导可以帮助我们优化内存使用,避免OOM(Out of Memory)错误。

    • 估算每个操作的内存需求。这包括输入、输出和临时缓冲区的大小。对于某些操作(如卷积),临时缓冲区可能很大。

    • 确定张量的生命周期。知道每个张量何时创建、何时最后使用,可以帮助我们及时释放内存。

    • 识别内存复用机会。生命周期不重叠的张量可以共享内存。这在内存受限的设备上特别重要。

    • 生成内存分配计划。决定何时分配、何时释放、如何复用,生成一个高效的内存管理方案。

形状推导的数学表示: \(\text{Shape}(Y) = f_{\text{op}}(\text{Shape}(X_1), \text{Shape}(X_2), ..., \text{attrs})\)

例如,对于矩阵乘法: \(\text{Shape}(C) = \text{MatMul}([M, K], [K, N]) = [M, N]\)

对于带有广播的逐元素操作: \(\text{Shape}(C) = \text{Broadcast}(\text{Shape}(A), \text{Shape}(B))\)

其中广播规则遵循NumPy语义:

  1. 从右向左对齐维度
  2. 维度大小必须相同或其中一个为1
  3. 缺失的维度视为1

符号形状的处理:

输入: x.shape = [?, 224, 224, 3]  # ?表示batch维度
卷积: conv.weight = [64, 3, 3, 3]
输出: y.shape = [?, 224, 224, 64]

约束: ? > 0 且 ? % 8 == 0  # 硬件对齐要求

2.2 静态图vs动态图

2.2.1 静态图的特点

静态图(Static Graph)在执行前完整构建整个计算图,这种”定义即编译”的模式带来了独特的优化机会:

优势:

劣势:

静态图的编译流程:

定义阶段:
  用户代码 -> AST构建 -> 图构建 -> 图验证
                                    |
                                    v
                              形状/类型推导
                                    |
                                    v
优化阶段:                    图优化passes
  代数简化 -> 算子融合 -> 内存优化 -> 并行化
                                    |
                                    v
代码生成:                    目标代码生成
  设备分配 -> kernel选择 -> 代码发射 -> 二进制

执行阶段:
  加载模型 -> 绑定输入 -> 执行计划 -> 返回输出

2.2.2 动态图的特点

动态图(Dynamic Graph)采用即时构建和执行的模式,每个操作立即执行:

优势:

劣势:

动态图的执行模式:

for batch in data_loader:
    # 每次迭代都会:
    1. 构建操作节点
    2. 检查输入合法性
    3. 分配输出内存
    4. 调用kernel执行
    5. 更新梯度tape
    6. 清理临时对象
    
    output = model(batch)  # 立即执行
    loss = criterion(output, target)
    loss.backward()  # 动态构建反向图

2.2.3 混合执行模式

现代AI框架趋向于结合静态图和动态图的优势,主要策略包括:

  1. 即时编译(JIT Compilation):
    @torch.jit.script  # PyTorch的JIT装饰器
    def optimized_layer(x, w, b):
        # 第一次执行时追踪
        # 后续执行使用编译版本
        y = torch.matmul(x, w) + b
        return torch.relu(y)
    
  2. 追踪(Tracing):
    • 记录一次执行的操作序列
    • 生成对应的静态图
    • 假设控制流不变
      # 追踪模式
      traced_model = torch.jit.trace(model, example_input)
      # 之后可以像静态图一样优化和部署
      
  3. 符号化执行(Symbolic Execution):
    • 延迟执行直到需要具体值
    • 构建符号表达式树
    • 支持符号化的控制流
      # TensorFlow的tf.function
      @tf.function
      def dynamic_rnn(x, length):
        for i in tf.range(length):  # 符号化的循环
            x = cell(x)
        return x
      
  4. 分阶段执行(Staged Execution):
    # JAX的分阶段编程
    @jax.jit
    def train_step(params, batch):
        # 静态编译的训练步骤
        def loss_fn(params):
            # 动态的损失计算
            return compute_loss(params, batch)
           
        grads = jax.grad(loss_fn)(params)
        return update_params(params, grads)
    

混合模式的实现策略:

执行流程:
1. 动态执行并收集trace
2. 识别热点代码路径(执行次数 > 阈值)
3. 提取可静态化的子图
4. 对子图进行编译优化
5. 缓存编译结果
6. 后续执行使用编译版本

关键技术:
- Trace缓存:基于输入签名的缓存
- Guards:检查假设是否仍然成立
- Bailout:假设失败时回退到解释执行
- Recompilation:适应新的执行模式

2.2.4 自动驾驶场景的权衡

在自动驾驶系统中,不同模块对图表示有不同需求,体现了实际系统的复杂性:

感知模块(如目标检测):

规划模块(如行为规划):

预测模块(如轨迹预测):

控制模块(如MPC控制器):

2.3 中间表示(IR)设计

2.3.1 IR的设计目标

中间表示是编译器前端和后端的桥梁,其设计需要平衡多个目标:

  1. 表达能力:能够表示所有必要的操作和控制流
  2. 优化友好:便于进行各种变换和优化
  3. 可扩展性:易于添加新操作和新优化
  4. 可移植性:与具体硬件无关,但能高效映射到硬件

2.3.2 多级IR架构

现代AI编译器通常采用多级IR设计:

前端语言 (Python/C++)
        |
        v
    Graph IR (高层抽象)
        |
        v
   Tensor IR (张量程序)
        |
        v
    Loop IR (循环嵌套)
        |
        v
   Machine IR (机器指令)

每一级IR都有其优化重点:

2.3.3 SSA形式与数据流分析

许多IR采用静态单赋值(SSA)形式,便于数据流分析:

SSA的特点:

SSA形式的优势在张量别名分析中尤为明显:

# 非SSA形式
x = conv(input)
x = relu(x)      # x被重新赋值
y = pool(x)

# SSA形式
x1 = conv(input)
x2 = relu(x1)    # 新变量名
y = pool(x2)

2.3.4 类型系统与形状推导

IR的类型系统需要捕获张量的完整信息:

TensorType := {
    dtype: DataType,      // 数据类型:f32, i8, bf16等
    shape: Shape,         // 形状:可以包含符号维度
    layout: Layout,       // 内存布局:NCHW, NHWC等
    device: Device        // 设备:CPU, GPU, NPU等
}

形状可以是:

形状推导规则示例(广播): \(\text{broadcast}([M, 1, N], [1, K, N]) = [M, K, N]\)

2.3.5 IR的可扩展性设计

可扩展性是IR长期演进的关键:

操作注册机制:

OpRegistry {
    name: string
    inputs: List[TensorType]
    outputs: List[TensorType]
    attributes: Dict[str, Any]
    shape_fn: Function
    lower_fn: Function
}

自定义操作支持:

2.4 具身智能案例:机械臂控制的计算图建模

2.4.1 场景描述

考虑一个七自由度机械臂执行抓取任务的场景:

  1. 输入:RGB-D图像、关节位置、目标位置
  2. 输出:关节扭矩指令
  3. 约束:碰撞避免、关节限位、实时性要求(10ms)

这个系统包含多个计算模块:

2.4.2 混合计算图设计

机械臂控制系统展现了静态图和动态图混合的必要性:

静态部分(感知网络):

RGB图像 -> CNN特征提取 -> 目标检测 -> 位姿估计
深度图像 -> 点云生成 -> 配准 -> 障碍物地图

动态部分(规划与控制):

while not reached_target:
    if obstacle_detected:
        path = replan_trajectory()
    else:
        path = interpolate_path()
    
    torques = compute_control(path, current_state)
    apply_torques(torques)

2.4.3 实时性约束的表达

在IR中表达实时性约束:

@deadline(10ms)
@priority(high)
def control_loop():
    perception = perception_net(image)  # 预编译的静态图,3ms
    plan = motion_planner(perception)   # 动态图,变长计算
    control = controller(plan)          # 静态图,1ms
    return control

编译器需要:

  1. 时间预算分配:为每个子图分配时间预算
  2. 优先级调度:确保关键路径优先执行
  3. 降级策略:超时时使用简化算法

2.4.4 变长输入的处理

机械臂需要处理变长输入(如不同数量的障碍物):

填充(Padding)策略:

obstacles = pad_to_max(detected_obstacles, MAX_OBSTACLES)
mask = create_mask(num_actual_obstacles)
distances = compute_distances(obstacles, mask)

动态批处理策略:

for obstacle in detected_obstacles:
    distance = compute_distance(obstacle)
    if distance < threshold:
        avoidance_forces.append(compute_force(obstacle))

稀疏表示策略:

sparse_obstacles = to_sparse(obstacle_map)
collisions = sparse_collision_check(trajectory, sparse_obstacles)

2.4.5 计算图的分区与调度

针对机械臂的异构计算资源:

设备分配:
- GPU: 视觉感知网络(高并行度)
- CPU: 运动规划(复杂逻辑)
- FPGA: 控制器(低延迟)

数据流:
GPU -> CPU: 目标位置、障碍物信息(每帧30Hz)
CPU -> FPGA: 轨迹点(每10ms)
FPGA -> Actuator: 控制信号(每1ms)

编译器优化策略:

  1. 流水线并行:感知、规划、控制三级流水线
  2. 数据预取:提前准备下一帧数据
  3. 增量计算:只更新变化的部分

2.5 本章小结

本章深入探讨了AI编译器中计算图的表示与抽象:

核心概念:

关键权衡:

设计原则:

实践洞察:

2.6 常见陷阱与错误

陷阱1:过早固定图结构

问题:在开发早期就采用纯静态图,限制了后续的灵活性 解决:先用动态图快速迭代,性能瓶颈处再静态化

陷阱2:忽视内存别名

问题:错误的别名分析导致in-place操作破坏数据 解决:采用SSA形式,显式跟踪张量生命周期

陷阱3:形状推导的符号维度处理

问题:符号维度的约束传播不完整 解决:建立完整的约束系统,使用SMT求解器验证

陷阱4:控制流与数据流混淆

问题:将控制依赖当作数据依赖,导致错误的并行化 解决:明确区分两种依赖,使用不同的边类型

陷阱5:IR层次选择不当

问题:在错误的IR层次进行优化,效果不佳 解决:理解每层IR的优化时机,遵循分层优化原则

陷阱6:动态图的性能退化

问题:动态图中的热点路径没有被优化 解决:实现追踪和JIT机制,自动识别和优化热点

2.7 练习题

🟢 练习2.1:计算图构建

给定如下神经网络层:y = ReLU(BatchNorm(Conv2d(x, w) + b)),画出对应的计算图,标注每个节点的操作类型和边的数据类型。

💡 提示:注意BatchNorm包含多个子操作(均值、方差、归一化、缩放、偏移)

📝 参考答案 计算图结构: ``` x w | | v v [----Conv2d----] | v b | v [--Add--] | v [--BatchNorm--] / | \ Mean Var Scale/Shift \ | / [--Normalize--] | v [--ReLU--] | v y ``` 节点类型: - Conv2d: 卷积操作 - Add: 逐元素加法 - BatchNorm: 复合操作(可进一步分解) - ReLU: 激活函数 边的数据类型: - x: [N, C_in, H, W] - w: [C_out, C_in, K, K] - b: [C_out] - Conv2d输出: [N, C_out, H', W'] - 最终输出y: [N, C_out, H', W']

🟢 练习2.2:静态图vs动态图选择

对于以下场景,分析应该选择静态图、动态图还是混合模式,并说明理由:

  1. 实时视频流的目标跟踪
  2. 自然语言的序列到序列翻译
  3. 固定大小图像的分类
  4. 强化学习的策略网络

💡 提示:考虑输入形状、控制流复杂度、性能要求

📝 参考答案 1. **实时视频流的目标跟踪**:静态图或混合模式 - 输入形状固定(视频分辨率) - 性能要求高(实时性) - 可能需要动态的跟踪逻辑 2. **序列到序列翻译**:动态图或混合模式 - 输入长度可变 - 需要条件解码(beam search) - 注意力机制的动态计算 3. **固定大小图像分类**:静态图 - 输入形状完全固定 - 无控制流 - 可以最大程度优化 4. **强化学习策略网络**:动态图 - 需要与环境交互 - 探索策略可能变化 - 需要在线更新

🟡 练习2.3:形状推导

实现一个简化的形状推导系统,处理以下操作的形状传播:

考虑符号维度(如batch_size = ?)的情况。

💡 提示:使用符号表达式表示未知维度,建立约束方程

📝 参考答案 形状推导系统的核心设计: 1. **符号维度表示**: - 使用符号变量:`Sym(name)` - 常量维度:`Const(value)` - 表达式:`Mul(Sym('M'), Const(2))` 2. **约束收集**: ``` MatMul约束: - input1.shape[-1] == input2.shape[0] - output.shape = [input1.shape[0], input2.shape[1]] Reshape约束: - prod(input.shape) == prod(output.shape) ``` 3. **广播规则**: - 从右向左对齐维度 - 1可以广播到任意大小 - 相同大小可以匹配 - 不兼容则报错 4. **约束求解**: - 收集所有等式约束 - 使用统一化(unification)算法 - 检查约束一致性

🟡 练习2.4:SSA转换

将以下计算序列转换为SSA形式,并分析哪些操作可以原地执行:

x = input()
x = conv(x, w1)
y = relu(x)
x = pool(y)
z = x + y
output(z)

💡 提示:跟踪每个变量的生命周期,判断何时可以复用内存

📝 参考答案 SSA形式: ```python x0 = input() x1 = conv(x0, w1) y0 = relu(x1) x2 = pool(y0) z0 = x2 + y0 output(z0) ``` 生命周期分析: - x0: input后即可释放(被x1消费) - x1: relu后可释放(被y0消费) - y0: 需要保留到加法操作 - x2: 需要保留到加法操作 - z0: output后释放 原地操作机会: - conv可以覆盖x0的内存(如果尺寸兼容) - relu可以原地执行(x1 → y0) - pool不能原地(y0还需要用于加法) - 加法结果z0需要新内存

🔴 练习2.5:JIT编译策略

设计一个JIT编译策略,用于识别和优化动态图中的热点路径。考虑:

  1. 如何识别热点路径?
  2. 何时触发编译?
  3. 如何处理输入形状变化?
  4. 编译缓存如何管理?

💡 提示:考虑追踪(tracing)、计数器、形状特化等技术

📝 参考答案 JIT编译策略设计: 1. **热点识别**: - 执行计数器:每个子图维护执行次数 - 时间profiling:记录执行时间占比 - 模式匹配:识别重复的计算模式 - 阈值:执行次数>100或时间占比>10% 2. **编译触发时机**: - 达到热点阈值时异步编译 - 空闲时预编译可能的路径 - 第一次执行时的快速路径检测 3. **形状变化处理**: - 形状特化:为常见形状生成特化版本 - 形状类别:将形状分组(小/中/大) - 动态形状:保留符号维度,运行时绑定 - 重编译:形状变化超过阈值时重新编译 4. **缓存管理**: - LRU缓存:限制编译代码总大小 - 形状签名:(形状,设备,精度)作为key - 版本控制:模型更新时invalidate - 预热:从持久化存储加载常用编译结果 5. **自适应策略**: - 监控编译收益vs开销 - 动态调整编译阈值 - 根据内存压力调整缓存大小

🔴 练习2.6:混合精度图变换

设计一个图变换算法,将FP32计算图转换为混合精度(FP16/FP32)计算图。需要考虑:

  1. 哪些操作适合FP16?
  2. 如何插入必要的类型转换?
  3. 如何保证数值稳定性?

💡 提示:某些操作(如损失计算)需要保持FP32精度

📝 参考答案 混合精度转换算法: 1. **操作分类**: - FP16安全:Conv, MatMul, ReLU(计算密集) - FP32必需:Loss, BatchNorm统计, Adam更新 - 条件转换:根据数值范围决定 2. **转换规则**: ``` 白名单操作 → FP16 黑名单操作 → FP32 灰名单操作 → 动态决策 ``` 3. **类型转换插入**: - 前向转换:FP32→FP16在白名单操作前 - 后向转换:FP16→FP32在黑名单操作前 - 最小化转换:合并相邻的转换 4. **数值稳定性保证**: - Loss scaling:梯度放大2^16避免下溢 - 动态范围检查:监控激活值范围 - 主权重保持FP32:更新在FP32进行 5. **优化策略**: - 转换融合:与其他操作融合减少开销 - 张量复用:避免额外内存分配 - 批量转换:多个小张量合并转换

🟢 练习2.7:计算图可视化

设计一个计算图的文本可视化方案,要求:

  1. 显示节点类型和连接关系
  2. 标注张量形状
  3. 区分前向和反向边

💡 提示:使用ASCII art或简化的DOT语言

📝 参考答案 文本可视化方案: ``` === 计算图结构 === [Input:x] (shape=[8,3,224,224]) | v [Conv2d:conv1] (weight=[64,3,7,7]) | (shape=[8,64,112,112]) v [ReLU:relu1] | (shape=[8,64,112,112]) v [MaxPool2d:pool1] (kernel=3x3) | (shape=[8,64,56,56]) v [Output:y] === 反向传播路径 === [Loss:grad_output] <--- [Output:y] ^ | [MaxPool2d_backward] <--- [MaxPool2d:pool1] ^ | [ReLU_backward] <--- [ReLU:relu1] ^ | [Conv2d_backward] <--- [Conv2d:conv1] ^ | [Input:grad_x] ``` 特点: - 节点格式:[类型:名称] - 形状标注:(shape=[...]) - 前向边:实线 | - 反向边:虚线 ^ - 属性显示:关键参数

🟡 练习2.8:设备分配优化

给定一个包含10个操作的计算图和2个GPU,设计一个设备分配策略,最小化:

  1. 设备间数据传输
  2. 计算负载不均衡
  3. 总执行时间

💡 提示:这是一个图分割(graph partitioning)问题,考虑使用启发式算法

📝 参考答案 设备分配策略: 1. **成本模型**: - 计算成本:每个操作在不同设备的执行时间 - 传输成本:跨设备边的数据传输时间 - 内存成本:每个设备的内存使用 2. **分割算法**: ``` 初始化:随机分配或基于拓扑排序 迭代优化: for node in graph: 计算node在每个设备的总成本 成本 = 计算成本 + 传输成本 将node分配到成本最小的设备 直到收敛或达到迭代上限 ``` 3. **负载均衡**: - 统计每个设备的总计算时间 - 如果不均衡度>阈值,迁移边界节点 - 考虑关键路径,优先均衡瓶颈 4. **优化技巧**: - 操作聚类:将强连接的操作分组 - 数据复制:热点数据在多设备复制 - 流水线:重叠计算和通信 5. **实际考虑**: - 某些操作只能在特定设备运行 - 内存限制可能强制分割 - 通信拓扑影响传输成本