[{"data":1,"prerenderedAt":44},["ShallowReactive",2],{"chapter:kernels\u002Fpatterns\u002Fsoftmax.json":3},{"project":4,"route":5,"title":6,"titleHtml":6,"navTitle":6,"part":7,"sourcePath":8,"editUrl":9,"html":10,"toc":11,"hasMermaid":37,"prev":38,"next":41},"kernels","\u002Fkernels\u002Fpatterns\u002Fsoftmax","Softmax: Your First Reduction","Real Patterns","patterns\u002Fsoftmax.md","https:\u002F\u002Fgithub.com\u002Fteenygrad\u002Fteenygrad\u002Fedit\u002Fmain\u002Fbooks\u002Fkernels\u002Fsrc\u002Fpatterns\u002Fsoftmax.md","\u003Cp>Every kernel so far has been embarrassingly parallel: lane \u003Ccode>i\u003C\u002Fcode> reads element\n\u003Ccode>i\u003C\u002Fcode>, does arithmetic, writes element \u003Ccode>i\u003C\u002Fcode>. No lane needed to know anything about\nany other.\u003C\u002Fp>\n\u003Cp>Softmax breaks that. To compute one output you need a sum over the whole row,\nwhich means the lanes have to combine their values. That operation is a\n\u003Cstrong>reduction\u003C\u002Fstrong>, and it is the first genuinely new idea in this book.\u003C\u002Fp>\n\u003Ch2 id=\"the-operation\">The operation\u003C\u002Fh2>\n\u003Cp>Softmax turns a row of numbers into a probability distribution:\u003C\u002Fp>\n\u003Cpre class=\"code-panel\" data-lang=\"text\">\u003Ccode>softmax(x)_i = exp(x_i) \u002F sum_j exp(x_j)\n\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>Every output depends on every input in its row. The denominator is the\nreduction.\u003C\u002Fp>\n\u003Ch2 id=\"why-the-obvious-version-is-wrong\">Why the obvious version is wrong\u003C\u002Fh2>\n\u003Cp>Write that formula directly and it breaks. \u003Ccode>exp(x)\u003C\u002Fcode> overflows \u003Ccode>f32\u003C\u002Fcode> at about\n\u003Ccode>x = 88\u003C\u002Fcode>, and logits above 88 are entirely ordinary. You get \u003Ccode>inf \u002F inf\u003C\u002Fcode>, which\nis \u003Ccode>NaN\u003C\u002Fcode>, and the \u003Ccode>NaN\u003C\u002Fcode> spreads through the rest of your model.\u003C\u002Fp>\n\u003Cp>The fix relies on softmax being invariant to shifts. Subtract any constant from\nevery element and the result is unchanged, because the constant cancels:\u003C\u002Fp>\n\u003Cpre class=\"code-panel\" data-lang=\"text\">\u003Ccode>exp(x_i - c) \u002F sum_j exp(x_j - c)\n\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>Choose \u003Ccode>c = max(x)\u003C\u002Fcode>. Now the largest exponent is \u003Ccode>exp(0) = 1\u003C\u002Fcode>, nothing\noverflows, and the terms that underflow to zero were negligible anyway.\u003C\u002Fp>\n\u003Cp>That is what “numerically stable softmax” means, and it costs a second\nreduction: one for the maximum, one for the sum.\u003C\u002Fp>\n\u003Ch2 id=\"the-kernel\">The kernel\u003C\u002Fh2>\n\u003Cp>Here is the library’s implementation:\u003C\u002Fp>\n\u003Cpre data-lang=\"rust\" class=\"shiki teeny-datasheet\" style=\"background-color:#16181a;color:#e6e8e3\" tabindex=\"0\">\u003Ccode>\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">#[\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">kernel\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">]\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">pub\u003C\u002Fspan>\u003Cspan style=\"color:#FF5F9E\"> fn\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\"> softmax_forward\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Triton\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Float\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#FF5F9E\"> const\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> BLOCK_SIZE\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>(\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#E6E8E3\">    x_ptr\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Pointer\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#E6E8E3\">    y_ptr\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Pointer\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#E6E8E3\">    _n_rows\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#E6E8E3\">    n_cols\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">)\u003C\u002Fspan>\u003Cspan style=\"color:#FF5F9E\"> where\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">    T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">I32Tensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> types\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Tensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> 1\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">    T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">I32Tensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Comparison\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> BoolTensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\"> =\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">BoolTensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">    T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Pointer\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>:\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> AddOffsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">i32\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> 1\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">I32Tensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Output\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\"> =\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Tensor\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Pointer\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>>>,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">{\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> pid \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">program_id\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">Axis\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">X\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> row_offset \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> pid \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">*\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> n_cols\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">;\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> col_offsets \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">arange\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\">0\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> BLOCK_SIZE\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> offsets \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> col_offsets \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">+\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> row_offset\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">;\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> x \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">load\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#E6E8E3\">        x_ptr\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">.\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">add_offsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">offsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">        None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">        None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">        &#x26;[],\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">        None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">        None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">        None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#B79AD4\">        false\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">    );\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#7F877D;font-style:italic\">    \u002F\u002F Triton's builtin: numerically-stable softmax (max subtraction, exp, sum, div).\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">    let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> y \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">softmax\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">x\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> false\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> false\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#6FBF98\">    T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">store\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">y_ptr\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">.\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">add_offsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">offsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> y\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\"> &#x26;[],\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#8A9088\">}\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>The shape of it is different from anything in Part 2. \u003Cstrong>One program handles one\nwhole row.\u003C\u002Fstrong> \u003Ccode>pid\u003C\u002Fcode> is the row index, not a slice index, and \u003Ccode>row_offset\u003C\u002Fcode> jumps\nto the start of that row.\u003C\u002Fp>\n\u003Cp>There is no mask, and no \u003Ccode>T::arange(0, BLOCK_SIZE) + block_start\u003C\u002Fcode> either —\n\u003Ccode>col_offsets\u003C\u002Fcode> covers the entire row in one go.\u003C\u002Fp>\n\u003Ch2 id=\"the-constraint-and-why-it-is-there\">The constraint, and why it is there\u003C\u002Fh2>\n\u003Cp>Look at the doc comment: \u003Ccode>BLOCK_SIZE\u003C\u002Fcode> must equal \u003Ccode>n_cols\u003C\u002Fcode>. The caller is\nrequired to round the row length up to the next power of two and pass that as\nthe block size.\u003C\u002Fp>\n\u003Cp>That is a real burden pushed onto the caller. In exchange:\u003C\u002Fp>\n\u003Cul>\n\u003Cli>\u003Cstrong>No mask is needed\u003C\u002Fstrong>, because the block exactly covers the row.\u003C\u002Fli>\n\u003Cli>\u003Cstrong>No loop is needed\u003C\u002Fstrong>, because the whole row is in registers at once.\u003C\u002Fli>\n\u003Cli>\u003Cstrong>The reduction is a single tree\u003C\u002Fstrong>, with no partial-result bookkeeping.\u003C\u002Fli>\n\u003C\u002Ful>\n\u003Cp>The cost is that a row wider than the largest workable block size cannot use\nthis kernel at all, and a row of 513 elements pays for 1024.\u003C\u002Fp>\n\u003Cp>This is a fair trade and a common one, but it is exactly the kind of constraint\nthat must be shouted rather than buried. If you write a kernel with a\nprecondition like this, say so in the doc comment, as this one does.\u003C\u002Fp>\n\u003Ch2 id=\"doing-the-reduction\">Doing the reduction\u003C\u002Fh2>\n\u003Cp>The kernel uses \u003Ccode>T::softmax\u003C\u002Fcode>, a builtin that does the whole stable sequence.\nWritten out, it is:\u003C\u002Fp>\n\u003Cpre data-lang=\"rust\" class=\"shiki teeny-datasheet\" style=\"background-color:#16181a;color:#e6e8e3\" tabindex=\"0\">\u003Ccode>\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> row_max \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">max\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">x\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Some\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\">0\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> true\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003Cspan style=\"color:#7F877D;font-style:italic\">      \u002F\u002F reduce\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> shifted \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> x \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">-\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> row_max\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">;\u003C\u002Fspan>\u003Cspan style=\"color:#7F877D;font-style:italic\">                    \u002F\u002F broadcast back\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> numerator \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">exp\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">shifted\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> denominator \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">sum\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">numerator\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Some\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\">0\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> true\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003Cspan style=\"color:#7F877D;font-style:italic\">  \u002F\u002F reduce\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> y \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> numerator \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">\u002F\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> denominator\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">;\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>Five lines, two of which are reductions. Three things about them:\u003C\u002Fp>\n\u003Cp>\u003Cstrong>The \u003Ccode>axis\u003C\u002Fcode> argument selects what to reduce.\u003C\u002Fstrong> \u003Ccode>Some(0)\u003C\u002Fcode> reduces along\ndimension 0; \u003Ccode>None\u003C\u002Fcode> reduces everything to a scalar.\u003C\u002Fp>\n\u003Cp>\u003Cstrong>\u003Ccode>keep_dims\u003C\u002Fcode> decides the shape of the result.\u003C\u002Fstrong> With \u003Ccode>true\u003C\u002Fcode>, reducing a\n\u003Ccode>[128]\u003C\u002Fcode> tensor gives \u003Ccode>[1]\u003C\u002Fcode> rather than a scalar — which is what lets \u003Ccode>x - row_max\u003C\u002Fcode> broadcast back across the row. With \u003Ccode>false\u003C\u002Fcode> you get the scalar, and the\nsubtraction will not line up. This is the single most common mistake in a first\nreduction.\u003C\u002Fp>\n\u003Cp>\u003Cstrong>You do not write the reduction.\u003C\u002Fstrong> In CUDA, \u003Ccode>T::sum\u003C\u002Fcode> would be a shared-memory\ntree: each thread writes a partial, barrier, half the threads combine pairs,\nbarrier, repeat. Here the compiler emits all of that. Chapter 2 promised this\nwould be the payoff of the block model, and this is it.\u003C\u002Fp>\n\u003Ch2 id=\"watch-the-masked-lanes\">Watch the masked lanes\u003C\u002Fh2>\n\u003Cp>The softmax kernel avoids masks entirely, which sidesteps a trap. Most reduction\nkernels cannot, and then the \u003Ccode>other\u003C\u002Fcode> argument from Chapter 7 becomes essential:\u003C\u002Fp>\n\u003Cpre data-lang=\"rust\" class=\"shiki teeny-datasheet\" style=\"background-color:#16181a;color:#e6e8e3\" tabindex=\"0\">\u003Ccode>\u003Cspan class=\"line\">\u003Cspan style=\"color:#7F877D;font-style:italic\">\u002F\u002F Summing: masked lanes must be 0, the identity for addition.\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> zeros \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">zeros\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::&#x3C;\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\">D\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">>(&#x26;[\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\">BLOCK_SIZE\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">]);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> x \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">load\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">ptr\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">.\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">add_offsets\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">offs\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Some\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">mask\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Some\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">zeros\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\"> &#x26;[],\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> None\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> false\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003Cspan style=\"color:#FF5F9E\">let\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\"> total \u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">=\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> T\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">::\u003C\u002Fspan>\u003Cspan style=\"color:#7FB6D9\">sum\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#E6E8E3\">x\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">,\u003C\u002Fspan>\u003Cspan style=\"color:#6FBF98\"> Some\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">(\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\">0\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">),\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> true\u003C\u002Fspan>\u003Cspan style=\"color:#8A9088\">);\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>If you leave \u003Ccode>other\u003C\u002Fcode> as \u003Ccode>None\u003C\u002Fcode>, the masked lanes hold undefined values, and\nthose undefined values go into the sum. The result is wrong in a way that\ndepends on whatever was in memory — so it will be right in testing and wrong in\nproduction.\u003C\u002Fp>\n\u003Cp>The identity depends on the reduction:\u003C\u002Fp>\n\u003Ctable>\n\u003Cthead>\n\u003Ctr>\n\u003Cth>Reduction\u003C\u002Fth>\n\u003Cth>Fill masked lanes with\u003C\u002Fth>\n\u003C\u002Ftr>\n\u003C\u002Fthead>\n\u003Ctbody>\n\u003Ctr>\n\u003Ctd>\u003Ccode>sum\u003C\u002Fcode>\u003C\u002Ftd>\n\u003Ctd>\u003Ccode>0\u003C\u002Fcode>\u003C\u002Ftd>\n\u003C\u002Ftr>\n\u003Ctr>\n\u003Ctd>\u003Ccode>max\u003C\u002Fcode>\u003C\u002Ftd>\n\u003Ctd>the most negative representable value\u003C\u002Ftd>\n\u003C\u002Ftr>\n\u003Ctr>\n\u003Ctd>\u003Ccode>min\u003C\u002Fcode>\u003C\u002Ftd>\n\u003Ctd>the most positive representable value\u003C\u002Ftd>\n\u003C\u002Ftr>\n\u003Ctr>\n\u003Ctd>product\u003C\u002Ftd>\n\u003Ctd>\u003Ccode>1\u003C\u002Fcode>\u003C\u002Ftd>\n\u003C\u002Ftr>\n\u003C\u002Ftbody>\n\u003C\u002Ftable>\n\u003Cp>For a maximum, \u003Ccode>T::full(&amp;[BLOCK], D::from_f64(f64::NEG_INFINITY))\u003C\u002Fcode>.\u003C\u002Fp>\n\u003Ch2 id=\"running-it\">Running it\u003C\u002Fh2>\n\u003Cp>The library kernel has tests, including one that runs on a device:\u003C\u002Fp>\n\u003Cpre data-lang=\"bash\" class=\"shiki teeny-datasheet\" style=\"background-color:#16181a;color:#e6e8e3\" tabindex=\"0\">\u003Ccode>\u003Cspan class=\"line\">\u003Cspan style=\"color:#7FB6D9\">cargo\u003C\u002Fspan>\u003Cspan style=\"color:#D8A76B\"> test\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> -p\u003C\u002Fspan>\u003Cspan style=\"color:#D8A76B\"> teeny-kernels\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> --features\u003C\u002Fspan>\u003Cspan style=\"color:#D8A76B\"> cuda\u003C\u002Fspan>\u003Cspan style=\"color:#B79AD4\"> --test\u003C\u002Fspan>\u003Cspan style=\"color:#D8A76B\"> test_softmax\u003C\u002Fspan>\u003C\u002Fspan>\n\u003Cspan class=\"line\">\u003C\u002Fspan>\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>There is also a snapshot test that needs \u003Ccode>teenyc\u003C\u002Fcode> but no GPU, which compiles the\nkernel and checks its MLIR — the Chapter 9 pattern.\u003C\u002Fp>\n\u003Ch2 id=\"the-backward-pass\">The backward pass\u003C\u002Fh2>\n\u003Cp>Softmax has an unusually neat gradient. Given the saved output \u003Ccode>y\u003C\u002Fcode> and the\nupstream gradient \u003Ccode>dy\u003C\u002Fcode>:\u003C\u002Fp>\n\u003Cpre class=\"code-panel\" data-lang=\"text\">\u003Ccode>dx_i = y_i * (dy_i - sum_j(y_j * dy_j))\n\u003C\u002Fcode>\u003C\u002Fpre>\n\u003Cp>That inner sum is another row-wide reduction, and it is a scalar broadcast back\nacross the row — the same shape of computation as the forward pass. The library\nimplements it as \u003Ccode>softmax_backward\u003C\u002Fcode> in the same file, and Chapter 22 covers how\na backward kernel gets wired to its forward.\u003C\u002Fp>\n\u003Cp>Next: the reduction’s opposite problem — a kernel where the arithmetic, not the\nmemory, is the cost.\u003C\u002Fp>\n",[12,16,19,22,25,28,31,34],{"id":13,"text":14,"level":15},"the-operation","The operation",2,{"id":17,"text":18,"level":15},"why-the-obvious-version-is-wrong","Why the obvious version is wrong",{"id":20,"text":21,"level":15},"the-kernel","The kernel",{"id":23,"text":24,"level":15},"the-constraint-and-why-it-is-there","The constraint, and why it is there",{"id":26,"text":27,"level":15},"doing-the-reduction","Doing the reduction",{"id":29,"text":30,"level":15},"watch-the-masked-lanes","Watch the masked lanes",{"id":32,"text":33,"level":15},"running-it","Running it",{"id":35,"text":36,"level":15},"the-backward-pass","The backward pass",false,{"title":39,"titleHtml":39,"route":40},"Compiling and Reading the Output","\u002Fkernels\u002Ffirst-kernel\u002Fcompiling",{"title":42,"titleHtml":42,"route":43},"Matrix Multiplication","\u002Fkernels\u002Fpatterns\u002Fmatmul",1786271829695]