|
6 | 6 | "source": [ |
7 | 7 | "# Triton 编程范式 - 课后练习\n", |
8 | 8 | "\n", |
9 | | - "本 Notebook 包含三个练习,帮助你巩固 Triton 的核心概念。\n", |
10 | | - "\n", |
11 | | - "**学习目标**:\n", |
12 | | - "- 掌握 Triton 的基本语法和向量化操作\n", |
13 | | - "- 理解 `BLOCK_SIZE` 对性能的影响\n", |
14 | | - "- 学会用向量化方式处理复杂的数据访问模式" |
| 9 | + "本 Notebook 包含俩个练习,帮助你巩固 Triton 的核心概念" |
15 | 10 | ] |
16 | 11 | }, |
17 | 12 | { |
|
212 | 207 | " print(f\"Torch: {y_torch[:5].cpu().numpy()}\")" |
213 | 208 | ] |
214 | 209 | }, |
215 | | - { |
216 | | - "cell_type": "markdown", |
217 | | - "metadata": {}, |
218 | | - "source": [ |
219 | | - "**思考题**(高级):\n", |
220 | | - "1. 为什么这种方法效率不高?(提示:重复加载)\n", |
221 | | - "2. 如何优化?(提示:加载更大的块然后切片)" |
222 | | - ] |
223 | | - }, |
224 | 210 | { |
225 | 211 | "cell_type": "markdown", |
226 | 212 | "metadata": {}, |
|
229 | 215 | "\n", |
230 | 216 | "## 总结\n", |
231 | 217 | "\n", |
232 | | - "完成这三个练习后,你应该掌握了 Triton kernel 的基本写法\n", |
| 218 | + "完成这两个练习后,你应该掌握了 Triton kernel 的基本写法\n", |
233 | 219 | "\n", |
234 | | - "**下一步**:学习 Triton 的 Shared Memory 和 Block Reduction 操作!\n", |
| 220 | + "**下一步**:学习 Triton 的内存与数据搬运\n", |
235 | 221 | "\n", |
236 | 222 | "## 课后答案\n", |
237 | 223 | "\n", |
|
277 | 263 | " mask = offsets < n_elements\n", |
278 | 264 | " \n", |
279 | 265 | " x_center = tl.load(x_ptr + offsets, mask=mask, other=0.0)\n", |
280 | | - " x_left = tl.load(x_ptr + offsets - 1, mask=offsets > 0, other=0.0)\n", |
| 266 | + " x_left = tl.load(x_ptr + offsets - 1, mask=mask & (offsets > 0), other=0.0)\n", |
281 | 267 | " x_right = tl.load(x_ptr + offsets + 1, mask=offsets < n_elements - 1, other=0.0)\n", |
282 | 268 | " \n", |
283 | 269 | " y = x_left + x_center + x_right\n", |
|
0 commit comments