diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/.gitignore" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/.gitignore" new file mode 100644 index 00000000..9acbc7c3 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/.gitignore" @@ -0,0 +1,9 @@ +build/ +__pycache__/ +*.pyc +*.so +*.o +*.ncu-rep +outputs/ +inputs/ +*.egg-info/ diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/README.md" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/README.md" new file mode 100644 index 00000000..cd977ae6 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/README.md" @@ -0,0 +1,119 @@ +# 低精度量化/反量化(CUDA 方向) +选题二,实现 NVFP4(e2m1)和 MXFP8(e4m3)两种格式的量化与反量化 +cpu端和gpu端都实现了,以此作为对比 + +## 1. 低精度格式 +NVFP4(e2m1):最大值6 +MXFP8(e4m3):最大值448 + +## 2. 缩放策略 +量化时除以缩放因子s,再将其映射到对应精度的格点上,格点全部算出来存储在数组上 + +### 2.1 量化/反量化公式 +NVFP4:MAX = 6.0,格点表 [0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0] 及负值 +MXFP8:MAX= 448,E4M3 位模式生成 128 个非负格点 +''' +s = |max| / MAX +code = round(x / s) +dequant = TABLE[code] * s +''' + +### 2.2 缩放因子 +tensor:整个张量共享一个缩放因子, s = |全局max| / MAX +block:每个block共享一个缩放因子, s = |块max| / MAX + +## 3. 打包 +NVFP4:每2个4bit的打包成一个字节,低4位在前,高4位在后,一个字节两个元素 +MXFP8:一个元素8bit,所以一个字节就是一个元素 + + +## 4. 运行 +``` +python setup.py build_ext --inplace + +# 生成三种输入矩阵(random / normal / outlier) +python gen_data.py --matrix random +python gen_data.py --matrix normal +python gen_data.py --matrix outlier + +# 跑量化/反量化,--config 选格式,--matrix 选数据分布 +python main.py --config configs/mxfp8_block.txt --matrix random +python main.py --config configs/nvfp4_block.txt --matrix normal +``` + +## 5. 产物 +``` +outputs///cpu/report.log +outputs///cpu/quantized.bin #量化文件 +outputs///cpu/dequant.bin + +outputs///cuda/report.log +outputs///cuda/quantized.bin +outputs///cuda/dequantized.bin +``` +report.log 包含:max_abs_error、mae、mse、compression_ratio、quant_time_s、dequant_time_s、quant_bandwidth_gbps + +## 6. 目录结构 +``` +王玉环/ +├── main.py # 程序入口(CPU/GPU 量化、误差、对比、保存) +├── gen_data.py # 生成输入矩阵 +├── io_utils.py # 文件读写、误差计算、日志 +├── quant_cpu.py # CPU 参考实现 +├── quant_common.py # 格式格点表、转换、打包/解包 +├── quant_cuda.py # CUDA 扩展封装 +├── setup.py # 构建脚本 +├── cuda/ +│ ├── nvfp4.cu # NVFP4 量化/反量化内核 +│ └── mxfp8.cu # MXFP8 量化/反量化内核 +├── configs/ # 参数文件 +├── inputs/ # 输入矩阵 +└── outputs/ # 运行产物 +``` + +## 7. 运行结果(RTX 3090, 4096×4096 FP32 输入) +### MXFP8(block_size=32, scale_mode=tensor) + +| 分布 | max_abs | MAE | MSE | 压缩比 | quant 时间 (GPU) | dequant 时间 (GPU) | 带宽 (GB/s) | +|---|---|---|---|---|---|---|---| +| random | 0.0357 | 0.0111 | 2.17e-4 | 3.56 | 0.0286 s | 0.00167 s | 3.00 | +| normal | 0.1866 | 0.0180 | 7.01e-4 | 3.56 | 0.0274 s | 0.00168 s | 3.14 | +| outlier | 0.2744 | 0.0180 | 7.01e-4 | 3.56 | 0.0264 s | 0.00167 s | 3.26 | + +CPU 端量化约 42 s(GPU 的 ~1500 倍),反量化约 2.9 s。 + +### NVFP4(block_size=16, scale_mode=block) + +| 分布 | max_abs | MAE | MSE | 压缩比 | quant 时间 (GPU) | dequant 时间 (GPU) | 带宽 (GB/s) | +|---|---|---|---|---|---|---|---| +| random | 0.167 | 0.0429 | 3.38e-3 | 5.33 | 0.0241 s | 0.00512 s | 3.30 | +| normal | 0.684 | 0.0686 | 8.86e-3 | 5.33 | 0.0229 s | 0.00488 s | 3.48 | +| outlier | 3.169 | 0.0686 | 8.86e-3 | 5.33 | 0.0223 s | 0.00482 s | 3.57 | + +CPU 端量化约 8.5 s(GPU 的 ~350 倍),反量化约 2.9 s。 + +## 8. 实现说明 + +### 软件模拟部分(不依赖特定硬件) + +- **E2M1 / E4M3 编码**:纯位运算 + 查找表,普通 CUDA kernel 实现,不使用 FP8/FP4 原生指令 +- **量化 kernel**:nearest rounding,block max 由线程串行计算,tensor max 由 host 端归约后传入 +- **打包/解包**:手动位操作,4bit 元素每两个打包成一个字节 +- **反量化 kernel**:查表 + 乘 scale,普通 CUDA kernel +- 以上代码我是在RTX3090上运行的 + +### 依赖的第三方库 + +- **PyTorch**:使用 `torch.utils.cpp_extension` 构建 CUDA extension,Tensor 作为数据容器;host 端全局 max 调用 `Tensor::abs().max().item()` +- **CUDA Runtime**:仅使用标准 CUDA Runtime API(kernel launch、`cudaMemcpyToSymbol`),未使用 cuBLAS/cuDNN/cuBLASLt 等 + +## 9.Nsight Compute 性能分析 + +使用 `ncu --set full --kernel-name regex:quant_` 对四个 kernel 进行 profile: + +| kernel | Duration | Memory Throughput | Compute Throughput | Achieved BW | Occupancy | +|---|---|---|---|---|---| +| quant_mxfp8 | 2.69 ms | 48.7% | 38.8% | 389 GB/s | 94.9% | +| dequant_mxfp8 | 2.77 ms | 52.7% | 6.2% | 334 GB/s | 92.7% | +| quant_nvfp4 | 0.17 ms | 91.9% | 38.2% | 508 GB/s | 82.9% | +| dequant_nvfp4 | 0.44 ms | 80.6% | 5.5% | 451 GB/s | 92.4% | \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/mxfp8_block.txt" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/mxfp8_block.txt" new file mode 100644 index 00000000..1c5ba336 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/mxfp8_block.txt" @@ -0,0 +1,6 @@ +# MXFP8 配置 +format = mxfp8 +block_size = 32 +scale_mode = tensor +output_type = fp32 +rounding = nearest \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/nvfp4_block.txt" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/nvfp4_block.txt" new file mode 100644 index 00000000..a915eb6b --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/configs/nvfp4_block.txt" @@ -0,0 +1,6 @@ +# NVFP4 配置 +format = nvfp4 +block_size = 16 +scale_mode = block +output_type = fp16 +rounding = nearest \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/__init__.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/__init__.py" new file mode 100644 index 00000000..e69de29b diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/mxfp8.cu" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/mxfp8.cu" new file mode 100644 index 00000000..3de5c37f --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/mxfp8.cu" @@ -0,0 +1,143 @@ +#include +#include +#include +#include +#include + +__constant__ float _E4M3_TABLE[128]; + +__device__ int _float_to_e4m3(float val) { + int idx = 0; + if (val == 0) return idx; + int sign = (val < 0) ? 0x80 : 0; + float a =fabsf(val); + int best_idx = 0; + float best_d = fabsf(a - _E4M3_TABLE[0]); + for (idx = 1; idx < 128; idx++) { + float d = fabsf(a - _E4M3_TABLE[idx]); + if (d < best_d) { + best_d = d; + best_idx = idx; + } + } + return sign | best_idx; +} + +__global__ void quant_mxfp8_kernel( + const float* data, + float* scales, + uint8_t* packed, + const int rows, const int cols, + const int block_size, + const int scale_mode, + const float gmax +) { + int r = blockIdx.y * blockDim.x + threadIdx.x; + int c = blockIdx.x; + int stride = r*cols + c*block_size; + int num = (cols + block_size - 1) / block_size; // 这个是缩放因子的个数 + + if (r >= rows) return; + + float global_scale = 0.0f; + if (scale_mode) { + global_scale = (gmax > 0) ? (gmax / 448.0) : 1.0; + } else { + float bmax = 0.0f; + for (int i = 0; i < block_size; i++) { + float temp = fabsf(data[i+stride]); + if (bmax < temp) bmax = temp; + } + global_scale = (bmax > 0) ? (bmax / 448.0) : 1.0; + } + scales[r * num + c] = global_scale; + + // 量化 + for (int i = 0; i < block_size; i++) { + packed[r*cols + c*block_size + i] = _float_to_e4m3(data[i+stride] / scales[r*num + c]); + } +} + +__global__ void dequant_mxfp8_kernel( + const uint8_t* packed, + float* out, + const float* scales, + const int rows, const int cols, const int block_size +) { + int r = blockIdx.y * blockDim.x + threadIdx.x; + int c = blockIdx.x; + + if (r >= rows) return; + + int stride = r*cols + c*block_size; + int num = (cols + block_size - 1) / block_size; + float s = scales[r*num + c]; + + for (int i = 0; i < block_size; i++) { + int code = packed[r*cols + c*block_size + i]; + float val = _E4M3_TABLE[code & 0x7F]; + out[i + stride] = s * ((code & 0x80) ? -val : val); + } +} + +std::tuple quant_mxfp8( + torch::Tensor data, + int64_t rows, int64_t cols, int64_t block_size, + int64_t scale_mode +) { + int n = (cols + block_size - 1) / block_size; + auto scales = torch::empty({rows*n}, data.options()); + auto packed = torch::empty({rows*cols}, data.options().dtype(torch::kUInt8)); + + float host_table[128]; + for (int i = 0; i < 128; i++) { + int e = (i >> 3) & 0xF; + int m = i & 0x7; + if (e == 15 && m == 7) host_table[i] = 448.0f; + else if (e == 0) host_table[i] = m * powf(2, -9); + else host_table[i] = (1.0f + m / 8.0f) * powf(2, e-7); + } + cudaMemcpyToSymbol(_E4M3_TABLE, host_table, sizeof(host_table)); + + //求一下最大值 + float gmax = 0.0f; + gmax = data.abs().max().item(); + + //同样的一个线程处理一个block_size的元素 + dim3 block(256); + dim3 grid(n, (rows + 255)/256, 1); + quant_mxfp8_kernel<<>>( + data.data_ptr(), + scales.data_ptr(), + packed.data_ptr(), + (int)rows, (int)cols, (int)block_size, + (int)scale_mode, gmax + ); + return {packed, scales}; +} + +torch::Tensor dequant_mxfp8( + torch:: Tensor packed, + torch::Tensor scales, + int64_t rows, int64_t cols, + int64_t block_size, std::string& output_type +) { + auto out = torch::empty({rows*cols}, scales.options().dtype(torch::kFloat32)); + dim3 block(256); + dim3 grid((int)((cols + block_size - 1) / block_size),(rows + 255) / 256, 1); + dequant_mxfp8_kernel<<>>( + packed.data_ptr(), + out.data_ptr(), + scales.data_ptr(), + (int)rows, (int)cols, (int)block_size + ); + + if (output_type == "fp16") return out.to(torch::kHalf); + if (output_type == "bf16") return out.to(torch::kBFloat16); + return out; +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("quant", &quant_mxfp8, "MXFP8 quantize -> (packed, scales)"); + m.def("dequant", &dequant_mxfp8, "MXFP8 dequantize -> dequant"); +} \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/nvfp4.cu" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/nvfp4.cu" new file mode 100644 index 00000000..e548874c --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/cuda/nvfp4.cu" @@ -0,0 +1,142 @@ +#include +#include +#include + +__device__ const float _E2M1_TABLE[16] = {0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, + 0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0}; + +__device__ int _float_to_e2m1(float val) { + if (val == 0) return 0; + float a = fabsf(val); + + int idx = 0; + if (a >= 5.0) idx = 7; + else if (a >= 3.5) idx = 6; + else if (a >= 2.5) idx = 5; + else if (a >= 1.75) idx = 4; + else if (a >= 1.25) idx = 3; + else if (a >= 0.75) idx = 2; + else if (a >= 0.25) idx = 1; + else idx = 0; + return (val > 0) ? idx : idx + 8; + +} + +__global__ void quant_nvfp4_kernel( + const float* data, + uint8_t* packed, + float* scales, + const int rows, const int cols, + const int block_size, + const int scale_mode, + const float gmax +) { + int r = blockIdx.x; + int c = threadIdx.x; + + if (r >= rows) return; + int stride = r * cols + c * block_size; + int num = (cols + block_size -1) / block_size; + + float global_scale = 0.0f; + if (scale_mode) { + global_scale = (gmax > 0) ? (gmax / 6.0) : 1.0; + } else { + float bmax = 0.0f; + for (int i = 0; i < block_size; i++) { + if (bmax < fabsf(data[i+stride])) bmax = fabsf(data[i+stride]); + } + global_scale = (bmax > 0) ? (bmax / 6.0) : 1.0; + } + scales[r * num + c] = global_scale; + + // 量化 + 打包 + int half = cols / 2; + for (int i = 0; i < block_size; i += 2) { + float x0 = data[i+stride]; + float x1 = data[i+stride+1]; + int code0 = _float_to_e2m1(x0 / scales[r * num + c]); + int code1 = _float_to_e2m1(x1 / scales[r * num + c]); + packed[r*half + c*(block_size/2) + i/2] = (uint8_t)(code0 & 0xF) | ((code1 & 0xF) << 4); + } +} + + +std::tuple quant_nvfp4( + torch::Tensor data, int64_t rows, int64_t cols, + int64_t block_size, int64_t scale_mode +) { + int num = (int)(cols + block_size -1) / block_size; + int scale_num = (int)(rows * num); + auto scales = torch::empty({scale_num}, data.options()); + // 打包的时候注意一个字节可以放两个元素 + auto packed = torch::empty({rows * (cols/2)}, data.options().dtype(torch::kUInt8)); + + //求一下全局最大值 + float gmax = 0.0f; + gmax = data.abs().max().item(); + + // 一个线程一次性处理blocksize个元素 + dim3 block((cols + block_size - 1) / block_size); + dim3 grid(rows, 1, 1); + quant_nvfp4_kernel<<>>( + data.data_ptr(), + packed.data_ptr(), + scales.data_ptr(), + (int)rows, (int)cols, + (int)block_size, (int)scale_mode, gmax + ); + + return {packed, scales}; +} + +__global__ void dequant_nvfp4_kernel( + const uint8_t* packed, + float* out, + const float* scales, + const int rows, const int cols, + const int block_size +) { + int r = blockIdx.x; + int c = threadIdx.x; + + if (r >= rows) return; + int stride = r * cols + c * block_size; + int num = (cols + block_size -1) / block_size; + int half = cols / 2; + float s = scales[r*num + c]; + //拆包 + for (int i = 0; i < block_size; i += 2) { + int codes = packed[r*half + c*(block_size/2) + i/2]; + int code0 = codes & 0xF; + int code1 = (codes >> 4) & 0xF; + out[i + stride] = _E2M1_TABLE[code0] * s; + out[i + stride + 1] = _E2M1_TABLE[code1] * s; + } +} + +torch::Tensor dequant_nvfp4( + torch::Tensor packed, + torch::Tensor scales, + int64_t rows, int64_t cols, + int64_t block_size, std::string& output_type +) { + auto out = torch::empty({rows*cols}, scales.options().dtype(torch::kFloat32)); + dim3 block((cols + block_size - 1) / block_size); + dim3 grid(rows, 1, 1); + dequant_nvfp4_kernel<<>>( + packed.data_ptr(), + out.data_ptr(), + scales.data_ptr(), + (int)rows, (int)cols, (int)block_size + ); + if (output_type == "fp16") return out.to(torch::kHalf); + if (output_type == "bf16") return out.to(torch::kBFloat16); + return out; +} + + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("quant", &quant_nvfp4, "NVFP4 quantize -> (packed, scales)"); + m.def("dequant", &dequant_nvfp4, "NVFP4 dequantize -> dequant"); +} \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/gen_data.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/gen_data.py" new file mode 100644 index 00000000..0057ead7 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/gen_data.py" @@ -0,0 +1,74 @@ +'''生成三种类型矩阵''' +import os +import argparse +import random +import struct + +# 处理输入的数据类型 +def _cast(data, dtype): + if dtype == "fp16": + out = [] + for row in data: + new_row = [] + for x in row: + try: + val = struct.unpack("e", struct.pack("e", float(x)))[0] + except OverflowError: + val = 65504.0 if x > 0 else -65504.0 + new_row.append(val) + out.append(new_row) + return out + + return data + + +# 随机矩阵 +def _gen_random(rows, cols, dtype = "fp32"): + rr = random.Random(0) + data = [[rr.random() * 2 - 1 for _ in range(cols)] for _ in range(rows)] + return _cast(data, dtype) + + +# 正态分布矩阵 +def _gen_normal(rows, cols, dtype = "fp32"): + rr = random.Random(1) + data = [[rr.gauss(0.0, 1.0) for _ in range(cols)] for _ in range(rows)] + return _cast(data, dtype) + +# 含有异常值矩阵 +def _gen_outlier(rows, cols, dtype = "fp32"): + rr = random.Random(2) + data = [[rr.gauss(0.0, 1.0) for _ in range(cols)] for _ in range(rows)] + # 随机挑选几个异常值 + for _ in range(8): + r, c = rr.randrange(rows), rr.randrange(cols) + data[r][c] = 1000.0 if rr.choice([True, False]) else -1000.0 + return _cast(data, dtype) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--matrix", choices=["random", "normal", "outlier"], help="生成三种矩阵", default="random") + parser.add_argument("--rows", type=int, default=4096) + parser.add_argument("--cols", type=int, default=4096) + parser.add_argument("--dtype", choices=["fp32", "fp16"], default="fp32") + args = parser.parse_args() + out = "inputs" + os.makedirs(out, exist_ok=True) + path = os.path.join(out, f"{args.matrix}.bin") + + if args.matrix == "random": + data = _gen_random(args.rows, args.cols, args.dtype) + elif args.matrix == "normal": + data = _gen_normal(args.rows, args.cols, args.dtype) + else: + data = _gen_outlier(args.rows, args.cols, args.dtype) + with open(path, "wb") as f: + f.write(f"num_rows: {args.rows}\nnum_cols: {args.cols}\ndtype: {args.dtype}\n[data]\n".encode()) + fmt = "f" if args.dtype == "fp32" else "e" + for row in data: + for x in row: + f.write(struct.pack(fmt, float(x))) + +if __name__ == "__main__": + main() \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/io_utils.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/io_utils.py" new file mode 100644 index 00000000..7f66eb19 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/io_utils.py" @@ -0,0 +1,84 @@ +'''处理一些输入输出文件,以及保存''' + +import struct + +def read_matrix_file(path): + with open(path, "rb") as f: + rows = cols = 0 + dtype = "fp32" + while True: + line = f.readline() + if not line: break + s = line.decode().strip() + if s.startswith("num_rows:"): + rows = int(s.split(":")[1].strip()) + elif s.startswith("num_cols:"): + cols = int(s.split(":")[1].strip()) + elif s.startswith("dtype:"): + dtype = s.split(":")[1].strip() + elif s == "[data]": + break + n = rows * cols + fmt = "f" if dtype == "fp32" else "e" + size = 4 if dtype == "fp32" else 2 + data = f.read(n * size) + data = [struct.unpack_from(fmt, data, i * size)[0] for i in range(n)] + + return data, rows, cols, dtype + +def read_config(path): + cfg = {} + with open(path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line or line.startswith("#"): + continue + key, value = line.split("=", 1) + cfg[key.strip()] = value.strip() + return cfg + +def save_file(path, scales, packed, rows, cols, block_size, format): + with open(path, "wb") as f: + f.write(("format: %s\n" % format).encode()) + f.write(("rows: %d\n" % rows).encode()) + f.write(("cols: %d\n" % cols).encode()) + f.write(("block_size: %d\n" % block_size).encode()) + f.write(("packed_len: %d\n" % len(packed)).encode()) + f.write(("scale_count: %d\n" % len(scales)).encode()) + f.write(b"[packed]\n") + f.write(bytes(packed)) + f.write(b"[scales]\n") + for s in scales: + f.write(struct.pack("> 3) & 0xF + m = code & 0x07 + if e == 15 and m == 7: + val = 448.0 + elif e == 0: + # 次正规数 + val = (m/8.0) * (2.0 ** -6) + # 正规值 + else: val = (1.0 + m/8.0) * (2.0 ** (e-7)) + _E4M3_TABLE.append(val) + +def float_to_e2m1(x, rng): + if x == 0.0: + return 0 + a = abs(x) + low_idx = 0 + for i in range(1, 8): + if a >= _E2M1_TABLE[i]: low_idx = i + high_idx = min(low_idx+1, 7) + low_val = _E2M1_TABLE[low_idx] + high_val = _E2M1_TABLE[high_idx] + + if rng is not None: + # stochastic + pi = (a - low_val) / (high_val - low_val) + idx = high_idx if rng.random() < pi else low_idx + else: + # nearest:离谁近选谁 + idx = high_idx if (a - low_val) >= (high_val - a) else low_idx + return idx if x > 0 else idx + 8 + +def float_to_e4m3(x, rng): + if x == 0.0: + return 0 + + sign = 0x80 if x < 0 else 0 + a = abs(x) + low_idx = 0 + for i in range(1, 128): + if a >= _E4M3_TABLE[i]: low_idx = i + high_idx = min(low_idx+1, 127) + low_val = _E4M3_TABLE[low_idx] + high_val = _E4M3_TABLE[high_idx] + + if rng is not None: + pi = (a - low_val) / (high_val - low_val) + idx = high_idx if rng.random() < pi else low_idx + else: + idx = high_idx if (a - low_val) > (high_val - a) else low_idx + + return sign | idx + +def pack(codes, element_bits): + if element_bits == 8: return bytes(codes) + if len(codes) % 2 == 1: + codes = codes + [0] # 奇数个就补一个 0 + out = bytearray() + for i in range(0, len(codes), 2): + out.append((codes[i] & 0xF) | ((codes[i + 1] & 0xF) << 4)) + return bytes(out) + +def unpack(packed, element_bits, numbers): + if element_bits == 8: + return list(packed[:numbers]) + codes = [] + for byte in packed: + codes.append(byte & 0xF) + codes.append((byte >> 4) & 0xF) + return codes[:numbers] + +def e2m1_to_float(idx): + return _E2M1_TABLE[idx] + +def e4m3_to_float(idx): + val = _E4M3_TABLE[idx & 0x7F] + return -val if (idx & 0x80) else val diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cpu.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cpu.py" new file mode 100644 index 00000000..9010cdc8 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cpu.py" @@ -0,0 +1,46 @@ +import quant_common as qc +import random + +def quant(data, rows, cols, format, block_size, scale_mode, rounding, seed = 2026): + codes = [] + scales = [] + rng = random.Random(seed) if rounding == "stochastic" else None + global_scale = None + if scale_mode == "tensor": + gmax = max(abs(x) for x in data) + global_scale = gmax / qc.MAX_FENMU[format] if gmax > 0 else 1.0 + + for r in range(rows): + for b in range(0, cols, block_size): + start = r*cols + b + blk = data[start:start+block_size] + + if scale_mode == "block": + bmax = max(abs(x) for x in blk) + s = bmax / qc.MAX_FENMU[format] if bmax > 0 else 1.0 + else: s = global_scale + scales.append(s) + + # 对应的数字 + for x in blk: + if format == "nvfp4": + codes.append(qc.float_to_e2m1(x / s, rng)) + else: codes.append(qc.float_to_e4m3(x / s, rng)) + + element_bits = 4 if format == "nvfp4" else 8 + return qc.pack(codes, element_bits), scales + +def dequant(packed, scales, rows, cols, format, block_size): + element_bits = 4 if format == "nvfp4" else 8 + codes = qc.unpack(packed, element_bits, rows*cols) + + # 反量化直接乘上缩放因子就行 + out = [] + for i, code in enumerate(codes): + r = i // cols + c = i % cols + s = scales[i // block_size] + if format == "nvfp4": + out.append(qc.e2m1_to_float(code) * s) + else: out.append(qc.e4m3_to_float(code) * s) + return out diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cuda.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cuda.py" new file mode 100644 index 00000000..20557293 --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/quant_cuda.py" @@ -0,0 +1,22 @@ +import torch +from cuda import nvfp4_cuda, mxfp8_cuda + +def quant_dequant(vals, rows, cols, format, scale_mode, block_size, output_type): + data = torch.tensor(vals, dtype=torch.float32, device="cuda") + + fmt = nvfp4_cuda if format == "nvfp4" else mxfp8_cuda + scale_mode = 1 if scale_mode == "tensor" else 0 + + start = torch.cuda.Event(enable_timing=True) + mid = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + packed, scales = fmt.quant(data, rows, cols, block_size, scale_mode) + mid.record() + dequant = fmt.dequant(packed, scales, rows, cols, block_size, output_type) + end.record() + end.synchronize() + t_quant = start.elapsed_time(mid) / 1000.0 + t_dequant = mid.elapsed_time(end) / 1000.0 + + return packed.cpu().numpy().tobytes(), scales.cpu().tolist(), dequant.cpu().tolist(), t_quant, t_dequant \ No newline at end of file diff --git "a/02_quant_dequant/\347\216\213\347\216\211\347\216\257/setup.py" "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/setup.py" new file mode 100644 index 00000000..751dddbb --- /dev/null +++ "b/02_quant_dequant/\347\216\213\347\216\211\347\216\257/setup.py" @@ -0,0 +1,12 @@ +from setuptools import setup, find_packages +from torch.utils.cpp_extension import BuildExtension, CUDAExtension + +setup( + name="cuda", + packages=find_packages(), + ext_modules=[ + CUDAExtension("cuda.nvfp4_cuda", ["cuda/nvfp4.cu"]), + CUDAExtension("cuda.mxfp8_cuda", ["cuda/mxfp8.cu"]), + ], + cmdclass={"build_ext": BuildExtension} +) \ No newline at end of file