Skip to content

Repository files navigation

DeepGEMM

DeepGEMM is a library designed for clean and efficient General Matrix Multiplications (GEMMs) on GCU (General Compute Unit) accelerators. It provides FP8 GEMM kernels for both normal and Mixture-of-Experts (MoE) grouped scenarios, as well as FP8/FP4 attention logits kernels for multi-query attention (MQA).

This project is derived from deepseek-ai/DeepGEMM and adapted for the GCU backend.

Project Structure

DeepGEMM/
├── deep_gemm/             # Core Python package
│   ├── __init__.py        # Public API exports
│   ├── attention.py       # Triton-based attention reference implementation
│   ├── testing/           # Benchmarking and numeric testing utilities
│   └── utils/             # Math and layout utilities
├── csrc/                  # C++ extension source (pybind11)
│   ├── apis/              # GCU kernel API bindings (GEMM, attention, layout)
│   └── utils/             # C++ utility headers
├── cmake/                 # CMake build configuration
├── 2nd/                   # Third-party dependency management
├── tests/                 # GCU unit tests
│   └── deepgemm_official_tests/  # Extended GCU test suite
├── setup.py               # Python package build configuration
├── develop.sh             # Development build script
└── install.sh             # Installation script (builds wheel)

Requirements

  • Python >= 3.10
  • C++17 compatible compiler
  • PyTorch >= 2.8
  • torch_gcu (GCU PyTorch extension)
  • topsaten (GCU ATen-compatible SDK)
  • topsruntime (GCU runtime)
  • triton_kernel_gcu (optional, for Triton-based GCU kernels)

Installation

Docker (Recommended)

The easiest way to get started is using the pre-built Docker image with all dependencies included.

  1. Pull and start the container:

    IMAGE=artifact.enflame.cn/enflame_docker_release/public_deepgemm:torch_2.11_gcc11_cxx3.4.30_v2
    
    docker run --name deep_gemm -d \
      -v /home:/home \
      --shm-size 8G \
      --ipc=host --network host \
      --cap-add SYS_PTRACE \
      --security-opt seccomp=unconfined \
      --privileged \
      "$IMAGE" \
      tail -f /dev/null
  2. Update the host GCU driver (to match the image's software version):

    # Extract the matching driver from the container
    docker cp deep_gemm:/enflame/driver ./
    
    # Install the driver on the host
    sudo driver/enflame-x86_64-gcc-*.run
    
    # Restart the container to pick up the new driver
    docker restart deep_gemm
  3. Clone the source code:

    cd /home
    git clone git@github.com:EnflameTechnology/DeepGEMM.git
  4. Enter the container and build:

    docker exec -it deep_gemm bash
    cd /home/DeepGEMM
    ./install.sh

Build from Source

git clone <repo_url>
cd DeepGEMM
./install.sh

This builds a wheel package and installs it via pip.

Development Build

For development with in-place shared library:

cd DeepGEMM
./develop.sh

Then import deep_gemm in your Python project.

Quick Start

import torch
import deep_gemm

# Grouped FP8 GEMM (masked layout for MoE inference)
a = ...  # (num_groups, max_m, k) FP8 tensor
b = ...  # (num_groups, n, k) FP8 tensor
d = torch.empty((num_groups, max_m, n), device='gcu', dtype=torch.bfloat16)
masked_m = ...  # (num_groups,) int tensor with valid row counts
deep_gemm.m_grouped_fp8_gemm_nt_masked(a, b, d, masked_m, expected_m_per_group)

# FP8 MQA attention logits
logits = deep_gemm.fp8_mqa_logits(q, kv, weights, cu_seq_len_k_start, cu_seq_len_k_end)

# FP8/FP4 paged MQA logits (for paged KV cache)
logits = deep_gemm.fp8_fp4_paged_mqa_logits(
    q, kv_cache, weights, context_lens, block_table,
    schedule_meta, max_context_len
)

Interfaces

Grouped GEMMs

API Description
m_grouped_fp8_gemm_nt_contiguous M-grouped FP8 GEMM with contiguous layout
m_grouped_fp8_gemm_nt_masked M-grouped FP8 GEMM with masked layout for MoE decoding

Attention Kernels

API Description
fp8_mqa_logits FP8 multi-query attention logits (non-paged)
fp8_fp4_mqa_logits FP8/FP4 multi-query attention logits (non-paged)
fp8_paged_mqa_logits FP8 paged multi-query attention logits
fp8_fp4_paged_mqa_logits FP8/FP4 paged multi-query attention logits
get_paged_mqa_logits_metadata Compute scheduling metadata for paged MQA
get_num_sms Query number of streaming multiprocessors

Optional Triton Kernels

When triton_kernel_gcu is installed, the following additional APIs are available:

API Description
tf32_hc_prenorm_gemm TF32 high-compute pre-normalized GEMM
transform_sf_into_required_layout Scale factor layout transformation
per_block_cast_to_fp8 Per-block FP8 quantization
fp8_einsum FP8 Einstein summation

Utilities

  • deep_gemm.get_mk_alignment_for_contiguous_layout: Get the group-level alignment requirement for grouped contiguous layout

Testing

Run GCU unit tests:

# Run a specific test
python tests/test_m_grouped_fp8_gemm_masked_random_gcu.py

# Run with pytest
pytest tests/ -v

Acknowledgement

This project is based on deepseek-ai/DeepGEMM. We thank the DeepSeek team for their excellent open-source work.

License

This project is licensed under the MIT License. See the original DeepGEMM LICENSE for details.

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages