跳到主要内容
P小二 P小二
← 返回文章 • AI Research • • 约 16 分钟

deepseek开源周第三天:DeepGEMM深度分析

今天带来deepseek开源DeepGEMM深度分析

今天,deepseek开源周第三天,9点准时发布,这次是DeepGEMM。

今天,deepseek开源周第三天,9点准时发布,这次是DeepGEMM。

发布后,目前项目3.3k star。

发布后,目前项目3.3k star。

官方介绍 DeepGEMM是一款支持 FP8 的 GEMM 库,兼容稠密和 MoE GEMM,用于 V3/R1 的训练和推理。

  • ⚡ 在 Hopper GPU 上可达 1350+ FP8 TFLOPS

  • ✅ 无繁重依赖,简单如教程

  • ✅ 完全即时编译 (Just-In-Time compiled)

  • ✅ 核心逻辑仅 ~300 行代码,却在大多数矩阵规模上超越专家优化的内核性能

  • ✅ 支持稠密布局和两种 MoE 布局

GEMM简单介绍

通用矩阵乘法(GEMM)是深度学习和科学计算中最基础也是最重要的计算操作之一。GEMM 代表 General Matrix Multiplication,即两个矩阵 A 和 B 相乘得到结果矩阵 C 的操作,通常表示为 C = A × B。

通用矩阵乘法(GEMM)是深度学习和科学计算中最基础也是最重要的计算操作之一。G

在深度学习中,GEMM 是全连接层、卷积层和注意力机制等核心组件的基础。例如,在 Transformer 架构中,自注意力和前馈网络层都大量依赖矩阵乘法。随着模型规模的增长,GEMM 操作占据了模型训练和推理过程中的大部分计算时间,因此 GEMM 的性能直接影响整个深度学习系统的效率。

现代 GPU 架构专门为加速矩阵乘法而设计,如 NVIDIA 的 Tensor Core 技术。随着 模型规模的不断扩大,对 GEMM 性能的要求也越来越高。特别是在LLM和MoE中,高效的 GEMM 实现对于实现实时推理和降低训练成本至关重要。

在deepseek发表的论文DeepSeek LLM Scaling Open-Source Language Models with Longtermism 中有提到GEMM,不过都是和他们另外一篇Fire-Flyer AI-HPC: A Cost-Effective Software-Hardware Co-Design for Deep Learning 中介绍的HAI-LLM训练系统有关,有兴趣可以去读一下Fire-Flyer AI-HPC这篇论文。

在deepseek发表的论文DeepSeek LLM Scaling Open

今天开源的DeepGEMM 是支持FP8的,主要是用于DeepSeek V3/R1的训练和推理的,在V3论文中deepseek为FP8训练做了许多优化。

今天开源的DeepGEMM 是支持FP8的,主要是用于DeepSeek V3/R

使用FP8做训练的主要挑战是在于精度与误差处理,deepseek为了支持FP8训练做了以下的优化措施:

  • 细粒度量化:将数据分成更小的组,每个组使用特定乘数来保持高精度。

  • 在线量化:在线计算每1x128激活块或128x128权重块的权重值,在线推算缩放因子,激活或者在线转换为FP8格式

  • 提高累加精度:FP8大量累加容易出现随机误差,将中间结果存储在FP32中,累加之后在转化回来。

  • 低精度/混合精度存储于通信:训练MoE模型时,混合使用FP8和BF16/FP32,确定模型的动态稳定。

详细的优化措施,有兴趣可以去看deepseek V3的论文。

详细的优化措施,有兴趣可以去看deepseek V3的论文。

DeepGEMM介绍

总结它的主要特点:

  • 支持FP8:DeepGEMM采用了 CUDA 核心两级累加(解决不精确的问题)

  • 支持分组GEMM:主要是改进了CUTLASS的分组GEMM,对MoE模型针对性优化

  • 即时编译: 通过 JIT 技术,代码可以在运行时动态生成和优化,进一步提升性能和灵活性

  • FFMA SASS 交错:deepseek深入分析了SASS编译结果,在FFMA/FADD中调整SASS指令,提高了细粒度 FP8 GEMM效率

性能

FFMA SASS 交错:deepseek深入分析了SASS编译结果,在FFMA

所有指标都有提高,最高一项有2.7倍加速。deepseek认为在某些地方表现并不太好,有兴趣优化还可以给他们提PR。

在DeepGEMM 项目的README文档中,deepseek团队还对DeepGEMM的优化做了详细的介绍,有兴趣可以结合代码仔细阅读和测试。

今天主要介绍项目中的一个文件,在jit目录下有个interleave_ffma.py文件,里面有一些骚操作。

代码如下:

import argparseimport mmapimport osimport reimport subprocessfrom torch.utils.cpp_extension import CUDA_HOMEdef run_cuobjdump(file_path):
    command = [f'{
CUDA_HOME}/bin/cuobjdump', '-sass', file_path]
    result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
    assert result.returncode == 0
    return result.stdoutdef extract_ffma(sass):
    lines = sass.splitlines()
    collected = []
    current = []
    arch_name, func_name = 'N/A', 'N/A'
    skip_next_line = False
    for line in lines:
        if 'code for' in line:
            arch_name = line.lstrip().lstrip('code for ').rstrip()
        elif 'Function :' in line:
            func_name = line.lstrip().lstrip('Function :').rstrip()
        elif 'FFMA' in line:
            current.append(line)
            skip_next_line = True
        elif skip_next_line:
            current.append(line)
            skip_next_line = False
        else:
            if len(current) >= 16:
                assert len(current) % 2 == 0
                collected.append((f'{
arch_name}::{
func_name}', current))
            current = []
    if os.getenv('DG_PRINT_REG_REUSE', None):
        print(f"Found {
len(collected)} FFMA segments")
    return collecteddef extract_hex_from_line(line):
    match = re.search(r'/\*\s*(0x[0-9a-fA-F]+)\s*\*/', line)
    assert match
    return int(match.group(1), 16)def validate(m, offset, le_bytes, num_lines):
    assert len(le_bytes) == num_lines // 2
    assert m[offset:offset + 16] == le_bytes[0]
    for i in range(1, num_lines // 2):
        if m[offset + i * 16:offset + i * 16 + 16] != le_bytes[i]:
            return False
    return Truedef parse_registers(line):
    import re
    line = re.sub(r'/\*.*?\*/', '', line)
    line = line.replace(';', '')
    tokens = line.strip().split(',')
    registers = []
    for token in tokens:
        token = token.strip()
        words = token.split()
        for word in words:
            if word.startswith('R'):
                reg = word.split('.')[0]
                registers.append(reg)
    return registersdef modify_segment(m, name, ffma_lines):
    num_lines = len(ffma_lines)
    assert num_lines % 2 == 0
    le_bytes, new_le_bytes = [], []
    reused_list = []
    dst_reg_set = set()
    last_reused, last_dst_reg = False, ''
    num_changed = 0
    for i in range(num_lines // 2):
        dst_reg = parse_registers(ffma_lines[i * 2])[-2]
        low_line, high_line = ffma_lines[i * 2], ffma_lines[i * 2 + 1]
        low_hex, high_hex = extract_hex_from_line(low_line), extract_hex_from_line(high_line)
        le_bytes.append(low_hex.to_bytes(8, 'little') + high_hex.to_bytes(8, 'little'))
        reused = (high_hex & 0x0800000000000000) != 0
        if reused:
            is_first_occurred = dst_reg not in dst_reg_set
            if is_first_occurred or (last_reused and dst_reg == last_dst_reg):
                # Modify the `reuse` and `yield` bits
                assert high_hex & 0x0800200000000000, f"{
hex(high_hex)}"
                high_hex ^= 0x0800200000000000
                reused = False
                num_changed += 1
            else:
                reused_list.append(i)
        dst_reg_set.add(dst_reg)
        new_le_bytes.append(low_hex.to_bytes(8, 'little') + high_hex.to_bytes(8, 'little'))
        last_reused, last_dst_reg = reused, dst_reg
    if os.getenv('DG_PRINT_REG_REUSE', None):
        print(f" > segment `{
name}` new reused list ({
num_changed} changed): {
reused_list}")
    # Find the offset
    offsets = []
    offset = m.find(le_bytes[0])
    while offset != -1:
        offsets.append(offset)
        offset = m.find(le_bytes[0], offset + 1)
    offsets = list(filter(lambda x: validate(m, x, le_bytes, num_lines), offsets))
    # Replace with `new_le_bytes`
    for offset in offsets:
        for i in range(num_lines // 2):
            m[offset + i * 16:offset + i * 16 + 16] = new_le_bytes[i]def process(path):
    if os.getenv('DG_PRINT_REG_REUSE', None):
        print(f'Processing {
path}')
    output = run_cuobjdump(path)
    segments = extract_ffma(output)
    with open(path, 'r+b') as f:
        mm = mmap.mmap(f.fileno(), 0, access=mmap.ACCESS_WRITE)
        for segment in segments:
            modify_segment(mm, *segment)
        mm.close()if __name__ == "__main__":
    parser = argparse.ArgumentParser(description='Interleave FFMA reg reuse')
    parser.add_argument('--so', help='Path to the SO file')
    args = parser.parse_args()
    process(args.so)

这个文件主要用于优化CUDA编译后的汇编代码中的FFMA(Fused Floating-point Multiply-Add)指令的寄存器重用模式。主要目的是通过修改二进制文件来改善GPU指令的执行效率。

主要函数功能:

SASS代码提取

def run_cuobjdump(file_path):
    # 使用CUDA工具cuobjdump来提取SASS(汇编)代码
    command = [f'{
CUDA_HOME}/bin/cuobjdump', '-sass', file_path]

FFMA指令分析

def extract_ffma(sass):
    # 从SASS代码中提取包含FFMA指令的代码段
    # 收集架构名称、函数名称和相关FFMA指令序列

寄存器使用分析

def parse_registers(line):
    # 解析指令中使用的寄存器
    # 提取以'R'开头的寄存器标识符

二进制修改

def modify_segment(m, name, ffma_lines):
    # 修改FFMA指令的重用位和yield位
    # 通过修改特定位模式(0x0800200000000000)来优化寄存器重用

工作流程

  • 首先读取编译后的CUDA共享库(.so文件)

  • 使用cuobjdump工具提取SASS代码

  • 识别并收集所有FFMA指令序列

  • 分析每个FFMA指令的寄存器使用模式

  • 根据特定规则修改寄存器重用标志

  • 将修改后的指令写回原文件

优化策略

该工具主要针对以下情况进行优化:

  • 当寄存器首次使用时

  • 当连续重用相同目标寄存器时

  • 通过修改重用(reuse)和yield位来优化指令调度

使用方式

python interleave_ffma.py --so path/to/cuda_lib.so

分析

这个工具是一个后处理优化器,它在CUDA代码编译后运行,通过修改生成的二进制文件来优化GPU指令的执行效率。主要关注点是FFMA指令的寄存器重用模式,这对于深度学习等计算密集型应用的性能有重要影响。

文件里面有个正则很多人看不懂

def extract_hex_from_line(line):
    match = re.search(r'/\\*\\s*(0x[0-9a-fA-F]+)\\s*\\*/', line)
    assert match
    return int(match.group(1), 16)

其实是在CUDA SASS汇编代码中,指令通常以这种格式出现:

FFMA R8, R8, R6, R4;
                  /* 0x5c98078000870808 */

这个正则表达式会从上面的代码中提取出 0x5c98078000870808,这个十六进制数代表了实际的机器指令编码。

这个函数:

  • 从汇编代码行中提取十六进制指令编码

  • 使用match.group(1)获取第一个捕获组(即十六进制数)

  • 将其转换为整数用于后续的指令修改

这是整个工具中重要的一步,因为它需要这些指令编码来:

  • 定位需要修改的指令

  • 修改指令中的特定位(如重用标志和yield位)

  • 将修改后的指令写回文件

只能感叹,deepseek工程师比一般的英伟达员工更懂CUDA!你看,他们API又降价了。

只能感叹,deepseek工程师比一般的英伟达员工更懂CUDA!你看,他们API

关注我,明天继续带来深度分析,我们一起猛学。

往期文章:

参考链接: