Buckets:

hf-doc-build/doc-dev / kernels /pr_779 /en /builder /triton-autotune.html
download
raw
33.8 kB
<meta charset="utf-8" /><meta name="hf:doc:metadata" content="{&quot;title&quot;:&quot;Ship Triton autotune configurations&quot;,&quot;local&quot;:&quot;ship-triton-autotune-configurations&quot;,&quot;sections&quot;:[{&quot;title&quot;:&quot;Shipping data files with a kernel&quot;,&quot;local&quot;:&quot;shipping-data-files-with-a-kernel&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Configuration file layout&quot;,&quot;local&quot;:&quot;configuration-file-layout&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Looking up configurations at runtime&quot;,&quot;local&quot;:&quot;looking-up-configurations-at-runtime&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Generating the configurations&quot;,&quot;local&quot;:&quot;generating-the-configurations&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Impact&quot;,&quot;local&quot;:&quot;impact&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2}],&quot;depth&quot;:1}"/>
<link href="/docs/kernels/main/en/_app/immutable/entry/start.BYCNYcT7.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/Cyqu85x0.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/CkB8YxkR.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/entry/app.DJ1rrz4k.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/DSBQ3IWM.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/D3DefaBM.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/DsnmJJEf.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/BJRjUCyW.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/nodes/0.-IKFj0nD.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/nodes/2.DqPxVHrz.js" rel="modulepreload">
<link href="/docs/kernels/main/en/_app/immutable/chunks/CPQPM7-z.js" rel="modulepreload">
<!--w9xnlz--><meta name="hf:doc:metadata" content="{&quot;title&quot;:&quot;Ship Triton autotune configurations&quot;,&quot;local&quot;:&quot;ship-triton-autotune-configurations&quot;,&quot;sections&quot;:[{&quot;title&quot;:&quot;Shipping data files with a kernel&quot;,&quot;local&quot;:&quot;shipping-data-files-with-a-kernel&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Configuration file layout&quot;,&quot;local&quot;:&quot;configuration-file-layout&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Looking up configurations at runtime&quot;,&quot;local&quot;:&quot;looking-up-configurations-at-runtime&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Generating the configurations&quot;,&quot;local&quot;:&quot;generating-the-configurations&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2},{&quot;title&quot;:&quot;Impact&quot;,&quot;local&quot;:&quot;impact&quot;,&quot;sections&quot;:[],&quot;depth&quot;:2}],&quot;depth&quot;:1}"/><!---->
<link href="/docs/kernels/main/en/_app/immutable/assets/0.tn0RQdqM.css" rel="modulepreload"> <!--[--><!--[0--><!--[--><!--[0--><!--[--><p></p> <div class="items-center shrink-0 min-w-[100px] max-sm:min-w-[50px] justify-end ml-auto flex" style="float: right; margin-left: 10px; display: inline-flex; position: relative; z-index: 10;"><div class="inline-flex rounded-md max-sm:rounded-sm"><button class="inline-flex items-center gap-1 h-7 max-sm:h-7 px-2 max-sm:px-1.5 text-sm font-medium text-gray-800 border border-r-0 rounded-l-md max-sm:rounded-l-sm border-gray-200 bg-white hover:shadow-inner dark:border-gray-850 dark:bg-gray-950 dark:text-gray-200 dark:hover:bg-gray-800" aria-live="polite"><span class="inline-flex items-center justify-center rounded-md p-0.5 max-sm:p-0 hover:text-gray-800 dark:hover:text-gray-200"><svg class="sm:size-3.5 size-3" xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----></span> <span>Copy page</span></button> <button class="inline-flex items-center justify-center w-6 max-sm:w-5 h-7 max-sm:h-7 disabled:pointer-events-none text-sm text-gray-500 hover:text-gray-700 dark:hover:text-white rounded-r-md max-sm:rounded-r-sm border border-l transition border-gray-200 bg-white hover:shadow-inner dark:border-gray-850 dark:bg-gray-950 dark:text-gray-200 dark:hover:bg-gray-800" aria-haspopup="menu" aria-expanded="false" aria-label="Open copy menu"><svg class="transition-transform text-gray-400 overflow-visible sm:size-3.5 size-3 rotate-0" width="1em" height="1em" viewBox="0 0 12 7" fill="none" xmlns="http://www.w3.org/2000/svg"><path d="M1 1L6 6L11 1" stroke="currentColor"></path></svg><!----></button></div> <!--[-1--><!--]--></div><!----> <!--[0--><h1 class="relative group"><a id="ship-triton-autotune-configurations" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#ship-triton-autotune-configurations"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Ship Triton autotune configurations</span></h1><!--]--><!----> <p>Triton kernels typically have parameters, such as tile sizes, number of warps and
number of pipeline stages, whose optimal values depend on the GPU and the
problem shape. Autotuning finds good values for these parameters, but doing
it at runtime (e.g. with the <a href="https://triton-lang.org/main/python-api/generated/triton.autotune.html" rel="nofollow"><code>@triton.autotune</code></a> decorator) re-benchmarks every candidate configuration in each new process.</p> <p>However, one can run the tuner for the GPU models they like and store the
best found configurations as files for using them later. This effectively
reduces the potentially costly tuning time.</p> <p><code>kernel-builder</code> support this by packaging these configurations as JSON files
with the kernel. At runtime, the kernel looks up the configuration for the
current GPU and shape and falls back to sensible defaults when there is no
matching configuration.</p> <p>This is how, for example, the vLLM fused MoE kernel ships on the Hub: the <a href="https://huggingface.co/RedHatAI/moe/tree/main/torch-ext/moe/configs" rel="nofollow"><code>RedHatAI/moe</code></a> repository contains a <code>configs</code> directory with tuned configurations for many
GPUs. This page walks through a small, complete example of the same pattern:
the <a href="https://github.com/huggingface/kernels/tree/main/examples/kernels/gemm-triton-autotune" rel="nofollow"><code>gemm-triton-autotune</code></a> example kernel, a Triton GEMM published as <a href="https://huggingface.co/kernels-test/gemm-triton-autotune" rel="nofollow"><code>kernels-test/gemm-triton-autotune</code></a>.</p> <!--[1--><h2 class="relative group"><a id="shipping-data-files-with-a-kernel" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#shipping-data-files-with-a-kernel"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Shipping data files with a kernel</span></h2><!--]--><!----> <p>Since autotune files are plain JSON files, we can store them anywhere inside the
kernels main Python sources in <code>torch-ext/&lt;kernel_name></code>. For this example,
we will use <code>torch-ext/&lt;kernel_name>/configs/</code>. By default, only <code>py</code> and <code>pyi</code> files are picked up from the kenel’s Python source directory, so add <code>json</code> to the <a href="writing-kernels#torch-noarch"><code>pyext</code> option</a> in <code>build.toml</code>:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-toml "><!----><span class="hljs-section">[general]</span>
<span class="hljs-attr">name</span> = <span class="hljs-string">&quot;gemm-triton-autotune&quot;</span>
<span class="hljs-attr">version</span> = <span class="hljs-number">1</span>
<span class="hljs-attr">edition</span> = <span class="hljs-number">5</span>
<span class="hljs-attr">license</span> = <span class="hljs-string">&quot;Apache-2.0&quot;</span>
<span class="hljs-attr">backends</span> = [<span class="hljs-string">&quot;cuda&quot;</span>, <span class="hljs-string">&quot;rocm&quot;</span>, <span class="hljs-string">&quot;xpu&quot;</span>]
<span class="hljs-section">[general.hub]</span>
<span class="hljs-attr">repo-id</span> = <span class="hljs-string">&quot;kernels-test/gemm-triton-autotune&quot;</span>
<span class="hljs-section">[torch-noarch]</span>
<span class="hljs-attr">pyext</span> = [<span class="hljs-string">&quot;json&quot;</span>, <span class="hljs-string">&quot;py&quot;</span>]<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="configuration-file-layout" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#configuration-file-layout"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Configuration file layout</span></h2><!--]--><!----> <p>A GEMM computes <code>(M, K) @ (K, N)</code>. For a model, the weight dimensions <code>N</code> and <code>K</code> are known ahead of time, while <code>M</code> (e.g. the number of tokens)
varies at runtime. The example therefore stores one file per <code>(N, K)</code> shape
and GPU, following the same naming convention as the MoE kernel:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class=" "><!----><span class="hljs-attribute">configs</span>/N=<span class="hljs-number">4096</span>,K=<span class="hljs-number">4096</span>,device_name=NVIDIA_L4.json
<span class="hljs-attribute">configs</span>/N=<span class="hljs-number">14336</span>,K=<span class="hljs-number">4096</span>,device_name=NVIDIA_L4.json<!----></pre></div><!----> <p>Each file maps an <code>M</code> value to the best configuration that the autotuner
found for that <code>M</code>:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-json "><!----><span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;1&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_M&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">16</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_N&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">128</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_K&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">32</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;GROUP_SIZE_M&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">8</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;num_warps&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">4</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;num_stages&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">3</span>
<span class="hljs-punctuation">}</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;1024&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_M&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">128</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_N&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">128</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;BLOCK_SIZE_K&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">32</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;GROUP_SIZE_M&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">8</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;num_warps&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">4</span><span class="hljs-punctuation">,</span>
<span class="hljs-attr">&quot;num_stages&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-number">3</span>
<span class="hljs-punctuation">}</span>
<span class="hljs-punctuation">}</span><!----></pre></div><!----> <p>Since the device name is part of the file name, configurations tuned for one
GPU are never applied to another. A configuration that would exceed the
resources of a smaller GPU (e.g. shared memory) is therefore harmless to
ship.</p> <!--[1--><h2 class="relative group"><a id="looking-up-configurations-at-runtime" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#looking-up-configurations-at-runtime"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Looking up configurations at runtime</span></h2><!--]--><!----> <p>At kernel launch, the kernel checks whether a configuration file exists for
the current device and shape. If it does, the configuration with the nearest
tuned <code>M</code> is used; otherwise the kernel falls back to a conservative default
and logs a warning. The lookup is cached, so the file is read at most once
per process:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-meta">@functools.cache</span>
<span class="hljs-keyword">def</span> <span class="hljs-title function_">_load_tuned_configs</span>(<span class="hljs-params">N: <span class="hljs-built_in">int</span>, K: <span class="hljs-built_in">int</span></span>) -&gt; <span class="hljs-type">Optional</span>[<span class="hljs-type">Dict</span>[<span class="hljs-built_in">int</span>, <span class="hljs-type">Dict</span>[<span class="hljs-built_in">str</span>, <span class="hljs-built_in">int</span>]]]:
path = _CONFIGS_DIR / config_file_name(N, K)
<span class="hljs-keyword">if</span> path.exists():
logger.info(<span class="hljs-string">&quot;Using tuned GEMM configurations from %s.&quot;</span>, path)
<span class="hljs-keyword">with</span> <span class="hljs-built_in">open</span>(path) <span class="hljs-keyword">as</span> f:
<span class="hljs-keyword">return</span> {<span class="hljs-built_in">int</span>(m): config <span class="hljs-keyword">for</span> m, config <span class="hljs-keyword">in</span> json.load(f).items()}
logger.warning(
<span class="hljs-string">&quot;No tuned GEMM configuration found for this device and shape (%s). &quot;</span>
<span class="hljs-string">&quot;Falling back to heuristic defaults, performance may be suboptimal. &quot;</span>
<span class="hljs-string">&quot;Generate a configuration with `tune_gemm(N=%d, K=%d)`.&quot;</span>,
path.name,
N,
K,
)
<span class="hljs-keyword">return</span> <span class="hljs-literal">None</span>
<span class="hljs-keyword">def</span> <span class="hljs-title function_">get_config</span>(<span class="hljs-params">M: <span class="hljs-built_in">int</span>, N: <span class="hljs-built_in">int</span>, K: <span class="hljs-built_in">int</span></span>) -&gt; <span class="hljs-type">Dict</span>[<span class="hljs-built_in">str</span>, <span class="hljs-built_in">int</span>]:
tuned = _load_tuned_configs(N, K)
<span class="hljs-keyword">if</span> tuned <span class="hljs-keyword">is</span> <span class="hljs-keyword">not</span> <span class="hljs-literal">None</span>:
<span class="hljs-comment"># Tuned Ms are spaced logarithmically, so pick the nearest in log space.</span>
nearest_m = <span class="hljs-built_in">min</span>(tuned, key=<span class="hljs-keyword">lambda</span> m: <span class="hljs-built_in">abs</span>(math.log(M / m)))
<span class="hljs-keyword">return</span> tuned[nearest_m]
<span class="hljs-keyword">return</span> default_config(M, N, K)<!----></pre></div><!----> <p>The configuration is then passed to the Triton kernel as its <code>constexpr</code> and launch parameters (see <a href="https://github.com/huggingface/kernels/blob/main/examples/kernels/gemm-triton-autotune/torch-ext/gemm_triton_autotune/gemm.py" rel="nofollow"><code>gemm.py</code></a> in the example):</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">def</span> <span class="hljs-title function_">launch_gemm_kernel</span>(<span class="hljs-params">a, b, out, config</span>):
M, K = a.shape
N = b.shape[<span class="hljs-number">1</span>]
grid = (
triton.cdiv(M, config[<span class="hljs-string">&quot;BLOCK_SIZE_M&quot;</span>]) * triton.cdiv(N, config[<span class="hljs-string">&quot;BLOCK_SIZE_N&quot;</span>]),
)
_gemm_kernel[grid](a, b, out, M, N, K, ..., **config)<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="generating-the-configurations" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#generating-the-configurations"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Generating the configurations</span></h2><!--]--><!----> <p>The autotuner itself is ordinary benchmarking code: for each <code>M</code>, benchmark
every candidate configuration with <code>triton.testing.do_bench</code> and keep the
fastest. The example ships the tuner as part of the kernel (the <code>tune_gemm</code> function in <a href="https://github.com/huggingface/kernels/blob/main/examples/kernels/gemm-triton-autotune/torch-ext/gemm_triton_autotune/tuning.py" rel="nofollow"><code>tuning.py</code></a>),
so users can also generate configurations for GPUs that the kernel author
did not tune. Candidates that do not fit the device (Triton raises <code>OutOfResources</code>) are skipped.</p> <p>The example repository contains a small script, <a href="https://github.com/huggingface/kernels/blob/main/examples/kernels/gemm-triton-autotune/tune.py" rel="nofollow"><code>tune.py</code></a>,
that runs the tuner and writes the configuration files to the source tree:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-bash "><!---->$ python tune.py --n 4096 --k 4096
$ python tune.py --n 14336 --k 4096<!----></pre></div><!----> <p>Commit the generated files, rebuild, and the configurations ship with the
kernel. When tuning a kernel that is not yet published, build it locally
(see <a href="local-dev">Develop locally</a>) and point <code>LOCAL_KERNELS</code> at the
build:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-bash "><!---->$ LOCAL_KERNELS=kernels-test/gemm-triton-autotune=build python tune.py --n 4096 --k 4096<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="impact" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#impact"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Impact</span></h2><!--]--><!----> <p>Tuned configurations are cheap to ship and can make a large difference. On
an NVIDIA L4, the tuned configuration for a <code>(1024, 4096) @ (4096, 4096)</code> float16 GEMM is ~1.4× faster than the example’s heuristic default (0.50 ms
vs. 0.71 ms) — on par with cuBLAS for this shape.</p> <a class="!text-gray-400 !no-underline text-sm flex items-center not-prose mt-4" href="https://github.com/huggingface/kernels/blob/main/docs/source/builder/triton-autotune.md" target="_blank"><svg class="mr-1" xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M31,16l-7,7l-1.41-1.41L28.17,16l-5.58-5.59L24,9l7,7z"></path><path d="M1,16l7-7l1.41,1.41L3.83,16l5.58,5.59L8,23l-7-7z"></path><path d="M12.419,25.484L17.639,6.552l1.932,0.518L14.351,26.002z"></path></svg><!----> <span><span class="underline">Update</span> on GitHub</span></a><!----> <p></p><!--]--><!----><!--]--><!--]--><!--]--> <!--[-1--><!--]--><!--]-->
<script>
{
__sveltekit_5yfe0v = {
base: "/docs/kernels/main/en",
assets: "/docs/kernels/main/en"
};
const element = document.currentScript.parentElement;
Promise.all([
import("/docs/kernels/main/en/_app/immutable/entry/start.BYCNYcT7.js"),
import("/docs/kernels/main/en/_app/immutable/entry/app.DJ1rrz4k.js")
]).then(([kit, app]) => {
kit.start(app, element, {
node_ids: [0, 2],
data: [null,null],
form: null,
error: null
});
});
}
</script>

Xet Storage Details

Size:
33.8 kB
·
Xet hash:
9cfb8b7fe5aa75e7b0a2d81ad8208ef4f28e64f220c9cd067c9b411dcb9137dd

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.