deepseek开源周第三天:DeepGEMM深度分析
今天带来deepseek开源DeepGEMM深度分析
今天,deepseek开源周第三天,9点准时发布,这次是DeepGEMM。

发布后,目前项目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 是全连接层、卷积层和注意力机制等核心组件的基础。例如,在 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这篇论文。

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

使用FP8做训练的主要挑战是在于精度与误差处理,deepseek为了支持FP8训练做了以下的优化措施:
-
细粒度量化:将数据分成更小的组,每个组使用特定乘数来保持高精度。
-
在线量化:在线计算每1x128激活块或128x128权重块的权重值,在线推算缩放因子,激活或者在线转换为FP8格式
-
提高累加精度:FP8大量累加容易出现随机误差,将中间结果存储在FP32中,累加之后在转化回来。
-
低精度/混合精度存储于通信:训练MoE模型时,混合使用FP8和BF16/FP32,确定模型的动态稳定。
详细的优化措施,有兴趣可以去看deepseek V3的论文。

DeepGEMM介绍
总结它的主要特点:
-
支持FP8:DeepGEMM采用了 CUDA 核心两级累加(解决不精确的问题)
-
支持分组GEMM:主要是改进了CUTLASS的分组GEMM,对MoE模型针对性优化
-
即时编译: 通过 JIT 技术,代码可以在运行时动态生成和优化,进一步提升性能和灵活性
-
FFMA SASS 交错:deepseek深入分析了SASS编译结果,在FFMA/FADD中调整SASS指令,提高了细粒度 FP8 GEMM效率
性能

所有指标都有提高,最高一项有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又降价了。

关注我,明天继续带来深度分析,我们一起猛学。
往期文章:
参考链接: