当前位置: 首页 > news >正文

BLOCK_M, BLOCK_N 是干什么的

一句话

BLOCK_M / BLOCK_N = 每个 Block(Program)负责计算的那一小块子矩阵的尺寸。

它把一个大矩阵切成很多小 tile,每个 tile 分给一个 Block 去算。


以矩阵乘法为例

C = A × B A: [M, K] B: [K, N] C: [M, N] 比如 M=1024, N=1024, K=512

不可能一个 Block 算整个 1024×1024 的 C,所以切块

设 BLOCK_M = 128, BLOCK_N = 128 C [1024 × 1024] 被切成: BLOCK_N=128 ├────┤ ┌────┬────┬────┬────┬────┬────┬────┬────┐ │ │ │ │ │ │ │ │ │ ↑ │(0,0)│(0,1)│(0,2)│(0,3)│(0,4)│(0,5)│(0,6)│(0,7)│ │ │ │ │ │ │ │ │ │ │ │ ├────┼────┼────┼────┼────┼────┼────┼────┤ │ │ │ │ │ │ │ │ │ │ │ │(1,0)│(1,1)│ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ ├────┼────┼────┼────┼────┼────┼────┼────┤ │ BLOCK_M=128 │ │ │ │ │ │ │ │ │ │ │(2,0)│ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ ├────┼────┼────┼────┼────┼────┼────┼────┤ │ │ │ │ │ │ │ │ │ │ │ │... │ │ │ │ │ │ │ │ │ └────┴────┴────┴────┴────┴────┴────┴────┘ ↓ 共 (1024/128) × (1024/128) = 8 × 8 = 64 个 tile → 启动 64 个 Block

每个 Block 做什么?

Block (pid_m=2, pid_n=3) 的任务: 计算 C 的第 2 行 tile、第 3 列 tile: C[256:384, 384:512] = A[256:384, :] × B[:, 384:512] ↑ 128行 ↑ 128列 但 K 方向也要分块(BLOCK_K),迭代累加: for k in range(0, K, BLOCK_K): # K=512, BLOCK_K=32 → 迭代16次 加载 A_tile [128 × 32] 到 smem 加载 B_tile [32 × 128] 到 smem C_tile += A_tile × B_tile # [128×32] × [32×128] = [128×128]
┌─────────────────────────────────────────────────┐ │ Block (2,3) 的工作: │ │ │ │ A [128 × 512] B [512 × 128] │ │ ┌──┬──┬──┬──┬──┐ ┌──┐ │ │ │32│32│32│32│..│ │ │ │ │ │ │ │ │ │ │ × │ │ = C_tile [128×128]│ │ │ │ │ │ │ │ │ │ │ │ └──┴──┴──┴──┴──┘ └──┘ │ │ ←── BLOCK_K=32 ──→ ↑ │ │ 每次加载一小条 每次加载一小条 │ │ 迭代 16 次累加 │ └─────────────────────────────────────────────────┘

对应到代码

CUDA

#define BLOCK_M 128 #define BLOCK_N 128 #define BLOCK_K 32 __global__ void matmul(float* A, float* B, float* C, int M, int N, int K) { // 我是哪个 tile? int pid_m = blockIdx.y; int pid_n = blockIdx.x; // 我负责 C 的哪一块? int row_start = pid_m * BLOCK_M; // 比如 256 int col_start = pid_n * BLOCK_N; // 比如 384 // 在 K 方向迭代 float acc[BLOCK_M][BLOCK_N] = {0}; // 每个线程负责 acc 的一小部分 for (int k = 0; k < K; k += BLOCK_K) { // 加载 A[row_start : row_start+128, k : k+32] 到 smem // 加载 B[k : k+32, col_start : col_start+128] 到 smem // acc += A_smem × B_smem } // 写回 C[row_start:row_start+128, col_start:col_start+128] } // 启动 dim3 block(256); dim3 grid(N / BLOCK_N, M / BLOCK_M); // (8, 8) = 64 个 block matmul<<<grid, block>>>(A, B, C, M, N, K);

Triton

@triton.jit def matmul_kernel( A_ptr, B_ptr, C_ptr, M, N, K, BLOCK_M: tl.constexpr, # 128 BLOCK_N: tl.constexpr, # 128 BLOCK_K: tl.constexpr, # 32 ): pid_m = tl.program_id(0) pid_n = tl.program_id(1) # 我负责 C 的 [pid_m*128 : pid_m*128+128, pid_n*128 : pid_n*128+128] acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for k in range(0, K, BLOCK_K): a = tl.load(A_ptr + ...) # [BLOCK_M, BLOCK_K] = [128, 32] b = tl.load(B_ptr + ...) # [BLOCK_K, BLOCK_N] = [32, 128] acc += tl.dot(a, b) # [128, 128] tl.store(C_ptr + ..., acc) # 启动 grid = (M // BLOCK_M, N // BLOCK_N) # (8, 8) matmul_kernel[grid](A, B, C, M, N, K, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32)

在 Attention 中的含义

Attention: O = softmax(Q × K^T) × V Q: [seq_len, head_dim] 比如 [4096, 576] K: [seq_len, head_dim] 比如 [4096, 576] V: [seq_len, head_dim_v] 比如 [4096, 512] BLOCK_M = 64 → 每个 Block 处理 64 个 Query token BLOCK_N = 64 → 每次加载 64 个 KV token(= TOPK_BLOCK_SIZE) ┌──────────────────────────────────────────────────────────┐ │ Q [4096 × 576] │ │ ┌──┐ │ │ │64│ ← BLOCK_M: 这个 Block 负责的 64 个 query │ │ └──┘ │ │ │ │ × K^T [576 × 4096] │ │ ┌──┬──┬──┬──┬──┬──┐ │ │ │64│64│64│64│..│64│ ← BLOCK_N: 每次加载64个KV │ │ └──┴──┴──┴──┴──┴──┘ │ │ │ │ = Score [64 × 4096] → softmax → × V → O [64 × 512] │ │ │ │ 迭代 4096/64 = 64 次(或 TopK 后只迭代 32 次) │ └──────────────────────────────────────────────────────────┘

BLOCK_M / BLOCK_N 怎么选?

考虑因素BLOCK 大BLOCK 小
Shared Memory 用量大(可能超限)
计算访存比高(好)低(差)
Grid 大小(并行度)小(可能填不满 SM)大(好)
Register 压力大(每线程累加器多)
典型值128 / 25616 / 32
经验法则: BLOCK_M × BLOCK_N × sizeof(float) ≈ 每个线程的累加器大小 比如 BLOCK_M=128, BLOCK_N=128, 256 个线程: 每线程累加器 = 128×128 / 256 = 64 个 float = 64 个 register 加上其他变量,总共 ~128 registers/thread → 合理

总结

BLOCK_M = 每个 Block 在 M 方向(行/query)上处理多大 BLOCK_N = 每个 Block 在 N 方向(列/KV)上处理多大 BLOCK_K = 在 K 方向(reduction/内积)上每次加载多大 它们决定了: 1. Grid 大小 = (M/BLOCK_M) × (N/BLOCK_N) → 启动多少个 Block 2. Shared Memory 大小 = BLOCK_M×BLOCK_K + BLOCK_K×BLOCK_N 3. 每个线程的计算量 = (BLOCK_M × BLOCK_N) / num_threads 4. K 方向迭代次数 = K / BLOCK_K

本质上就是分治:大问题切小,每个 Block 只解决一小块,最后拼起来就是完整结果。

http://www.jsqmd.com/news/1300850/

相关文章:

  • 工业4-20mA信号转换:从经典运放到PWM方案的电路设计与实践
  • 嵌入式Linux与Qt构建汽车智能中控:架构、实战与车规级开发全解析
  • 买二手电脑哪个平台更便宜?一手货源+自动化质检兼得好品低价 - 品牌品鉴馆
  • 微信中进入定页面的,判断时通过扫二维码进入的,还是点小程序名称进入的
  • 2026保山昭通丽江感应门机组品牌怎么选?松下多玛西门子实测对比 - LYL仔仔
  • 2026 天窗冰甲防护膜:全景天幕夏天头顶暴晒的终极解决方案 - 章鱼智讯
  • 基于51单片机与AD1674的数字温度计设计:从热敏电阻到数码管显示全流程解析
  • Go语言高级并发模式实战与性能优化
  • Java素数查找算法优化与工程实践
  • 2026深圳办公室写字楼装修全流程跟了一次工地 - LYL仔仔
  • 新品发布|联软UniEDR服务器防护系统:无驱动、轻量化、AI降噪,专为服务器而生
  • Obsidian插件汉化终极指南:3分钟让英文插件秒变中文界面
  • Android PDF渲染解决方案:基于Pdfium的高性能原生渲染架构
  • AI学习进度如何不跑偏?3步动态校准法,让自学效率提升300%(附实时追踪模板)
  • Trie 树的结构优化与字符串检索加速方法7
  • 2026靠谱的物联网APP开发公司推荐 全维度选商指南 - 榜单测评
  • 2026 石家庄螺杆机组 速冻机组选购指南,工厂采购避坑干货 - LYL仔仔
  • Teable无代码数据库终极指南:5分钟从零到精通的开源数据管理平台
  • 2026年南京市鼓楼区水电维修选维小达 电路维修、水管漏水抢修、管道疏通、马桶维修、暖气维修一站式服务 - 一点传媒
  • Linux下TLS/SSL协议与密码套件探测:从OpenSSL到testssl.sh的实战指南
  • Godot虚拟摇杆实现:从原理到实战的移动端输入解决方案
  • 2026年7月新发布昆明餐饮服务机构:五家风格各异的专业机构深度解析 - 工业推荐榜
  • 论文降重天花板✅第五代改写模型真的能双检通关
  • 硬件电路设计:从需求翻译到模块化构建的系统性思维与实践
  • 2026权威测评:四川酒坛、酒缸、酒瓶5大高评分产品全维度对比 - 深度智识库
  • 拉孚 AI 审计系统打通预算/招投标/合同内审全链路合规
  • COMET翻译质量评估框架:如何用AI技术准确评估机器翻译效果
  • 南昌医疗损害强制执行律所推荐:跟进判决履行与财产查控 - 品牌深度评测
  • 开源视频修复神器untrunc终极指南:如何快速恢复损坏的MP4/MOV文件
  • 四川省居民电费计算全攻略:阶梯电价 + 峰谷电价详解