muFFT 开发者指南
muFFT 是基于 MUSA 开发的离散傅里叶变换库,在 MTGPU 上经过深度优化,在深度学习、计算物理、分子动力学、量子化学、地震和医学成像等领域被广泛使用。
概述
什么是 muFFT
muFFT(MUSA Fast Fourier Transform,快速傅里叶变换库)提供基于 MUSA 的快速傅里叶变换接口,支持在 MTGPU 上执行一维、二维和三维 FFT,可用于信号处理、科学计算和深度学习等工作负载。
关键特性
- 支持 R2C、C2C、C2R 等常用 FFT 变换类型。
- 支持一维、二维、三维以及批处理 FFT。
- 提供 plan 创建、执行和参数配置接口,便于复用变换配置。
- 针对 MTGPU 架构优化,适合高吞吐频域计算场景。
支持的变换类型
| 变换类型 | 说明 | 精度 |
|---|---|---|
| R2C | 实数到复数 (Real to Complex) | 单精度/双精度 |
| C2C | 复数到复数 (Complex to Complex) | 单精度/双精度 |
| C2R | 复数到实数 (Complex to Real) | 单精度/双精度 |
支持的维度
- 一维 FFT (1D)
- 二维 FFT (2D)
- 三维 FFT (3D)
批处理支持
支持批处理 FFT,可同时处理多个相同尺寸的变换。
快速开始
一维复数到复数变换 (1D C2C)
#include <musa_runtime.h>
#include <mufft.h>
#include <complex>
#include <iostream>
#include <vector>
int main() {
std::cout << "muFFT 1D single-precision complex-to-complex transform\n";
const int Nx = 8;
std::vector<std::complex<float>> h_x(Nx);
// 初始化输入数据
for (size_t i = 0; i < Nx; i++) {
h_x[i] = std::complex<float>(i, 0);
}
std::cout << "Input: ";
for (size_t i = 0; i < Nx; i++) {
std::cout << h_x[i] << " ";
}
std::cout << std::endl;
// 创建 MUSA 设备对象并复制数据到设备
size_t complex_bytes = sizeof(decltype(h_x)::value_type) * h_x.size();
mufftComplex* d_x;
musaMalloc(&d_x, complex_bytes);
musaMemcpy(d_x, h_x.data(), complex_bytes, musaMemcpyHostToDevice);
// 创建 plan
mufftHandle plan;
mufftPlan1d(&plan, // plan handle
Nx, // 变换长度
MUFFT_C2C, // 变换类型 (MUFFT_C2C 单精度)
1); // 变换数量
// 执行 plan (正向变换)
mufftExecC2C(plan, d_x, d_x, MUFFT_FORWARD);
// 复制结果回主机
musaMemcpy(h_x.data(), d_x, complex_bytes, musaMemcpyDeviceToHost);
std::cout << "Output: ";
for (size_t i = 0; i < Nx; i++) {
std::cout << h_x[i] << " ";
}
std::cout << std::endl;
// 释放资源
mufftDestroy(plan);
musaFree(d_x);
return 0;
}
编译和运行
# 使用 mcc 编译
mcc fft1d_example.cpp \
-L/usr/local/musa/lib \
-I/usr/local/musa/include \
-lmufft -lmusart \
-o fft1d_example
# 运行
./fft1d_example
API 参考
创建 Plan
一维 FFT Plan
mufftResult_t mufftPlan1d(mufftHandle* plan,
int nx,
mufftType_t type,
int batch);
参数:
| 参数 | 说明 |
|---|---|
plan | 输出的 plan handle |
nx | X 维度的尺寸 |
type | 变换类型 (MUFFT_C2C, MUFFT_R2C, MUFFT_C2R) |
batch | 批处理数量 |
二维 FFT Plan
mufftResult_t mufftPlan2d(mufftHandle* plan,
int nx,
int ny,
mufftType_t type);
三维 FFT Plan
mufftResult_t mufftPlan3d(mufftHandle* plan,
int nx,
int ny,
int nz,
mufftType_t type);