muBLAS 开发者指南
muBLAS 是基于 MUSA 开发的基础线性代数运算库,在 MTGPU 上经过深度优化,在 AI 和 HPC 场景下被广泛使用。
概述
什么是 muBLAS
muBLAS(MUSA Basic Linear Algebra Subprograms,基础线性代数算子库)提供 BLAS 标准风格的向量、矩阵与矩阵乘法接口,开发者可以通过统一的 handle、stream 和设备内存模型调用 MTGPU 上优化后的线性代数算子。
关键特性
- 覆盖 BLAS Level 1、Level 2 和 Level 3 常用计算。
- 支持单精度、双精度、复数等基础数据类型。
- 可与 MUSA stream 配合使用, 支持异步执行和多 stream 并发。
- 针对 MTGPU 架构优化,适用于 AI、HPC 和科学计算场景。
功能分类
按照计算复杂性,muBLAS 函数可分为三类:
| 类别 | 说明 | 示例函数 |
|---|---|---|
| 第一类 | 标量、向量和向量与向量间的运算 | axpy, dot, nrm2, scal |
| 第二类 | 向量与矩阵之间的运算 | gemv, ger |
| 第三类 | 矩阵与矩阵间的运算 | gemm, trsm |
数据类型支持
muBLAS 支持以下数据类型:
| 类型前缀 | 数据类型 |
|---|---|
s | 单精度浮点 (float) |
d | 双精度浮点 (double) |
c | 单精度复数 (complex float) |
z | 双精度复数 (complex double) |
快速开始
完整示例:SAXPY
SAXPY: (单精度 AX 加 Y)
#include <cstdio>
#include <cstdlib>
#include <musa_runtime.h>
#include <mublas.h>
#include <vector>
int main(int argc, char* argv[]) {
mublasHandle_t mublasH = NULL;
musaStream_t stream = NULL;
// 初始化数据
// A = [1.0, 2.0, 3.0, 4.0]
// B = [5.0, 6.0, 7.0, 8.0]
const std::vector<float> A = {1.0, 2.0, 3.0, 4.0};
std::vector<float> B = {5.0, 6.0, 7.0, 8.0};
const float alpha = 2.1;
const int incx = 1;
const int incy = 1;
float* d_A = nullptr;
float* d_B = nullptr;
// Step 1: 创建 mublas handle,绑定 stream
mublasCreate(&mublasH);
musaStreamCreateWithFlags(&stream, musaStreamNonBlocking);
mublasSetStream(mublasH, stream);
// Step 2: 复制数据到设备
musaMalloc(reinterpret_cast<void**>(&d_A), sizeof(float) * A.size());
musaMalloc(reinterpret_cast<void**>(&d_B), sizeof(float) * B.size());
musaMemcpyAsync(d_A, A.data(), sizeof(float) * A.size(),
musaMemcpyHostToDevice, stream);
musaMemcpyAsync(d_B, B.data(), sizeof(float) * B.size(),
musaMemcpyHostToDevice, stream);
// Step 3: 计算 B = alpha * A + B
// B = 2.1 * [1, 2, 3, 4] + [5, 6, 7, 8]
// B = [7.1, 10.2, 13.3, 16.4]
mublasSaxpy(mublasH, A.size(), &alpha, d_A, incx, d_B, incy);
// Step 4: 复制数据回主机
musaMemcpyAsync(B.data(), d_B, sizeof(float) * B.size(),
musaMemcpyDeviceToHost, stream);
musaStreamSynchronize(stream);
// 输出结果
printf("B = [");
for(int i = 0; i < B.size(); i++) {
printf("%f ", B[i]);
}
printf("]\n");
// 预期输出:B = [7.100000 10.200000 13.300000 16.400000]
// 释放资源
musaFree(d_A);
musaFree(d_B);
mublasDestroy(mublasH);
musaStreamDestroy(stream);
musaDeviceReset();
return EXIT_SUCCESS;
}
编译和运行
# 使用 mcc 编译
mcc saxpy_example.cpp \
-L/usr/local/musa/lib \
-I/usr/local/musa/include \
-lmublas -lmusart \
-o saxpy_example
# 运行
./saxpy_example
或使用 CMake:
cmake_minimum_required(VERSION 3.10)
project(SAXPY_Example LANGUAGES CXX)
add_executable(saxpy_example saxpy_example.cpp)
target_include_directories(saxpy_example PRIVATE /usr/local/musa/include)
target_link_directories(saxpy_example PRIVATE /usr/local/musa/lib)
target_link_libraries(saxpy_example PRIVATE mublas musart)
set_property(TARGET saxpy_example PROPERTY CXX_STANDARD 17)
编译步骤:
# 1. 创建构建目录并进入
mkdir build && cd build
# 2. 配置 CMake(需在包含 CMakeLists.txt 的目录下执行)
cmake ..
# 3. 编译
cmake --build .
# 或使用 make / ninja,与项目约定一致即可
注意:若 MUSA 未安装在默认路径 /usr/local/musa,需将 target_include_directories 和 target_link_directories 中的路径改为实际安装路径。
API 参考
向量与向量运算 (Level 1)
axpy - 向量线性组合
mublasStatus mublasSaxpy(mublasHandle_t handle,
int n,
const float* alpha,
const float* x,
int incx,
float* y,
int incy);
功能:
参数:
| 参数 | 说明 |
|---|---|
handle | muBLAS 句柄 |
n | 向量元素数量 |
alpha | 标量乘数 |
x | 输入向量 x |
incx | x 的步长 |
y | 输入/输出向量 y |
incy | y 的步长 |