P小二 P小二
← 返回文章 AI Research 约 17 分钟

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

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

往期文章:

参考链接: