cpu-kernel-authoring
Use when writing, optimizing, or benchmarking a C++ CPU kernel with AVX2 or AVX512 intrinsics for the Hugging Face kernels ecosystem. Not for CUDA kernels: use cuda.
Install
npx skills add https://github.com/OutlineDriven/outline-driven-development/tree/main/.devin/skills/cpu-kernel-authoring
claude plugin marketplace add https://llmmart.ai/marketplace.json && claude plugin install outlinedriven-outline-driven-development@llmmart
git clone https://github.com/OutlineDriven/outline-driven-development.git
The skills CLI installs just this skill, for any of its supported agents. Claude Code installs the whole outlinedriven/outline-driven-development collection as a plugin from our marketplace. Git is the plain clone.
Skill manifest
CPU kernel authoring
Contract
| Field | Bound contract |
|---|---|
| Trigger | A C++ CPU kernel for the Hugging Face kernels ecosystem must be written, optimized, or benchmarked with AVX2 or AVX512 intrinsics against a PyTorch baseline. |
| Authority | Reversible local. Writes C++ kernel sources, build.toml, and torch_binding.cpp under the kernel directory, a wheel under dist/, the installed kernel package in the active Python environment, and trial state under trials/<kernel_name>/ and output/. Rollback is version control for the sources, pip uninstall <package> for the package, and removal of dist/, trials/<kernel_name>/, and output/. No remote mutation. |
| Side effect | Kernel sources and build files change; a wheel is built and installed; trial directories and result records accumulate. |
| Done | The kernel passes the correctness check in scripts/benchmark_cpu.py, every trial up to max_trials has run or the speedup exceeded early_stop_speedup, and the best trial is finalized into output/ with its final measurement; or a failure class from the table below is reported with the recovery step taken. |
Inputs
- Kernel name (required): the trial-tree label, for example
my_rmsnorm. Used only bytrial_manager.py, which accepts it as a single directory name undertrials/, never a path. - Baseline file (required): a
baseline.pythat definesget_inputs()and eitherget_reference_output()or aModelclass (with optionalget_init_inputs()). It is the ground truth for correctness and the speed reference. - Operation name (required): the plain name
analyze_op.py --oplooks up, for examplerms_norm. - Input shapes (required): comma-separated shape strings for
analyze_op.py --shapes, for example"1024x4096,2048x8192". - Package and function path (required from step 5): the installed package name, for example
my_kernel, and its callable aspackage.function, for examplemy_kernel.rms_norm.benchmark_cpu.pyandcpu_profiler.pytake this path as their--op; it is not the operation name above. - Toolchain (required): Python 3.11+ (
validate_cpu_kernel.pyparsesbuild.tomlwith the standard-librarytomllib),kernel-builder,pip, PyYAML (imported byscripts/config.py),numactl(used by the pinned benchmark in step 8), a C++ compiler with AVX512 support, and PyTorch.perfis required only whenperf_stat_enabledis true.
The work has two phases. The correctness phase builds the tiers in order (generic ATen fallback, optional AVX2, AVX512) and each tier must pass correctness before the next starts. The performance phase iterates on the AVX512 tier through the trial tree until max_trials is exhausted or early_stop_speedup is exceeded.
Procedure
- Read
scripts/config.yamland notemax_trials,early_stop_speedup,perf_stat_enabled,vtune_enabled,build_command, andinstall_command. Use those two commands wherever this procedure builds or installs. Done when: every value is known. - Run
python scripts/analyze_op.py --op <op_name> --shapes <shapes>and read the compute and memory characteristics and the suggested SIMD strategy. Readreferences/workflow_details.mdfor the analysis and design steps. Done when: the kernel type is fixed as element-wise, reduction, GEMM, or attention. - Run
python scripts/trial_manager.py init <kernel_name> <baseline_file>. Done when:trials/<kernel_name>/exists and records the baseline. - Write the generic tier:
<kernel>_cpu/cpu_features.hppin the kernel's own namespace, the dispatcher<kernel>_cpu/<kernel>_cpu.cppwith an ATen-only fallback, the bridge<kernel>_cpu/<kernel>_cpu_torch.cpp,torch-ext/torch_binding.cppusing theregistration.hmacros, andbuild.tomlwith one[kernel.*]section per tier andinclude = ["<kernel>_cpu"]in every section. Readreferences/runtime_dispatch.yaml,references/build_system.md,references/implementation_reference.md, andreferences/correctness.yamlwhile writing. Runpython scripts/validate_cpu_kernel.py <kernel_dir>. Done when: validation reports no error. - Build and install with the configured commands, by default
kernel-builder build --releasethenpip install dist/*.whl --force-reinstall --no-deps. Done when:python -c "import <package>"succeeds. - Run
python scripts/benchmark_cpu.py <baseline_file> --kernel-package <package> --op <package>.<function>. The correctness check walks tuples, lists, and dicts element-wise, requires equal structure, dtype, and shape, and compares each tensor leaf in its own dtype: half an ulp relative for bf16 and fp16,atol=1e-6, rtol=1e-5for fp32,atol=1e-12, rtol=1e-9for fp64, exact for integer and bool. Widen with--atoland--rtolonly when the kernel's accumulation order legitimately differs from the reference, and record the reason in the trial's--strategy. Done when: correctness passes and the baseline and kernel times are recorded; on failure, go to the failure table. - Add the AVX512 tier in its own translation unit
<kernel>_cpu/<kernel>_avx512.cppwith its owncxx-flagssection (-mavx512f -mavx512bf16 -mavx512vlfor element-wise kernels; GEMM kernels add-mavx512dq -mavx512bw -mavx512vbmi -mamx-tile -mamx-bf16 -mamx-int8), and-fopenmpin every SIMD section. Add an AVX2 tier only when it gives an element-wise kernel a measurable benefit; GEMM kernels dispatch AVX512 to fallback. Repeat steps 4 to 6, then runpython scripts/trial_manager.py save <kernel_name> <kernel_dir> --strategy "<description>"and record the numbers withpython scripts/trial_manager.py result <kernel_name> <trial_id> --correctness pass --speedup <x> --baseline_us <us> --kernel_us <us>. Done when: the AVX512 tier is correct and trial t0 is recorded. This ends the correctness phase. - Pin the benchmark to one NUMA node for every later measurement:
numactl --cpunodebind=0 --membind=0 python scripts/benchmark_cpu.py ... --baseline-us <cached>, where the cached value comes frompython scripts/trial_manager.py baseline-us <kernel_name>. Done when: the pinned command is the one used from here on. - When
perf_stat_enabledis true, runpython scripts/cpu_profiler.py --kernel-package <package> --op <package>.<function>once after the first benchmarked trial. Read IPC together with the L1 and LLC miss rates: a pure AVX512 FMA loop has low IPC by design, so a memory bound is claimed only when a miss rate is also high. Done when: the profile is read and the next change is chosen fromreferences/optimization_strategies.md. - For each remaining trial up to
max_trials: change one thing in the AVX512 tier (blocking, prefetch, unrolling, threading, or a different algorithm fromreferences/simd_optimization_patterns.yaml,references/memory_patterns.yaml,references/threading_patterns.yaml,references/dtype_optimizations.yaml,references/brgemm_patterns.yaml,references/quantized_gemm_patterns.yaml, andreferences/optimization_levels.yaml), validate, build, benchmark, thensavewith--parent <best_or_current_id>andresult. A regression branches back to the best trial; a plateau after two trials changes the algorithm, data layout, or fusion instead of sweeping the same knobs. Stop early only when the speedup exceedsearly_stop_speedup. Done when:max_trialstrials are recorded or the early stop fired. - Run
python scripts/trial_manager.py finalize <kernel_name> output/, then re-run the pinnedbenchmark_cpu.pywithout--baseline-usfor the final measurement. Readreferences/huggingface-kernels-integration.mdif the kernel is to be published to the Hub. Done when:output/holds the best trial's sources and its final correctness and speedup are recorded.
Modify only .cpp and .hpp files, torch_binding.cpp, and build.toml. Do not write new benchmark or timing scripts; scripts/benchmark_cpu.py is the only timing source. When a script fails, report the error rather than working around it.
Failure and recovery
| Failure class | Behavior |
|---|---|
scripts/config.yaml missing |
Stop and report. Do not assume trial counts. |
analyze_op.py reports no matmul, reduction, or activation for the op |
The script recognizes norm, softmax, gemm, linear, matmul, attention, gelu, silu, relu, moe, and megablocks by name and classifies anything else as plain element-wise. Classify the kernel type by hand from the baseline and continue at step 3. |
validate_cpu_kernel.py reports an error |
Fix the named file or build.toml section; re-run validation. A validation fix does not count as a trial. |
kernel-builder build fails |
Read the compiler output; fix the source or the section's cxx-flags; rebuild. |
| Correctness fails | Read the leaf path, dtype, and worst-element values in the benchmark output. A wrong second output or a dtype or shape change is a binding or dispatcher bug; a one-ulp bf16 difference on many elements is a rounding-mode or conversion bug; a tail-only difference is missing tail handling; a large scattered difference is an alignment bug. Fix on the same branch, rebuild, re-benchmark. Do not enter the performance phase with a failing kernel. |
| Kernel slower than baseline on small tensors | Add a num_tokens threshold below which the dispatcher calls the ATen fallback; see references/threading_patterns.yaml. |
perf unavailable or perf stat returns no counters |
Continue without profiling and report it; choose the next trial from references/optimization_levels.yaml. |
| Speedup regressed | save the next trial with --parent set to the best trial id from python scripts/trial_manager.py best <kernel_name>. |
| Plateau after two or more trials | Change algorithm, data layout, or fusion strategy. Do not sweep the same parameters. |
max_trials reached below early_stop_speedup |
Finalize the best trial and report the speedup reached and the trial tree from trial_manager.py status. |
Output
- Kernel sources:
<kernel>_cpu/withcpu_features.hpp, the dispatcher, the bridge, the AVX512 implementation, and any AVX2 implementation, plustorch-ext/torch_binding.cppandbuild.toml. - Installed package: the wheel under
dist/and the installed<package>. - Trial tree:
trials/<kernel_name>/with each saved trial, its parent, strategy, correctness, and timing. - Correctness report: the
benchmark_cpu.pyoutput naming per-dtype tolerances and, on failure, the leaf path of each mismatch. - Performance report: baseline and kernel microseconds and speedup from the NUMA-pinned run.
- Final kernel:
output/holding the best trial's sources and its final measurement.
Files (outline-driven-development)
-
agents
-
openai.yaml 250 B
interface: display_name: "Cpu Kernel Authoring" short_description: "Use when writing, optimizing, or benchmarking a C++ CPU kernel with AVX2 or AVX512 intrinsics for the Hugging Face kernels ecosystem." policy: allow_implicit_invocation: false
-
-
references
-
brgemm_patterns.yaml 14.1 KB
# brgemm and AMX Patterns # # AMX is NOT used directly in HF kernels. Instead, kernels call # at::native::cpublas::brgemm() which wraps oneDNN brgemm, and # oneDNN internally dispatches to AMX tile instructions when available. # # Source: kernels-community flash-attn2, megablocks, quantization-bitsandbytes overview: name: "brgemm: AMX-accelerated GEMM via PyTorch/oneDNN" description: | AMX (Advanced Matrix Extensions) provides hardware tile multiply on Intel Xeon 4+, but it's too complex to call directly (tile configuration, register allocation, data layout constraints). Instead, CPU kernels use the brgemm API: Kernel code → at::native::cpublas::brgemm() → oneDNN brgemm → AMX tiles The kernel developer's responsibilities: 1. Pack weight data in VNNI format (2-element interleave for bf16) 2. Tile the GEMM loop to match AMX-friendly sizes (TILE_M=16, TILE_N=16, TILE_K=32) 3. Call brgemm() for each tile 4. Clean up via brgemm_release() For small M (≤ 4 for bf16), brgemm overhead is too high: fall back to tinygemm_kernel using AVX512 _mm512_dpbf16_ps intrinsics. compiler_flags: | # While brgemm/oneDNN dispatches to AMX internally, it is highly recommended # to include AMX flags for GEMM kernels (as done in megablocks and flash-attn2) # to ensure any explicit AMX packing instructions compile successfully: cxx-flags = ["-mavx512f", "-mavx512bf16", "-mavx512vl", "-mavx512dq", "-mavx512bw", "-mavx512vbmi", "-mamx-tile", "-mamx-bf16", "-mamx-int8", "-mfma", "-mf16c", "-fopenmp"] brgemm_api: name: "at::native::cpublas::brgemm API" signature: | void at::native::cpublas::brgemm( int64_t M, // rows of output tile int64_t N, // cols of output tile int64_t K, // reduction dimension int64_t lda, // leading dim of A (typically K) int64_t ldb, // leading dim of B (typically N) int64_t ldc, // leading dim of C (typically BLOCK_N) bool add_C, // if true, accumulate onto existing C; if false, overwrite const void* A, // input activation (bf16) const void* B, // weight, VNNI-packed (bf16) float* C // output accumulator (fp32) ); void at::native::cpublas::brgemm_release(); // cleanup, call after all GEMMs done example: | #include <ATen/native/CPUBlas.h> // GEMM: C[M×N] = A[M×K] × B[K×N], A is bf16, B is VNNI-packed bf16, C is fp32 at::native::cpublas::brgemm( m_size, // M n_size, // N K, // K K, // lda n_size, // ldb BLOCK_N, // ldc false, // overwrite C A_ptr, // bf16* B_vnni_ptr, // bf16*, VNNI-packed C_ptr); // float* // Always clean up after all brgemm calls are done at::native::cpublas::brgemm_release(); vnni_packing: name: "VNNI Data Layout" description: | AMX/VNNI requires weight matrix B to be packed in VNNI format: - BF16: 2 consecutive K elements interleaved per N position - INT8: 4 consecutive K elements interleaved per N position Standard layout B[K][N]: VNNI layout B[K/2][N][2] (bf16): [b00 b01 b02 b03] [(b00,b40) (b01,b41) (b02,b42) (b03,b43)] [b10 b11 b12 b13] [(b10,b50) (b11,b51) (b12,b52) (b13,b53)] [b20 b21 b22 b23] ... [b30 b31 b32 b33] [b40 b41 b42 b43] ... pack_function: | // Pack B from [K][N] to VNNI [K/vnni_block][N][vnni_block] template<typename T> void pack_vnni(T* dst, const T* src, int K, int N) { constexpr int VNNI_BLK = std::is_same_v<T, at::BFloat16> ? 2 : 4; for (int k = 0; k < K; k += VNNI_BLK) { for (int n = 0; n < N; n++) { for (int v = 0; v < VNNI_BLK; v++) { dst[(k / VNNI_BLK) * N * VNNI_BLK + n * VNNI_BLK + v] = src[(k + v) * N + n]; } } } } weight_conversion: name: "Weight Conversion for brgemm (CRITICAL)" description: | brgemm ALWAYS requires matrix B in VNNI-interleaved layout. Every kernel that calls brgemm must convert B beforehand. The strategy differs by kernel: strategies: megablocks_persistent: name: "Megablocks MoE: Persistent Conversion at First Forward" description: | MoE expert weights are static, so VNNI packing is done ONCE at the first forward call. The Python wrapper (cpu_moe_cpp.py) drives this: 1. Transpose weights: gate_up_proj.data.transpose(-1, -2).contiguous() 2. Call ops.convert_weight_packed(data): C++ does VNNI packing 3. Store result back: self.experts.gate_up_proj.data = packed_data 4. Set self.packed_weight = True 5. All subsequent forwards pass is_vnni=True, skipping conversion For MXFP4 quantized models, also call ops.convert_scale_packed(scales) to reorder scales from [E, N, G] → [E, NB, G, BLOCK_N] for cache efficiency during the GEMM inner loop. convert_weight_packed_dtypes: | | dtype | VNNI block | Layout after packing | Extra | |-------------|-----------|-------------------------------|------------------------------| | bf16/fp16 | 2 | [IC/2, N, 2] |: | | int8 | 4 | [IC/4, N, 4] | + s8s8 compensation suffix | | fp8_e4m3 | 2 | [IC/2, N, 2] (same as bf16) |: | | uint8 (mxfp4/int4) | special | nibble unpack + 32-way repack | get_row_size(K) = K >> 1 | python_code: | # In CPUMegaBlocksMoeMLP.forward(): first call only: if not self.packed_weight: data_1 = self.experts.gate_up_proj.data.transpose(-1, -2).contiguous() data_2 = self.experts.down_proj.data.transpose(-1, -2).contiguous() if self.use_mxfp4: self.experts.gate_up_proj.storage.data = ops.convert_weight_packed(data_1) self.experts.down_proj.storage.data = ops.convert_weight_packed(data_2) # Also convert scales for MXFP4 self.experts.gate_up_proj_precision_config.weight_scale.storage.data = \ ops.convert_scale_packed(scale_data.transpose(-1, -2).contiguous()) else: data_1 = data_1.to(torch.bfloat16) if data_1.dtype == torch.float32 else data_1 self.experts.gate_up_proj.data = ops.convert_weight_packed(data_1) self.packed_weight = True cpp_fallback: | // In C++ fused_experts: if is_vnni=false, convert on-the-fly (slow) auto packed_w1 = is_vnni ? w1 : convert_weight_packed(w1); quantized_fused_dequant: name: "GPTQ / BnB: Block-Interleaved Weight + Fused Dequant per-forward" description: | The C++ kernel receives weights ALREADY converted to block-interleaved format by the model framework (GPTQModel / bitsandbytes). The kernel does NOT convert raw checkpoint weights: that's done externally. ## External conversion (done ONCE at first forward): - GPTQ: GPTQModel's transform_cpu() unpacks int32→uint8, reorders by g_idx, transposes to [N,K]; then convert_weight_packed_zp() repacks to [N,K/2] with BLOCK_N=32 interleaving. See quantized_gemm_patterns.yaml for details. - BnB: bitsandbytes' _convert_weight_packed_for_cpu() unpacks nibbles→[N,K], repacks to [N,K/2] with same BLOCK_N=32 interleaving. Also transposes absmax to [K/blocksize, N] bf16. ## Kernel-side dequant (per-forward, after receiving pre-converted weights): Quantized GEMM uses two threshold variables to decide how to handle matrix B: - `use_brgemm` (e.g. M > 4 for bf16): Switch from tinygemm (fused dequant in loop) to brgemm. - `use_brgemm_dequant_out` (e.g. M > 100): Control WHEN the unpack_B (dequant+VNNI) happens. brgemm dequant paths: - use_brgemm_dequant_out = true (M > 100): Unpack ALL available blocks of B upfront into a single large Btmp tensor before the M*N loop. Then execute brgemm using this pre-dequantized buffer. Better for Prefill phase (large M). - use_brgemm_dequant_out = false (4 < M ≤ 100): Unpack per K-block (BLOCK_K=128) into a small inner Btmp_inner buffer during the GEMM loop. Executes brgemm per block. Better for Decode phase (small/medium M) to save L3 cache / memory bandwidth. unpack_B flow: 1. Load packed byte from block-interleaved layout → nibble split 2. Subtract zero_point per group (GPTQ) or skip (BnB: zero in LUT) 3. LUT lookup (NF4/FP4) or linear scale (INT4) 4. Multiply by scale 5. _mm512_cvtne2ps_pbh → bf16 pair (naturally VNNI 2-element) 6. Store to Btmp[K/2][N][2] The block-interleaved layout ensures 32 consecutive N-elements are together in memory, enabling efficient AVX512 loads via _mm512_permutexvar_epi8. flash_attention_per_tile: name: "Flash Attention: pack_vnni per tile (no persistent convert)" description: | K and V matrices change every forward, so no caching is possible. Before each Q@K^T: pack_vnni(Btmp, K_tile) Before each S@V: pack_vnni2(Btmp, V_tile) Uses AVX512 16×16 transpose for efficient tile packing. rmsnorm_none: name: "RMSNorm / Element-wise: No conversion needed" description: | Element-wise kernels (RMSNorm, activations, reductions) do not use brgemm and need no weight conversion. Weight tensor used as-is. decision_guide: | ┌─────────────────────────────────┬────────────────────────────────────────────┐ │ Weight type │ Strategy │ ├─────────────────────────────────┼────────────────────────────────────────────┤ │ Static bf16/fp16 (MoE experts) │ convert_weight_packed() once, cache result │ │ Static bf16/fp16 + MXFP4 scales │ convert_weight_packed() + convert_scale_packed() │ │ Quantized INT4/NF4/FP4 │ Pre-converted by framework; kernel dequants per-forward │ │ Dynamic matrices (K, V) │ pack_vnni() per tile, no caching │ │ Element-wise (RMSNorm etc) │ No conversion needed │ └─────────────────────────────────┴────────────────────────────────────────────┘ Key rule: if B doesn't change between forwards → convert ONCE and cache. If B changes (activations, K/V, or is dequantized per-forward) → convert each time. tinygemm_vs_brgemm: name: "Algorithm Selection: tinygemm vs brgemm" description: | Two GEMM paths exist for different M sizes: 1. **tinygemm** (M ≤ 4 for bf16): Hand-written AVX512 kernel - Uses _mm512_dpbf16_ps (VNNI dot-product, NOT AMX tiles) - Lower overhead, better for small batch / single token - Kernel author writes this directly with SIMD intrinsics 2. **brgemm** (M > 4 for bf16): oneDNN-backed, AMX-accelerated - Higher overhead but much higher throughput for larger M - Kernel author just calls the API, oneDNN handles AMX internally selection_pattern: | template<typename scalar_t> bool can_use_brgemm(int M) { if constexpr (std::is_same_v<scalar_t, at::BFloat16>) { return M > 4; // bf16: brgemm needs M > 4 } return M > 8; // other types: higher threshold } // In the kernel: if (can_use_brgemm<scalar_t>(M)) { // Tile loop calling brgemm() for (int m = 0; m < M; m += TILE_M) { for (int n = 0; n < N; n += TILE_N) { at::native::cpublas::brgemm( std::min(TILE_M, M - m), std::min(TILE_N, N - n), K, K, N, BLOCK_N, false, A + m * K, B_vnni + n * K_vnni, C + m * BLOCK_N + n); } } at::native::cpublas::brgemm_release(); } else { // Small M: use hand-written tinygemm with _mm512_dpbf16_ps tinygemm_kernel(M, N, K, A, B_vnni, C); } tile_sizes: name: "AMX-Friendly Tile Dimensions" description: | When tiling for brgemm, use sizes that match AMX hardware tiles: values: TILE_M: 16 # AMX processes 16 rows at a time TILE_N: 16 # AMX output is 16 cols (fp32) TILE_K: 32 # AMX-BF16 processes 32 bf16 elements per step note: | These are the hardware tile dimensions. The outer blocking (BLOCK_M, BLOCK_N) can be multiples of these for better cache utilization, e.g. BLOCK_M=64 means 4 AMX tiles in M. when_to_use_brgemm: applies: - "Flash Attention: QK^T and S×V matmuls (M = seq_len, can be large)" - "Quantized GEMM: after dequantizing to bf16, M > 4" - "MoE expert GEMM: avg tokens per expert > 4" does_not_apply: - "Element-wise ops (RMSNorm, activation): use AVX512 intrinsics directly" - "Small reductions (single-vector operations)" - "Single-token decode (M=1): use tinygemm with _mm512_dpbf16_ps" runtime_detection: | // Check AMX availability for brgemm // cpu_features.hpp already provides: static bool hasAMX() { unsigned int eax, ebx, ecx, edx; __cpuid_count(7, 0, eax, ebx, ecx, edx); bool amx_bf16 = (edx & (1 << 22)) != 0; bool amx_tile = (edx & (1 << 24)) != 0; return amx_bf16 && amx_tile; } // Note: brgemm() will still work without AMX: oneDNN falls back // to AVX512 internally. But performance won't be as good. // XCR0 check for XTILEDATA/XTILECFG (bits 17,18) is not implemented // in existing kernels: they rely on brgemm to handle this internally. -
build_system.md 7.7 KB
# build.toml multi-target CPU compilation ## Overview Each CPU kernel uses `build.toml` to define multiple compilation targets with different compiler flags. The `kernel-builder` CLI reads this file to produce a Python wheel with all SIMD tiers compiled separately. ## Rules 1. Every section needs `include`. kernel-builder does not add the kernel directory to the header search path, so without it headers fail to resolve across source files. 2. Each ISA tier (generic, AVX2, AVX512) gets its own `[kernel.*]` section and so its own translation unit and flag set. 3. The base section has no `cxx-flags`. It compiles with default flags and no SIMD intrinsics. 4. The AVX2 tier is optional. Only rmsnorm has one; GEMM kernels go from generic to AVX512. ## Example: element-wise kernel with AVX2 Based on `kernels-community/rmsnorm/build.toml`: ```toml [general] name = "rmsnorm" license = "Apache-2.0" version = 1 backends = ["cpu"] [general.hub] repo-id = "kernels-community/rmsnorm" [torch] src = ["torch-ext/torch_binding.cpp"] [kernel.rmsnorm_cpu] backend = "cpu" depends = ["torch"] include = ["rmsnorm_cpu"] src = [ "rmsnorm_cpu/rmsnorm_cpu_torch.cpp", "rmsnorm_cpu/rmsnorm_cpu.cpp", "rmsnorm_cpu/rmsnorm_cpu.hpp", "rmsnorm_cpu/cpu_features.hpp", ] [kernel.rmsnorm_cpu_avx2] backend = "cpu" cxx-flags = ["-mavx2", "-mfma", "-fopenmp", "-mf16c"] depends = ["torch"] include = ["rmsnorm_cpu"] src = [ "rmsnorm_cpu/rmsnorm_avx2.cpp", "rmsnorm_cpu/rmsnorm_avx2.hpp", "rmsnorm_cpu/cpu_types_avx2.hpp", ] [kernel.rmsnorm_cpu_avx512] backend = "cpu" cxx-flags = ["-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl"] depends = ["torch"] include = ["rmsnorm_cpu"] src = [ "rmsnorm_cpu/rmsnorm_avx512.cpp", "rmsnorm_cpu/rmsnorm_avx512.hpp", "rmsnorm_cpu/cpu_types_avx512.hpp", ] ``` Note: rmsnorm AVX512 only needs `-mavx512f -mavx512bf16 -mavx512vl`. GEMM kernels additionally need `-mavx512dq -mavx512bw -mavx512vbmi` for nibble manipulation. ## Example: GEMM kernel without AVX2 Based on `kernels-community/quantization-gptq/build.toml`: ```toml [general] name = "quantization-gptq" license = "MIT" version = 1 backends = ["cpu"] [general.hub] repo-id = "kernels-community/quantization-gptq" [torch] src = ["torch-ext/torch_binding.cpp"] [kernel.gptq_cpu] backend = "cpu" depends = ["torch"] include = ["gptq_cpu"] src = [ "gptq_cpu/gptq_cpu_torch.cpp", "gptq_cpu/gptq_cpu.cpp", "gptq_cpu/gptq_cpu.hpp", "gptq_cpu/cpu_features.hpp", ] [kernel.gptq_cpu_avx512] backend = "cpu" cxx-flags = ["-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl", "-mavx512dq", "-mavx512bw", "-mavx512vbmi", "-mamx-tile", "-mamx-bf16", "-mamx-int8"] depends = ["torch"] include = ["gptq_cpu"] src = [ "gptq_cpu/gptq_avx512.cpp", "gptq_cpu/gptq_avx512.hpp", ] ``` Note: for GEMM kernels, always include the `-mamx-tile`, `-mamx-bf16`, and `-mamx-int8` flags. Even if `brgemm` dispatches to AMX internally via oneDNN, these flags are required if the kernel (like flash-attn2 or megablocks) uses custom AMX definitions or `cpuid` based checks that compile conditionally. ## Section fields | Field | Required | Description | |-------|----------|-------------| | `backend` | Yes | Always `"cpu"` for CPU kernels | | `depends` | Yes | Always `["torch"]` | | `include` | Yes | Header search dirs, typically `["<kernel>_cpu"]` | | `src` | Yes | List of source files (`.cpp` and `.hpp`) | | `cxx-flags` | No | Compiler flags. Omit for generic (no-SIMD) sections | ## Compiler flag groups ### AVX2 (element-wise only) ```toml cxx-flags = ["-mavx2", "-mfma", "-mf16c", "-fopenmp"] ``` ### AVX512 (element-wise kernels like rmsnorm) ```toml cxx-flags = ["-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl"] ``` ### AVX512 (GEMM kernels; vbmi, dq, and bw are needed for nibble manipulation) ```toml cxx-flags = ["-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl", "-mavx512dq", "-mavx512bw", "-mavx512vbmi", "-mamx-tile", "-mamx-bf16", "-mamx-int8"] ``` Note: for GEMM kernels, always include the `-mamx-tile`, `-mamx-bf16`, and `-mamx-int8` flags. Even if `brgemm` dispatches to AMX internally via oneDNN, these flags are required if the kernel (like flash-attn2 or megablocks) uses custom AMX definitions or `cpuid` based checks that compile conditionally. ### `at::vec::Vectorized` needs a `CPU_CAPABILITY` macro; `-mavx512f` is not enough This only applies if your kernel uses ATen's portable vector wrapper `at::vec::Vectorized<T>` (e.g. `convert_from_float`, `Vectorized<float>::exp()`) instead of raw `_mm512_*` intrinsics. Raw intrinsics are fine with the flags above. `<ATen/cpu/vec/vec.h>` picks the vec512 or scalar implementation from a PyTorch preprocessor macro, not from the `-m` arch flags: ```cpp #if defined(CPU_CAPABILITY_AVX512) #include <ATen/cpu/vec/vec512/vec512.h> // real AVX512 + Sleef #else #include <ATen/cpu/vec/vec256/vec256.h> // and without CPU_CAPABILITY_AVX2 falls to scalar vec_base.h #endif ``` - `-mavx512f` only defines the compiler macro `__AVX512F__`. It does not define `CPU_CAPABILITY_AVX512`, so `at::vec` silently falls back to the scalar `vec_base.h`. `-march=native` has the same problem. - Consequence is worst for transcendentals: scalar `Vectorized<float>::exp()` calls `std::expf` element-by-element, while the vec512 path calls `Sleef_expf16_u10` (16-wide). In one fused silu kernel this alone cost about 2x against PyTorch until fixed. Fix: do one of these for any TU that uses `at::vec`: ```cpp // (a) Source-level, before including vec.h (what gptq_avx512.cpp does): #define CPU_CAPABILITY_AVX512 #include <ATen/cpu/vec/vec.h> ``` ```toml # (b) Or in build.toml cxx-flags for that section: cxx-flags = [..., "-DCPU_CAPABILITY_AVX512"] ``` `-DCPU_CAPABILITY=AVX512` (the inline-namespace name) is what upstream PyTorch also passes; it is optional for correctness and harmless to add alongside. ### Verify that the build vectorized After building, confirm that the `.so` emitted AVX512 and, if it uses `at::vec` transcendentals, linked Sleef. A scalar fallback is otherwise invisible: ```bash objdump -d *.so | grep -c '%zmm' # AVX512 active: expect > 0 nm -C *.so | grep -i sleef # at::vec exp and friends: expect U Sleef_expf16_u10 ``` If `%zmm` count is 0 or you see `U expf@GLIBC` (scalar libm) where you expected vectorized math, the `CPU_CAPABILITY_AVX512` macro above is missing. ## File naming conventions | File | Purpose | |------|---------| | `<kernel>_cpu/<kernel>_cpu.cpp` | Dispatcher: checks cpu_features and calls the best tier | | `<kernel>_cpu/<kernel>_cpu.hpp` | Shared declarations | | `<kernel>_cpu/<kernel>_cpu_torch.cpp` | Python to C++ bridge (torch tensor wrapping) | | `<kernel>_cpu/cpu_features.hpp` | CPUID detection (kernel's own namespace) | | `<kernel>_cpu/<kernel>_avx2.cpp` | AVX2 implementation (optional) | | `<kernel>_cpu/<kernel>_avx512.cpp` | AVX512 implementation | | `<kernel>_cpu/cpu_types_avx512.hpp` | Vector type abstractions (optional) | | `torch-ext/torch_binding.cpp` | Op registration with registration.h macros | ## torch_binding.cpp pattern Located at `torch-ext/torch_binding.cpp`, referenced by the `[torch]` section in build.toml: ```cpp #include "registration.h" // Forward declarations, guarded for multi-device kernels #if defined(CPU_KERNEL) torch::Tensor my_kernel_cpu_forward(torch::Tensor input, torch::Tensor weight, float eps); #endif TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { ops.def("forward(Tensor input, Tensor weight, float eps) -> Tensor"); ops.impl("forward", torch::kCPU, &my_kernel_cpu_forward); } REGISTER_EXTENSION(TORCH_EXTENSION_NAME) ``` Note: rmsnorm uses `c10::DispatchKey::CompositeExplicitAutograd` instead of `torch::kCPU` since it handles device routing internally. GPTQ/BnB use `torch::kCPU`. -
correctness.yaml 5.9 KB
# Correctness Constraints for CPU Kernels # # CRITICAL: Read this file before writing any CPU kernel. # Violating these constraints causes SEGFAULT, wrong results, or silent corruption. alignment: - id: unaligned_simd_load severity: critical description: | Use UNALIGNED load/store intrinsics unless you guarantee alignment. Aligned intrinsics (_mm512_load_ps) will SEGFAULT on unaligned data. rule: "Default to _mm512_loadu_ps / _mm512_storeu_si512 (note the 'u')" exception: "Use aligned versions only if you allocate with posix_memalign or torch's aligned allocator AND the stride guarantees alignment" - id: stack_alignment severity: high description: | AVX512 requires 64-byte stack alignment for local __m512 variables. Most compilers handle this, but be careful with alloca or VLAs. rule: "Use __attribute__((aligned(64))) for stack arrays holding SIMD data" tail_handling: - id: vector_width_remainder severity: critical description: | When hidden_size is not divisible by the SIMD vector width (16 for AVX512 fp32, 32 for AVX512 bf16), the tail elements must be handled separately. patterns: - name: "Scalar tail loop" code: | int j = 0; for (; j + VEC_ELEM_NUM <= hidden_size; j += VEC_ELEM_NUM) { // SIMD loop } for (; j < hidden_size; ++j) { // Scalar fallback for remaining elements } - name: "Masked SIMD" code: | __mmask16 tail_mask = (1 << (hidden_size % 16)) - 1; _mm512_mask_storeu_ps(out + j, tail_mask, result); - name: "TORCH_CHECK assertion" code: | // If the kernel ONLY supports aligned sizes, assert early: TORCH_CHECK(hidden_size % VEC_ELEM_NUM == 0, "hidden_size must be divisible by ", VEC_ELEM_NUM); floating_point: - id: ftz_daz severity: medium description: | PyTorch may set FTZ (Flush-to-Zero) and DAZ (Denormals-Are-Zero) mode. Do NOT assume IEEE 754 denormal behavior in your kernel. This affects near-zero values in RMSNorm variance computation. rule: "Accumulate in fp32 to minimize precision loss. Add epsilon BEFORE rsqrt." - id: bf16_precision severity: medium description: | bf16 has only 7 bits of mantissa (vs 23 for fp32). Operations like pow(2), mean, rsqrt accumulate error quickly in bf16. rule: | Always accumulate in fp32, convert to bf16 only for final store. benchmark_cpu.py compares each output leaf in its own dtype: half an ulp relative for bf16. Pass --rtol to widen it only when the accumulation order legitimately differs from the reference; the default catches a truncating bf16 convert that a float32 upcast with atol=1e-2 would accept. - id: fp32_accumulation severity: high description: | _mm512_dpbf16_ps takes bf16 inputs but accumulates in fp32. This is correct. Do NOT try to accumulate in bf16. rule: "Accumulator type must be __m512 (fp32), not __m512bh (bf16)" openmp: - id: omp_small_tensor_overhead severity: high description: | OpenMP thread fork/join has non-trivial overhead (~10-50us). For small tensors, this overhead exceeds the computation time. rule: | Guard parallel regions with a size threshold: if (num_tokens > OMP_THRESHOLD) { #pragma omp parallel for for (...) { ... } } else { for (...) { ... } // serial } typical_threshold: "4-16 tokens for RMSNorm, depends on hidden_size" - id: omp_false_sharing severity: medium description: | When threads write to adjacent memory locations, cache line bouncing (false sharing) kills performance. Each cache line is 64 bytes. rule: "Ensure each thread writes to a separate cache line. For reductions, use thread-local accumulators." - id: omp_num_threads severity: low description: | Respect OMP_NUM_THREADS environment variable. Do NOT hardcode thread count. rule: "Let OpenMP determine thread count unless benchmarking specific configurations" cpuid: - id: os_support_check severity: critical description: | A CPU may support AVX512 but the OS may not enable the state save/restore. This happens on WSL1, some VMs, and containers without proper CPU passthrough. Using AVX512 without OS support causes SIGILL. rule: | Always check XCR0 register via XGETBV after checking CPUID. See references/runtime_dispatch.yaml for the full pattern. - id: feature_caching severity: low description: | CPUID is slow (~100 cycles). Cache results in static variables. rule: | static bool supported = checkAVX512(); return supported; // checked once, cached forever build_system: - id: separate_translation_units severity: critical description: | The dispatcher (*_cpu.cpp) must be compiled WITHOUT -mavx* flags. SIMD implementations must be in SEPARATE .cpp files with their own flags. Mixing flags in one translation unit causes the compiler to emit AVX512 instructions even in the "generic" code path. rule: "One [kernel.*] section per SIMD tier in build.toml, each with its own cxx-flags" - id: include_guards severity: medium description: | Headers shared between translation units (cpu_features.hpp, cpu_types_*.hpp) must use #pragma once or include guards. Without them, the linker sees duplicate symbols from separate SIMD compilation units. rule: "Use #pragma once at the top of every .hpp file" - id: namespace_isolation severity: medium description: | Each SIMD tier should be in its own namespace to avoid ODR violations. The dispatcher calls into each namespace explicitly. rule: | namespace my_kernel_cpu { namespace avx2 { void impl(...); } namespace avx512 { void impl(...); } } -
dtype_optimizations.yaml 4.1 KB
# Data Type Optimizations for CPU Kernels bf16_on_cpu: name: "BFloat16 on CPU" description: | BF16 is the primary inference dtype for LLMs on CPU. AVX512-BF16 extension provides hardware support. key_instructions: - name: "_mm512_dpbf16_ps" description: "Dot product of bf16 pairs, accumulate in fp32" note: "Processes 16 bf16 pairs per call (k must be even)" - name: "_mm512_cvtne2ps_pbh" description: "Convert two fp32 vectors to one bf16 vector (round-to-nearest-even)" - name: "CVT_BF16_TO_FP32 macro" description: "Left-shift bf16 by 16 bits to get fp32 (exact, no rounding)" code: | #define CVT_BF16_TO_FP32(a) \ _mm512_castsi512_ps(_mm512_slli_epi32(_mm512_cvtepu16_epi32(a), 16)) precision_guidelines: - "Always accumulate in fp32 (_mm512 / float), never in bf16" - "Convert to bf16 only for final output store" - "Add epsilon BEFORE rsqrt, not after (avoids underflow)" - "Correctness: benchmark_cpu.py compares bf16 leaves in bf16 at half-ulp rtol; widen with --rtol only for accumulation-order differences" fp16_on_cpu: name: "Float16 on CPU" description: | FP16 has better precision than BF16 (10-bit mantissa vs 7-bit) but narrower dynamic range (5-bit exponent vs 8-bit). support: avx2: "F16C extension: _mm256_cvtph_ps / _mm256_cvtps_ph" avx512: "_mm512_cvtph_ps / _mm512_cvtps_ph" avx512_fp16: "AVX512-FP16: Native FP16 vector compute (_mm512_add_ph, _mm512_fmadd_ph)" amx_fp16: "AMX-FP16 (Intel Xeon 6+): Native FP16 tile compute (_tile_dpfp16ps)" note: | Some older Intel architectures lack native FP16 compute and require converting to fp32 for calculation. Newer architectures support AVX512-FP16 for vector operations, and Intel Xeon 6 (Granite Rapids) and newer support AMX-FP16 for matrix multiplication. BF16 remains highly prevalent in the ecosystem because AMX-BF16 support was adopted earlier. int8_quantization: name: "INT8 on CPU" description: | INT8 quantization for weight-only or W8A8 inference. instructions: avx512_vnni: "_mm512_dpbusd_epi32: uint8 × int8 → int32 dot product" amx_int8: "_tile_dpbssd: int8 × int8 → int32 tile multiply" pattern: | // W8A8: quantize activation, multiply, dequantize // 1. Quantize input: fp32/bf16 → int8 (per-token scale) // 2. GEMM: int8 × int8 → int32 // 3. Dequantize: int32 → fp32/bf16 (multiply by scale_a × scale_w) fp8_on_cpu: name: "FP8 on CPU" description: | FP8 (E4M3 or E5M2) conversion on CPU via bit manipulation. No native FP8 compute: convert to bf16 for computation. conversion: | // FP8 E4M3 -> BF16: rebase the exponent (bias 7 -> 127) and keep sign and mantissa inline __m512bh CVT_FP8_TO_BF16(__m256i a) { __m512i x = _mm512_cvtepu8_epi16(a); // sign<<7 | exp(4b)<<3 | man(3b) __m512i vsign = _mm512_slli_epi16(_mm512_and_si512(x, _mm512_set1_epi16(0x80)), 8); x = _mm512_and_si512(x, _mm512_set1_epi16(0x7F)); __m512i e = _mm512_srli_epi16(x, 3); // exponent in bits 3-0 __m512i m = _mm512_slli_epi16(_mm512_and_si512(x, _mm512_set1_epi16(0x7)), 4); // 3-bit mantissa into BF16 bits 6-4 __m512i r = _mm512_or_si512(_mm512_slli_epi16(_mm512_add_epi16(e, _mm512_set1_epi16(120)), 7), m); // bias delta 127-7, exponent at bits 14-7 r = _mm512_mask_mov_epi16(r, _mm512_cmpeq_epi16_mask(e, _mm512_setzero_si512()), _mm512_setzero_si512()); // e==0: flush subnormals to zero r = _mm512_mask_mov_epi16(r, _mm512_cmpeq_epi16_mask(x, _mm512_set1_epi16(0x7F)), _mm512_set1_epi16(0x7FC0)); // e=15, m=7: E4M3 NaN -> BF16 quiet NaN return (__m512bh)_mm512_or_si512(r, vsign); } dtype_selection_guide: description: "Choose dtype based on operation type and precision needs" rules: - "Normalization (RMSNorm, LayerNorm): compute in fp32, output matches input dtype" - "GEMM: bf16 input, fp32 accumulate, bf16 output" - "Quantized GEMM: int4/int8 weights, bf16 activation, fp32 accumulate" - "Attention: bf16 QK^T, fp32 softmax, bf16 output" - "Activation functions: compute in fp32 if accuracy matters (GELU, SiLU)" -
huggingface-kernels-integration.md 4.4 KB
# Hugging Face kernels integration (CPU) How to load and publish CPU kernels with the Hugging Face kernels library. ## Overview The [Hugging Face kernels](https://huggingface.co/docs/kernels/en/index) library loads pre-compiled kernels from the Hugging Face Hub at run time. CPU kernels are compiled with `kernel-builder` and distributed as platform-specific wheels. What the library gives a CPU kernel: - No local compilation: the user downloads a built wheel. - Version pinning: `get_kernel` loads a named revision. - One API for CUDA, XPU, and CPU kernels. - The wheel is matched to the user's PyTorch build. ## Installation ```bash pip install kernels torch ``` ## Core API ### get_kernel ```python from kernels import get_kernel kernel = get_kernel("kernels-community/rmsnorm", version=1) ``` ### has_kernel ```python from kernels import has_kernel if has_kernel("kernels-community/rmsnorm", version=1): kernel = get_kernel("kernels-community/rmsnorm", version=1) ``` ### get_local_kernel (Development) ```python from pathlib import Path from kernels import get_local_kernel # Load from local build (requires Path object and package name) kernel = get_local_kernel(Path("/path/to/my-kernel"), "my_kernel") ``` ## CPU Kernel Usage ### RMSNorm Example ```python import torch from kernels import get_kernel, has_kernel repo_id = "kernels-community/rmsnorm" if has_kernel(repo_id, version=1): rmsnorm_kernel = get_kernel(repo_id, version=1) x = torch.randn(2, 1024, 2048, dtype=torch.bfloat16) # CPU tensor weight = torch.ones(2048, dtype=torch.bfloat16) # CPU dispatch happens automatically via torch_binding.cpp out = rmsnorm_kernel.apply_rms_norm_forward(x, weight, eps=1e-6) ``` ### Integration with Transformers ```python import torch from kernels import get_kernel, has_kernel repo_id = "kernels-community/rmsnorm" if has_kernel(repo_id, version=1): rmsnorm_kernel = get_kernel(repo_id, version=1) def patch_rmsnorm(model): """Patch model's RMSNorm to use CPU kernel.""" patched = 0 for name, module in model.named_modules(): if 'RMSNorm' in type(module).__name__: eps = getattr(module, 'variance_epsilon', None) or getattr(module, 'eps', 1e-6) def make_forward(mod, epsilon): def forward(hidden_states): return rmsnorm_kernel.apply_rms_norm_forward(hidden_states, mod.weight, eps=epsilon) return forward module.forward = make_forward(module, eps) patched += 1 return patched ``` ## Publishing a CPU Kernel ### 1. Build with kernel-builder ```bash cd my-kernel/ kernel-builder build --release # Produces dist/my_kernel-*.whl ``` ### 2. Test locally ```bash pip install dist/*.whl --force-reinstall python -c "from pathlib import Path; from kernels import get_local_kernel; k = get_local_kernel(Path('.'), 'my_kernel'); print(dir(k))" ``` ### 3. Create Hub repository ```bash # Create repo on huggingface.co/kernels-community/ # Upload the kernel source and build.toml ``` ### 4. Multi-backend support A single kernel repo can support CUDA, XPU, and CPU. The `build.toml` defines all backends: ```toml # CUDA sections [kernel.my_kernel_cuda] backend = "cuda" # ... # CPU sections [kernel.my_kernel_cpu] backend = "cpu" # ... ``` The kernel-builder CI builds wheels for each backend. `get_kernel()` automatically selects the right wheel for the user's hardware. ## Kernel File Layout (Hub) ``` my-kernel/ ├── build.toml # Multi-target build config ├── torch-ext/ │ └── torch_binding.cpp # Op registration (registration.h) ├── my_kernel_cpu/ # CPU implementation │ ├── cpu_features.hpp │ ├── my_kernel_cpu.cpp │ ├── my_kernel_cpu.hpp │ ├── my_kernel_cpu_torch.cpp │ ├── my_kernel_avx512.cpp │ └── my_kernel_avx512.hpp ├── csrc/ # CUDA implementation (if any) │ └── ... └── README.md ``` ## Notes - CPU kernels use the same `torch_binding.cpp` and `registration.h` pattern as CUDA/XPU kernels - The `ops.impl("forward", torch::kCPU, &func)` call ensures CPU dispatch (element-wise kernels such as rmsnorm instead register with `c10::DispatchKey::CompositeExplicitAutograd`) - Multi-device kernels use `#if defined(CPU_KERNEL)` / `#elif defined(CUDA_KERNEL)` guards - CPU wheels are built for x86_64 Linux; ARM/macOS may require source builds -
implementation_reference.md 12.2 KB
# Implementation reference Code templates for C++ CPU kernels with AVX2 and AVX512 intrinsics. ## Template selection Start from the template that matches the kernel type: | Kernel type | Examples | Template | |---|---|---| | Element-wise | RMSNorm, activations | Direct AVX512 intrinsics with vector type abstractions | | GEMM | quantized GEMM, MoE | tinygemm and brgemm dual path, `Unroll<N>` template | | Attention | Flash-Attention | Tiled attention with brgemm for the matmul blocks | ## Core file structure Every CPU kernel uses this layout: ``` my_kernel/ ├── my_kernel_cpu/ │ ├── cpu_features.hpp # CPUID detection (own namespace) │ ├── my_kernel_cpu.cpp # Dispatcher │ ├── my_kernel_cpu.hpp # Shared declarations │ ├── my_kernel_cpu_torch.cpp # Python ↔ C++ bridge │ ├── my_kernel_avx512.cpp # AVX512 implementation │ └── my_kernel_avx512.hpp # AVX512 declarations ├── torch-ext/ │ └── torch_binding.cpp # Op registration └── build.toml # Multi-target compilation ``` ## cpu_features.hpp (per kernel, own namespace) Each kernel carries its own copy of `cpu_features.hpp` in its own namespace. Two kernels loaded into one process with the same symbol names would violate the one-definition rule: ```cpp #pragma once #include <cpuid.h> namespace my_kernel_cpu { class CPUFeatures { public: static bool hasAVX2() { unsigned int eax, ebx, ecx, edx; if (!__get_cpuid_count(7, 0, &eax, &ebx, &ecx, &edx)) return false; return (ebx >> 5) & 1; // AVX2 bit } static bool hasAVX512BF16() { unsigned int eax, ebx, ecx, edx; // Check AVX512F first if (!__get_cpuid_count(7, 0, &eax, &ebx, &ecx, &edx)) return false; if (!((ebx >> 16) & 1)) return false; // AVX512F // Check OS support via XCR0 unsigned int xcr0_lo, xcr0_hi; asm volatile("xgetbv" : "=a"(xcr0_lo), "=d"(xcr0_hi) : "c"(0)); if ((xcr0_lo & 0xe6) != 0xe6) return false; // OS support for ZMM // Check AVX512_BF16 (CPUID leaf 7, sub-leaf 1) if (!__get_cpuid_count(7, 1, &eax, &ebx, &ecx, &edx)) return false; return (eax >> 5) & 1; // AVX512_BF16 bit } // For GEMM kernels that need AMX via brgemm static bool hasAMX() { unsigned int eax, ebx, ecx, edx; if (!__get_cpuid_count(7, 0, &eax, &ebx, &ecx, &edx)) return false; return (edx >> 24) & 1; // AMX-TILE bit } // Composite check for kernels requiring multiple features static bool hasAllRequiredFeatures() { return hasAVX512BF16(); // Add hasAMX() for GEMM kernels } }; } // namespace my_kernel_cpu ``` ## Dispatcher pattern (my_kernel_cpu.cpp) Most kernels have two tiers: AVX512, then the ATen fallback: ```cpp #include "cpu_features.hpp" #include "my_kernel_avx512.hpp" #include <torch/torch.h> namespace my_kernel_cpu { void my_kernel_forward( torch::Tensor& output, const torch::Tensor& input, const torch::Tensor& weight, float eps ) { if (CPUFeatures::hasAVX512BF16()) { avx512::my_kernel_impl(output, input, weight, eps); } else { // ATen fallback: runs on any CPU auto variance = input.to(torch::kFloat32).pow(2).mean(-1, true); output = input * torch::rsqrt(variance + eps) * weight; } } } // namespace my_kernel_cpu ``` ## Bridge file (my_kernel_cpu_torch.cpp) Bridges Python-facing tensor API to internal C++ implementation: ```cpp #include <torch/torch.h> #include "my_kernel_cpu.hpp" torch::Tensor my_kernel_cpu_forward( torch::Tensor input, torch::Tensor weight, float eps ) { auto output = torch::empty_like(input); my_kernel_cpu::my_kernel_forward(output, input, weight, eps); return output; } ``` ## Element-wise kernel template (AVX512) ```cpp #include <immintrin.h> #include <torch/torch.h> #include <omp.h> namespace my_kernel_cpu { namespace avx512 { constexpr int VEC_ELEM_NUM = 16; // fp32 elements per __m512 void rmsnorm_impl( torch::Tensor& output, const torch::Tensor& input, const torch::Tensor& weight, float eps ) { auto num_tokens = input.size(0); auto hidden_size = input.size(1); auto input_data = input.data_ptr<float>(); auto weight_data = weight.data_ptr<float>(); auto output_data = output.data_ptr<float>(); #pragma omp parallel for schedule(static) for (int64_t i = 0; i < num_tokens; ++i) { const float* x = input_data + i * hidden_size; float* y = output_data + i * hidden_size; // 1. Compute variance (sum of squares) __m512 sum_sq = _mm512_setzero_ps(); int64_t j = 0; for (; j + VEC_ELEM_NUM <= hidden_size; j += VEC_ELEM_NUM) { __m512 v = _mm512_loadu_ps(x + j); sum_sq = _mm512_fmadd_ps(v, v, sum_sq); } float variance = _mm512_reduce_add_ps(sum_sq); // Handle tail for (; j < hidden_size; ++j) { variance += x[j] * x[j]; } variance /= hidden_size; // 2. Compute rsqrt(variance + eps) float inv_rms = 1.0f / sqrtf(variance + eps); __m512 inv_rms_vec = _mm512_set1_ps(inv_rms); // 3. Normalize and scale by weight j = 0; for (; j + VEC_ELEM_NUM <= hidden_size; j += VEC_ELEM_NUM) { __m512 v = _mm512_loadu_ps(x + j); __m512 w = _mm512_loadu_ps(weight_data + j); __m512 result = _mm512_mul_ps(_mm512_mul_ps(v, inv_rms_vec), w); _mm512_storeu_ps(y + j, result); } for (; j < hidden_size; ++j) { y[j] = x[j] * inv_rms * weight_data[j]; } } } } // namespace avx512 } // namespace my_kernel_cpu ``` ## GEMM kernel patterns ### Unroll<N> template (used in all GEMM kernels) Compile-time loop unrolling via template recursion. Uses `std::integral_constant` for compile-time index: ```cpp #define ALWAYS_INLINE __attribute__((always_inline)) inline template <int n> struct Unroll { template <typename Func, typename... Args> ALWAYS_INLINE void operator()(const Func &f, Args... args) const { Unroll<n - 1>{}(f, args...); f(std::integral_constant<int, n - 1>{}, args...); } }; template <> struct Unroll<1> { template <typename Func, typename... Args> ALWAYS_INLINE void operator()(const Func &f, Args... args) const { f(std::integral_constant<int, 0>{}, args...); } }; // Usage: Unroll<ROWS * COLS>{}(compute_lambda, k); // The lambda receives std::integral_constant<int, i> as first arg. ``` ### tinygemm_kernel_nn (small-M GEMM micro-kernel) A struct template with static `apply()` method. For M ≤ 4 with bf16, uses fused dequant + _mm512_dpbf16_ps: ```cpp template <typename scalar_t, int BLOCK_M, int BLOCK_N> struct tinygemm_kernel_nn { // Primary template: static_assert fires for unsupported types static inline void apply(...) { static_assert(sizeof(scalar_t) == 0, "unsupported"); } }; // BFloat16 specialization template <int BLOCK_M, int BLOCK_N> struct tinygemm_kernel_nn<at::BFloat16, BLOCK_M, BLOCK_N> { static inline void apply( const at::BFloat16* __restrict__ A, const unsigned char* __restrict__ B, at::BFloat16* __restrict__ C, const uint8_t* __restrict__ Bz, // zero-points (GPTQ only) const at::BFloat16* __restrict__ Bs, // scales int64_t K, int blocksize, int64_t lda, int64_t ldb, int64_t ldc, int64_t strideBz, int64_t strideBs ) { constexpr int COLS = BLOCK_N / 16; __m512 vc[BLOCK_M * COLS] = {}; // fp32 accumulators // pre_compute: load zeros, LUT, etc. auto compute = [&](auto i, int k) { // 1. nibble split + zero subtract + LUT lookup → bf16 // 2. Unroll<BLOCK_M>{}(load_a_and_dpbf16, k); }; auto scale_and_store = [&](auto i) { // scale fmadd per group, convert fp32 → bf16, store }; int64_t K2 = K >> 1; // dpbf16_ps processes 2 bf16 pairs for (int64_t k = 0; k < K2; ++k) { Unroll<BLOCK_M * COLS>{}(compute, (int)k); // scale at group boundaries } Unroll<BLOCK_M * COLS>{}(scale_and_store); } }; ``` ### parallel_2d threading (GEMM kernels) Custom 2D thread decomposition for matrix operations. Template function (not std::function): ```cpp inline int adjust_num_threads(int m) { int nth = at::get_num_threads(); if (m == 1) return 1; return std::max(1, (nth >> 1) * 2); // round to even } inline int div_up(int a, int b) { return (a + b - 1) / b; } template <typename func_t> inline void parallel_2d( int m, int n, const func_t& f ) { int nth = adjust_num_threads(m); // Factor nth into nth_m * nth_n based on M/N ratio int nth_m = 1, nth_n = nth; while (nth_m < nth && nth_m * 2 <= m) { nth_m *= 2; nth_n = nth / nth_m; } #pragma omp parallel num_threads(nth) { int tid = omp_get_thread_num(); int tm = tid / nth_n; int tn = tid % nth_n; int m_start = tm * div_up(m, nth_m); int m_end = std::min(m, (tm + 1) * div_up(m, nth_m)); int n_start = tn * div_up(n, nth_n); int n_end = std::min(n, (tn + 1) * div_up(n, nth_n)); if (m_start < m_end && n_start < n_end) { f(m_start, m_end, n_start, n_end); } } } ``` ### tinygemm vs brgemm selection The `tinygemm_kernel` function wraps both paths with `parallel_2d`: ```cpp template <typename scalar_t> void tinygemm_kernel( const scalar_t *A, const unsigned char *B, scalar_t *C, const uint8_t* Bz, const scalar_t *Bs, scalar_t *Btmp, float *Ctmp, int64_t M, int64_t N, int64_t K, int blocksize, int64_t lda, int64_t ldb, int64_t ldc, int64_t strideBz, int64_t strideBs, bool brg, bool use_brgemm_dequant_out = false ) { // brg = (M > 4), set by the caller parallel_2d(div_up(M, BLOCK_M), div_up(N, BLOCK_N), [&](int mb_start, int mb_end, int nb_start, int nb_end) { for (int mb = mb_start; mb < mb_end; ++mb) { for (int nb = nb_start; nb < nb_end; ++nb) { if (brg) { // brgemm path: dequant B block → brgemm brgemm<scalar_t>::apply(...); } else { // tinygemm path: fused dequant+GEMM tinygemm_kernel_nn<scalar_t, BLOCK_M, NB_SIZE>::apply(...); } } } } ); if (brg) at::native::cpublas::brgemm_release(); } // Caller (in gemm_int4_inference): const bool use_brgemm = M > 4; const bool use_brgemm_dequant_out = M > 100; // pre-dequant all B tinygemm_kernel<scalar_t>(..., use_brgemm, use_brgemm_dequant_out); ``` ## torch_binding.cpp registration Located at `torch-ext/torch_binding.cpp`: ```cpp #include "registration.h" #if defined(CPU_KERNEL) torch::Tensor my_kernel_cpu_forward(torch::Tensor input, torch::Tensor weight, float eps); #endif TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { ops.def("forward(Tensor input, Tensor weight, float eps) -> Tensor"); ops.impl("forward", torch::kCPU, &my_kernel_cpu_forward); } REGISTER_EXTENSION(TORCH_EXTENSION_NAME) ``` Note: rmsnorm registers with `c10::DispatchKey::CompositeExplicitAutograd` instead of `torch::kCPU`. ## Vector type abstractions (cpu_types_avx512.hpp) Used by element-wise kernels (rmsnorm). Not used by GEMM kernels or flash-attn2. ```cpp struct FP32Vec16 { __m512 reg; FP32Vec16(float v) : reg(_mm512_set1_ps(v)) {} FP32Vec16(__m512 r) : reg(r) {} FP32Vec16 operator*(const FP32Vec16& other) const { return FP32Vec16(_mm512_mul_ps(reg, other.reg)); } FP32Vec16 operator+(const FP32Vec16& other) const { return FP32Vec16(_mm512_add_ps(reg, other.reg)); } float reduce_sum() const { return _mm512_reduce_add_ps(reg); } }; struct BF16Vec32 { __m512i reg; BF16Vec32(__m512i r) : reg(r) {} // Convert to two FP32Vec16 void convert(FP32Vec16& lo, FP32Vec16& hi) const { __m256i lo_half = _mm512_castsi512_si256(reg); __m256i hi_half = _mm512_extracti64x4_epi64(reg, 1); lo = FP32Vec16(_mm512_cvtpbh_ps((__m256bh)lo_half)); hi = FP32Vec16(_mm512_cvtpbh_ps((__m256bh)hi_half)); } }; ``` -
memory_patterns.yaml 3.5 KB
# Memory Optimization Patterns for CPU Kernels cache_blocking: name: "Cache Blocking (Tiling)" description: | Partition the computation into tiles that fit in L1/L2 cache. For GEMM: tile over M, N, K dimensions. L1 is typically 32-48KB (data: 32KB on most AMD Zen, 48KB on recent Intel), L2 is 1.25-2MB, L3 is shared. pattern: | // Tile the K dimension to fit working set in L2 for (int64_t kb = 0; kb < K; kb += BLOCK_K) { int64_t kb_size = std::min(BLOCK_K, K - kb); // Process tile [M, kb:kb+BLOCK_K] x [kb:kb+BLOCK_K, N] } tuning: | BLOCK_K sizing rule of thumb: - A tile (BLOCK_M × BLOCK_K) + B tile (BLOCK_K × BLOCK_N) should fit in L2 - For bf16: each element is 2 bytes - Example: BLOCK_M=32, BLOCK_N=32, BLOCK_K=128 → 32×128×2 + 128×32×2 = 16KB, fits easily in L1 prefetch: name: "Software Prefetch" description: | Use _mm_prefetch to bring data into cache before it's needed. Most effective for streaming access patterns (sequential reads). hints: - "_MM_HINT_T0: prefetch to L1 (closest, smallest)" - "_MM_HINT_T1: prefetch to L2" - "_MM_HINT_T2: prefetch to L3" - "_MM_HINT_NTA: non-temporal (minimal cache pollution, for data read once)" guideline: | Prefetch distance = memory_latency × bandwidth / element_size Typical: 4-16 cache lines ahead for L1, 64+ for L2 alignment: name: "Memory Alignment" description: | PyTorch's memory allocator guarantees 64-byte alignment for the BASE pointer of a tensor. However, tensor views and slicing (e.g., `tensor[1:]`, or indexing specific attention heads) add offsets that frequently break 64-byte alignment. Using aligned instructions (`_mm512_load_ps`) on unaligned addresses causes an immediate hardware SEGFAULT. Furthermore, on modern Intel CPUs (Skylake/AVX-512 and newer), unaligned instructions (`_mm512_loadu_ps`) have ZERO performance penalty when the address happens to be aligned. rules: - "ALWAYS use unaligned loads/stores (_mm512_loadu_*, _mm512_storeu_*). Kernels-community exclusively uses `loadu`; zero `load` instructions exist in the repo." - "For local stack buffers, use `alignas(64)` so the CPU hardware can take the optimal fast-path when your `loadu` instructions read from them." - "For heap buffers, use `aligned_alloc(64, size)` or `posix_memalign`." numa: name: "NUMA Awareness" description: | On multi-socket systems, memory access across NUMA nodes is 2-3x slower. PyTorch allocates on the current NUMA node by default. rules: - "Use numactl --cpunodebind=0 --membind=0 for single-socket benchmarking" - "For production, let the PyTorch allocator handle NUMA placement" - "Avoid cross-socket OpenMP thread migration: OMP_PROC_BIND=close" streaming_store: name: "Non-Temporal Stores" description: | For large output buffers that won't be read back soon, use streaming stores to bypass cache and avoid polluting L1/L2. code: | _mm512_stream_si512((__m512i*)(output + j), result); warning: | Only use for large, sequential, write-only buffers. Requires 64-byte alignment. _mm_sfence() after all stores. data_layout: name: "Data Layout for SIMD" description: | SIMD works best on contiguous, stride-1 data. Ensure tensors are contiguous before kernel entry. pattern: | // In torch_binding.cpp or kernel entry: auto input_c = input.contiguous(); auto weight_c = weight.contiguous(); // Then pass data_ptr<scalar_t>() to the C++ kernel -
optimization_levels.yaml 5.8 KB
# Optimization Levels for CPU Kernel Development # # Use this framework when optimizing a kernel in Phase 2. # Most production kernels should reach at least Level 2. levels: - id: level_1_generic_baseline name: "Level 1: Generic ATen Baseline" phase: 1 description: | Get the kernel working with pure PyTorch/ATen operations. No SIMD intrinsics. This is the portable fallback. checklist: - Implement using torch/ATen ops only - Correct for all dtypes (float32, bfloat16, float16) - Handle arbitrary tensor shapes (no alignment requirements) - Register via torch_binding.cpp with proper device dispatch typical_speedup: "1.0x (this IS the baseline)" when_done: "Move to Level 2. This tier is just for correctness." - id: level_2_avx2 name: "Level 2: AVX2 SIMD (Optional/Disabled by default)" phase: 1 description: | Add 256-bit SIMD vectorization. Most x86 CPUs since 2013 support AVX2. Note: This tier is disabled/skipped by default in most kernels as they transition directly to AVX512. checklist: - Use _mm256_* intrinsics with FMA - Handle tail elements (hidden_size % 8 != 0) - Add -mavx2 -mfma -fopenmp to build.toml cxx-flags - Runtime dispatch via cpu_features.hpp - Verify correctness against Level 1 output typical_speedup: "1.5-2x vs ATen" when_done: "Move to Level 3. AVX2 is an optional stepping stone, often skipped in favor of AVX512 directly." - id: level_3_avx512_basic name: "Level 3: AVX512 Basic" phase: 2 description: | Add 512-bit SIMD with AVX512 Foundation. This is the entry point to Phase 2: performance optimization starts here. checklist: - Use _mm512_* intrinsics - Use vector type abstractions (FP32Vec16, BF16Vec32) from cpu_types_avx512.hpp - Add bf16 support via _mm512_dpbf16_ps (if AVX512BF16 available) - OpenMP parallelization with #pragma omp parallel for - Proper compiler flags in build.toml typical_speedup: "2-3x vs ATen" when_done: | Run perf stat / cpu_profiler.py. Check if memory-bound or compute-bound. If memory-bound → Level 4. If compute-bound → check tile sizes. - id: level_4_avx512_optimized name: "Level 4: AVX512 Optimized" phase: 2 description: | Optimize the AVX512 implementation for the target microarchitecture. This is where most of the Phase 2 trial loop effort goes. checklist: - Cache blocking (tile M, N, K dimensions to fit L1/L2) - Prefetch instructions (_mm_prefetch with appropriate distance) - Loop unrolling (manual or via Unroll<N>{} template) - Minimize register pressure (reuse accumulators) - OpenMP grain size tuning (avoid overhead for small tensors) - Data layout optimization (ensure contiguous access patterns) optimization_search_space: tile_size: "Try BLOCK_M in {16, 32, 64}, BLOCK_N in {16, 32, 64}" prefetch_distance: "Try 0, 64, 128, 256 bytes ahead" unroll_depth: "Try 2, 4, 8 iterations" omp_threshold: "Skip OpenMP for num_tokens < threshold (try 4, 8, 16)" typical_speedup: "3-5x vs ATen" when_done: | If speedup > early_stop_speedup → finalize. If not, consider Level 5 (AMX) or algorithmic changes. - id: level_5_amx name: "Level 5: AMX via brgemm" phase: 2 description: | Use Intel AMX for matrix-heavy kernels (attention, GEMM) via brgemm. AMX provides hardware tile matrix multiplication units (SPR+). Existing HF kernels access AMX indirectly through brgemm, NOT via direct tile intrinsics. oneDNN handles AMX tile configuration internally. checklist: - Call at::native::cpublas::brgemm() for large-M GEMM - Pre-dequantize weights to bf16 before brgemm (for quantized kernels) - Use parallel_2d for 2D thread decomposition over tiles - Call brgemm_release() after parallel region - No AMX flags needed in build.toml: oneDNN dispatches internally note: | Direct AMX intrinsics (_tile_loadconfig, _tile_dpbf16ps, etc.) are NOT used by any existing HF kernel. They are listed here for reference only: - _tile_loadconfig: configure tile registers - _tile_dpbf16ps: bf16 matrix multiply (16×16 for bf16, 16×64 for int8) - _tile_dpbssd: int8 matrix multiply - Only 8 tile registers (tmm0-tmm7) applies_when: - "Flash Attention forward pass" - "Large GEMM (M > 4 for bf16)" - "INT8 quantized inference" does_not_apply_when: - "Element-wise ops (RMSNorm, activations)" - "Small reductions" - "Target CPUs without AMX (pre-SPR)" typical_speedup: "5-10x vs ATen for GEMM-heavy kernels" decision_tree: name: "Try Harder Decision Tree" description: | Use this when you're stuck at a given speedup level. rules: - condition: "Speedup < 1.5x after Level 3" action: "Check if OpenMP is actually parallelizing (OMP_NUM_THREADS). Check for false sharing." reference: "references/threading_patterns.yaml" - condition: "Speedup 1.5-2.5x after Level 3" action: "Memory-bound. Add prefetch, try cache blocking, check alignment." reference: "references/memory_patterns.yaml" - condition: "Speedup 2.5-3x after Level 4" action: "Check IPC via perf stat. If IPC > 2, compute-bound: try AMX. If IPC < 1, still memory-bound: try different blocking." reference: "references/brgemm_patterns.yaml (Level 5 / AMX is reached via brgemm)" - condition: "Speedup 3-5x" action: "Good for most workloads. Level 5 (AMX) for critical-path kernels only." - condition: "Speedup > 5x" action: "Excellent. Finalize unless there's clear headroom." - condition: "For quantized GEMM specifically" action: "Higher speedup is expected (5-10x+) because baseline includes dequant overhead. Don't stop at 3x." -
optimization_strategies.md 4.8 KB
# Optimization strategies (CPU) ## Optimization levels Work the levels in order. The speedup column records what existing kernels in the Hugging Face kernels tree reached at each level against the PyTorch baseline; a new kernel on different shapes or hardware will land elsewhere. | Level | Focus | Speedup seen in existing kernels | |---|---|---| | L1 baseline AVX512 | Correct vectorization, unaligned loads, OpenMP threading | 1.5x to 3x | | L2 memory | Prefetch (L1, L2), cache blocking, streaming stores | 2x to 4x | | L3 compute | FMA use, loop unrolling, brgemm for GEMM | 3x to 6x | | L4 expert | 2D thread decomposition, tinygemm micro-kernel, VNNI packing | 5x to 10x and above | ## Decision tree | Observation after a level | Next action | |---|---| | Speedup under 2x after L1 | Apply L2. Confirm the memory bound with `cpu_profiler.py` miss rates before adding prefetch. | | Speedup 2x to 3x after L2 | Check L3: inspect the disassembly for FMA instructions. | | Speedup 3x to 5x | Enough for most workloads. Apply L4 only to a GEMM kernel on the critical path. | | Speedup above `early_stop_speedup` | Stop. | ## Element-wise kernels (RMSNorm, activations) 1. Use `_mm512_loadu_ps` and `_mm512_storeu_ps` for every memory access. 2. Handle tail elements with a scalar loop or masked operations. 3. Use `_mm512_fmadd_ps` for multiply-add. 4. Thread over rows with `#pragma omp parallel for schedule(static)`. 5. Prefetch the next row with `_mm_prefetch(ptr, _MM_HINT_T1)`. 6. Wrap intrinsics in vector types (`FP32Vec16`, `BF16Vec32`) for readability. 7. Use grain size 1024 for `at::parallel_for`. ## GEMM kernels (quantized GEMM, MoE) 1. Implement the dual path: tinygemm for M at or below 4, brgemm above. 2. Use the `Unroll<N>` template for compile-time loop unrolling. 3. Use `_mm512_dpbf16_ps` for bf16 dot-product accumulation in tinygemm. 4. Use `at::native::cpublas::brgemm()` for large-M GEMM. 5. Pack brgemm inputs in VNNI layout: interleave bf16 pairs for AMX. 6. Decompose threads in 2D with `parallel_2d(m, n, fn)`. 7. Unroll the K loop by 4 (`#pragma GCC unroll 4`). 8. Budget 1 MB of L2 (half of a 2 MB L2) for N-blocking. ## Attention kernels (Flash-Attention) 1. Tile attention with BLOCK_M=256 and BLOCK_N=768. 2. Use brgemm for the Q K^T and softmax(S) V matmul blocks. 3. Use online softmax with a running max and sum. 4. Thread over batch, heads, and M tiles. 5. Requires AVX512 and AMX (through brgemm). ## Constraints - Use `_mm512_loadu_*`; never `_mm512_load_*`. PyTorch does not guarantee 64-byte alignment of a tensor's data pointer. - Do not mix AVX2 and AVX512 intrinsics in one translation unit. - Do not call AMX instructions directly; use the brgemm wrapper. - Do not use `double` in hot paths; use float or bf16. - Handle the tail for hidden sizes that are not a multiple of the vector width. - Keep the ATen fallback tier. ## Profiling quick reference These bands are heuristics for the Xeon cores the existing kernels ran on. Read IPC together with the miss rates: a pure AVX512 FMA loop retires few wide instructions per cycle by design, so a low IPC with low miss rates is not a memory bound. | Metric | Reading | Action | |---|---|---| | IPC under 1.0 with high L1 or LLC miss rate | Memory bound | Add prefetch, reduce tile size | | L1 miss rate above 10% | Working set exceeds L1 | Reduce blocking to fit L1 (48 KB per core on the target Xeon) | | LLC miss rate above 20% | Working set exceeds L3 | Add cache blocking for L2 (1 MB budget) | | Branch miss rate above 5% | Unpredictable branches | Use SIMD masking or `__builtin_expect` | ## Reference index | Question | File | |---|---| | Starting a kernel | `references/implementation_reference.md` | | Build system | `references/build_system.md` | | Runtime dispatch | `references/runtime_dispatch.yaml` | | GEMM kernel | `references/brgemm_patterns.yaml` and `references/quantized_gemm_patterns.yaml` | | SIMD patterns | `references/simd_optimization_patterns.yaml` | | Memory issues | `references/memory_patterns.yaml` | | Threading | `references/threading_patterns.yaml` | | Wrong results | `references/correctness.yaml` | | Data types | `references/dtype_optimizations.yaml` | | More speedup | `references/optimization_levels.yaml` | ## Checklist - [ ] Operation type identified (element-wise, reduction, GEMM, attention) - [ ] `cpu_features.hpp` created in the kernel's own namespace - [ ] ATen fallback implemented in the dispatcher - [ ] AVX512 implementation compiled with its own flags - [ ] Tail handling for sizes that are not a multiple of the vector width - [ ] `torch_binding.cpp` uses the `registration.h` macros - [ ] `build.toml` has an `include` directive in every section - [ ] Validated with `python scripts/validate_cpu_kernel.py .` - [ ] Benchmarked with `python scripts/benchmark_cpu.py` - [ ] Profiled with `python scripts/cpu_profiler.py` after the first correct trial -
quantized_gemm_patterns.yaml 18.6 KB
# Quantized GEMM Patterns for CPU (INT4, NF4, FP4, FP8, MXFP4) # # All 4-bit quantized GEMM kernels in kernels-community share the same # architecture. This reference documents the common skeleton and the # parameterized "slots" that differ between quantization formats. # # Source: Extracted from kernels-community/quantization-gptq, # quantization-bitsandbytes, megablocks overview: description: | Quantized GEMM on CPU fuses dequantization with matrix multiplication to avoid materializing full-precision weights. The core pipeline is: packed INT4 bytes → nibble split → LUT lookup → scale × dequantized → _mm512_dpbf16_ps accumulate → bf16 output This is implemented as a "tinygemm" micro-kernel for small M (decode), and falls back to "brgemm" (unpack + ATen BLAS) for large M (prefill). common_skeleton: | 1. unpack_B() : dequantize packed weights to bf16 using LUT 2. tinygemm_kernel_nn<BLOCK_M, BLOCK_N> : fused dequant+GEMM micro-kernel 3. brgemm : unpack whole block then call ATen cpublas::brgemm 4. gemm_4bit() : top-level dispatcher (selects tinygemm vs brgemm) parameterized_slots: - name: "Lookup Table (LUT)" description: | Maps 4-bit integer values to their dequantized bf16 representation. This is the ONLY component that differs between quantization formats. variants: gptq_int4: description: "Linear INT4: value = (int4_val - zero_point) * scale" lut_content: | __m512i lut = _mm512_set_epi16( 0x0000, 0x4170, 0x4160, 0x4150, 0x4140, 0x4130, 0x4120, 0x4110, 0x4100, 0x40E0, 0x40C0, 0x40A0, 0x4080, 0x4040, 0x4000, 0x3F80, 0x0000, -0x4080, -0x4000, -0x3FC0, -0x3F80, -0x3F60, -0x3F40, -0x3F20, -0x3F00, -0x3EF0, -0x3EE0, -0x3ED0, -0x3EC0, -0x3EB0, -0x3EA0, -0x3E90); has_zero_point: true zero_point_type: "per-group uint8" bitsandbytes_nf4: description: "NormalFloat4: 16 non-uniform values from normal distribution quantiles" lut_values: [-1.0, -0.6962, -0.5251, -0.3949, -0.2844, -0.1848, -0.0911, 0.0, 0.0796, 0.1609, 0.2461, 0.3379, 0.4407, 0.5626, 0.7230, 1.0] has_zero_point: false note: "Zero is encoded in the LUT at index 7, no separate zero_point tensor" bitsandbytes_fp4: description: "4-bit floating point" lut_values: [0.0, 0.0052, 0.6667, 1.0, 0.3333, 0.5, 0.1667, 0.25, 0.0, -0.0052, -0.6667, -1.0, -0.3333, -0.5, -0.1667, -0.25] has_zero_point: false megablocks_fp8: description: "FP8 E4M3: 8-bit float, converted via bit manipulation (exponent rebase + masks)" conversion: | // FP8 E4M3 -> BF16: rebase the exponent (bias 7 -> 127) and keep sign and mantissa __m512i x = _mm512_cvtepu8_epi16(a); __m512i vsign = _mm512_slli_epi16(_mm512_and_si512(x, _mm512_set1_epi16(0x80)), 8); x = _mm512_and_si512(x, _mm512_set1_epi16(0x7F)); __m512i e = _mm512_srli_epi16(x, 3); __m512i m = _mm512_slli_epi16(_mm512_and_si512(x, _mm512_set1_epi16(0x7)), 4); __m512i r = _mm512_or_si512(_mm512_slli_epi16(_mm512_add_epi16(e, _mm512_set1_epi16(120)), 7), m); r = _mm512_mask_mov_epi16(r, _mm512_cmpeq_epi16_mask(e, _mm512_setzero_si512()), _mm512_setzero_si512()); // e==0: flush subnormals r = _mm512_mask_mov_epi16(r, _mm512_cmpeq_epi16_mask(x, _mm512_set1_epi16(0x7F)), _mm512_set1_epi16(0x7FC0)); // E4M3 NaN -> BF16 quiet NaN __m512i result = _mm512_or_si512(r, vsign); has_zero_point: false scale_type: "per-block float" megablocks_mxfp4: description: "Microscaling FP4 E2M1: 4-bit float with shared exponent" conversion: | // MXFP4 to BF16 via LUT + scale multiplication const __m512 values = _mm512_set_ps(MXFP4_VALUES); // 16 float values const __m512i lut = (__m512i)(_mm512_cvtne2ps_pbh(values, values)); x0 = _mm512_permutexvar_epi16(x0, lut); // Emulate bf16 mul with shared scale: add exponent in integer domain x0 = _mm512_add_epi16(x0, scale_bf16); has_zero_point: false scale_type: "per-block bf16 (shared exponent)" - name: "Zero-Point Handling" description: | GPTQ uses explicit per-group zero_points (uint8). BnB/FP8/MXFP4 encode the zero mapping in the LUT itself, eliminating the zero_point tensor. gptq_pattern: | // Load 32 zero points, repeat-interleave to match weight layout __m256i zraw = _mm256_loadu_si256((const __m256i*)(Bz + kgs * strideBz + n)); zeros_lo = _mm256_permutexvar_epi8(z_idx0, zraw); zeros_hi = _mm256_permutexvar_epi8(z_idx1, zraw); // Subtract before LUT w_lo = _mm256_sub_epi8(w_lo, zeros_lo); bitsandbytes_pattern: | // No zero_point tensor: just shift to LUT index range w_lo = _mm256_add_epi8(w_lo, fifteen); // shift [0,15] → [0,30] for LUT - name: "Scale Application" all_formats: | // Load scales (bf16) → convert to fp32 → broadcast to match tile layout __m512i scales_bf16 = _mm512_loadu_si512((const __m512i*)(Bs + kgs * strideBs + n)); scale_lo_fp32 = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32(scales_bf16, 0)); scale_hi_fp32 = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32(scales_bf16, 1)); // Permute to match 16-element groups scales[0] = _mm512_permutexvar_ps(s_idx0, scale_lo_fp32); scales[1] = _mm512_permutexvar_ps(s_idx1, scale_lo_fp32); unpack_b_template: name: "unpack_B: Weight Dequantization" description: | Dequantizes packed INT4 weights to bf16. Called either: - Per-block inside tinygemm (fused, for small M) - Once for full weight matrix before brgemm (materialized, for large M) code: | inline void unpack_B( at::BFloat16* __restrict__ Btmp, const unsigned char* __restrict__ packed_B, const at::BFloat16* __restrict__ Bs, // scales int64_t N, int64_t K, int blocksize, int64_t ldb, int64_t ldb_tmp, int64_t strideBs) { const int64_t K2 = K >> 1; // 2 weights per byte const int64_t gs2 = blocksize >> 1; __m256i mask = _mm256_set1_epi8(0xF); for (int64_t n = 0; n < N; n += 32) { for (int64_t k = 0; k < K2; ++k) { if (k % gs2 == 0) { // Load scales for this group // ... (see parameterized_slots.scale_application) } // Load 32 packed bytes → 64 INT4 values __m256i w_u4 = _mm256_loadu_si256((const __m256i*)(packed_B + k * ldb + n)); // Split nibbles __m256i w_lo = w_u4 & mask; __m256i w_hi = _mm256_srli_epi16(w_u4, 4) & mask; // === FORMAT-SPECIFIC SLOT: zero_point + LUT === // GPTQ: subtract zero, then LUT // BnB: shift to range, then LUT // LUT lookup → bf16 __m512i w_lo_bf16 = _mm512_permutexvar_epi16( _mm512_cvtepi8_epi16(w_lo), lut); // Scale and convert __m512 w_fp32 = CVT_BF16_TO_FP32(...) * scales[...]; __m512bh packed = _mm512_cvtne2ps_pbh(w_hi_fp32, w_lo_fp32); _mm512_storeu_si512(Btmp + offset, (__m512i)packed); } } } tinygemm_template: name: "tinygemm_kernel_nn: Fused Dequant+GEMM Micro-Kernel" description: | For small M (typical in decode/generation), fuse dequantization into the GEMM inner loop. Each iteration: load A tile, dequant B tile, dpbf16_ps. Template parameters: - BLOCK_M: output rows (typically 1-32 for decode) - BLOCK_N: output columns (typically 32, matching nibble pack width) - Unroll<ROWS * COLS>{}: compile-time loop unrolling helper key_instructions: - "_mm512_dpbf16_ps(acc, a_bf16, b_bf16) // dot product bf16→fp32 accumulate" - "_mm512_cvtne2ps_pbh(hi, lo) // pack two fp32 vectors to bf16" - "_mm512_permutexvar_epi16(idx, lut) // LUT lookup for dequantization" four_stage_loop: | // Stage 1: loadc: zero accumulators auto loadc = [&](auto i) { vc_master[i] = _mm512_set1_ps(0.f); }; Unroll<ROWS * COLS>{}(loadc); for (int64_t k = 0; k < K2; k += gs2) { // Stage 2: pre_compute: load scales (and zeros for GPTQ) auto pre_compute = [&](auto i, int64_t kgs) { if constexpr (row == 0 && col % 2 == 0) { // Load and broadcast scales for this group } }; Unroll<ROWS * COLS>{}(pre_compute, k / gs2); for (int64_t k_offset = 0; k_offset < gs2; ++k_offset) { // Stage 3: compute: dequant + dpbf16_ps auto compute = [&](auto i, int64_t k) { if constexpr (col == 0) { va = (__m512bh)(_mm512_set1_ps(a_ptr[row * lda2 + k])); } if constexpr (row == 0 && col % 2 == 0) { // Dequant: load INT4, split nibbles, LUT, scale } vc[i] = _mm512_dpbf16_ps(vc[i], va, vb[col]); }; Unroll<ROWS * COLS>{}(compute, k + k_offset); } // Stage 4: post_compute: apply group scale to accumulator auto post_compute = [&](auto i, int64_t kgs) { vc_master[i] = _mm512_fmadd_ps(vc[i], scales[i % COLS], vc_master[i]); }; Unroll<ROWS * COLS>{}(post_compute, k / gs2); } // Stage 5: storec: convert fp32 accumulators to bf16 and store auto storec = [&](auto i) { if constexpr (col % 2 == 0) { _mm512_storeu_si512(C + row * ldc + col * 16, (__m512i)(_mm512_cvtne2ps_pbh(vc_master[i+1], vc_master[i]))); } }; Unroll<ROWS * COLS>{}(storec); algorithm_selection: name: "tinygemm vs brgemm Selection & Dequant Policy" description: | For small M (M ≤ 4 for bf16, typical decode), use tinygemm (fused dequant+GEMM). For large M (M > 4), use brgemm via parallel_2d. Within the brgemm path, we further differentiate the dequantization policy based on M: - use_brgemm_dequant_out = true (M > 100, Prefill): Pre-dequantize all N×K weights into VNNI format upfront in parallel, then run GEMM. - use_brgemm_dequant_out = false (4 < M ≤ 100, Decode/Medium): Dequantize per K-block (chunk) within the inner loop on-the-fly. Both paths are wrapped in `tinygemm_kernel` which uses `parallel_2d` for 2D thread decomposition over (M/BLOCK_M, N/BLOCK_N) tiles. code: | void gemm_4bit(const torch::Tensor& input, const torch::Tensor& weight, const torch::Tensor& scales, torch::Tensor& out, int blocksize) { int64_t M = input.size(0); if (CPUFeatures::hasAVX512BF16()) { const bool use_brgemm = M > 4; const bool use_brgemm_dequant_out = M > 100; // pre-dequant all B tinygemm_kernel<scalar_t>( A, B, C, Bz, Bs, Btmp, Ctmp, M, N, K, blocksize, lda, ldb, ldc, strideBz, strideBs, use_brgemm, use_brgemm_dequant_out ); } else { // Generic fallback: unpack with PyTorch ops, then torch::matmul auto unpacked = unpack_with_torch(weight, zeros, scales, blocksize); torch::matmul_out(out, input, unpacked.t()); } } weight_format: name: "Weight Format and Conversion (CRITICAL)" description: | The CPU kernel expects weight in a specific block-interleaved format: qweight: [N, K/2] uint8, BLOCK_N=32 interleaved This is NOT the standard GPTQ/BnB checkpoint format. Each framework does its own conversion BEFORE calling the kernel for the first time. kernel_api: gptq: | gemm_int4_forward(input, weight, zeros, absmax, blocksize) input: [M, K] bf16 weight: [N, K/2] uint8: block-interleaved (BLOCK_N=32) zeros: [num_groups, N] uint8: per-group zero-points absmax: [num_groups, N] bf16: per-group scales (called "absmax" in API, actually scales) output: [M, N] bf16 bitsandbytes: | gemm_4bit_forward(input, weight, absmax, blocksize, quant_type) input: [M, K] bf16 weight: [N, K/2] uint8: block-interleaved (BLOCK_N=32) absmax: [K/blocksize, N] bf16: per-block scales (TRANSPOSED vs GPTQ!) blocksize: int (typically 64) quant_type: 0=NF4, 1=FP4 output: [M, N] bf16 NOTE: No zeros param: NF4/FP4 LUT encodes zero mapping block_interleaved_format: | ## Block-Interleaved Packing (BLOCK_N=32) Both GPTQ and BnB use the SAME packing algorithm to convert [N, K] uint8 (individual 4-bit values) → [N, K/2] uint8 (packed): 1. Reshape to [N/32, 32, K/2, 2] : split N into blocks of 32 2. Transpose to [N/32, K/2, 32, 2] : put K/2 before BLOCK_N 3. Flatten to [-1, 64], split high 32 / low 32 4. Pack: ((high << 4) | low).to(uint8) : two nibbles per byte 5. Reshape back to [N, K/2] This interleaving groups 32 N-elements together in memory, which maps directly to AVX512 register width (32 × 16-bit = 512 bits after dequant). The C++ kernel's unpack_B() reverses this with _mm512_permutexvar_epi8. gptq_conversion: name: "GPTQ Weight Conversion (in GPTQModel repo)" description: | Performed by HFKernelLinear.transform() at first inference forward. Two steps: transform_cpu() then convert_weight_packed_zp(). where: "GPTQModel/gptqmodel/nn_modules/qlinear/gemm_hf_kernel.py" triggered_by: "HFKernelLinear.forward() → self.transform(x.device.type)" step1_transform_cpu: | def transform_cpu(self): # 1. Convert scales to bf16 self.scales = self.scales.to(torch.bfloat16).contiguous() # 2. Unpack qweight from int32 packed → uint8 individual values # Original: [N/pack_factor, K] int32 (each int32 holds 8 int4 values) weight = bitwise_and(bitwise_right_shift(qweight.unsqueeze(1).expand(...), wf), maxq) # 3. Reorder by g_idx for desc_act support ret_idx = self._build_ret_idx() # inverse permutation of g_idx weight = weight.reshape(...).index_select(0, ret_idx) # 4. Transpose: [K, N] → [N, K] uint8 weight = weight.t() self.qweight = weight.contiguous() # [N, K] uint8 (0-15) # 5. Unpack qzeros similarly → [num_groups, N] uint8 zeros = bitwise_and(bitwise_right_shift(qzeros.unsqueeze(2).expand(...), wf), maxq) self.qzeros = zeros.reshape(num_groups, N).contiguous() step2_convert_weight_packed_zp: | def convert_weight_packed_zp(self, block_n=32): # Takes [N, K] uint8 → [N, K/2] uint8 block-interleaved # (see block_interleaved_format above for algorithm) bnb_conversion: name: "BitsAndBytes Weight Conversion (in bitsandbytes repo)" description: | Performed by _convert_weight_packed_for_cpu() called from Linear4bit.forward() when device is CPU. where: "bitsandbytes/bitsandbytes/functional.py" triggered_by: "Linear4bit.forward() → _convert_weight_packed_for_cpu(qweight, quant_state)" steps: | def _convert_weight_packed_for_cpu(qweight, quant_state, block_n=32): # 1. View as uint8 if stored as different dtype qweight = qweight.view(torch.uint8) # 2. Unpack BnB nibble format → [N, K] uint8 # BnB stores: high nibble = even position, low nibble = odd position unpacked_w = empty(qweight.shape[0] * 2) unpacked_w[1::2] = qweight & 0xF # low nibble → odd unpacked_w[::2] = qweight >> 4 # high nibble → even qweight_final = unpacked_w.reshape(N, K).to(uint8) # 3. Block-interleaved packing: [N, K] → [N, K/2] # (same algorithm as GPTQ's convert_weight_packed_zp) # 4. Denest absmax (if nested/double quantization) if quant_state.nested: absmax = dequantize_blockwise(quant_state.absmax, quant_state.state2) absmax += quant_state.offset # 5. Reshape + TRANSPOSE absmax → [K/blocksize, N] bf16 # NOTE: GPTQ keeps scales as [num_groups, N], # BnB TRANSPOSES to [K/blocksize, N] quant_state.absmax = ( quant_state.absmax .reshape(N, K // blocksize) # [N, num_groups] .T # [num_groups, N] .to(torch.bfloat16) .contiguous() ) key_differences: name: "GPTQ vs BnB Format Differences" table: | | Aspect | GPTQ | BitsAndBytes | |---------------------|----------------------------------------|----------------------------------------| | Original format | int32 packed (8×int4 per int32) | uint8 packed (2×nibble per byte) | | Unpacking | bitwise_right_shift + mask | qw >> 4 (even), qw & 0xF (odd) | | g_idx reorder | Yes (desc_act support via ret_idx) | No | | Final qweight | [N, K/2] uint8 block-interleaved | [N, K/2] uint8 block-interleaved | | Zero-points | [num_groups, N] uint8 | None (encoded in NF4/FP4 LUT) | | Scales shape | [num_groups, N] bf16 | [K/blocksize, N] bf16 (TRANSPOSED) | | Scales name in API | absmax | absmax | | Conversion location | GPTQModel repo (gemm_hf_kernel.py) | bitsandbytes repo (functional.py) | | Triggered by | first forward (self.linear_mode=None) | first forward on CPU device | | Reversible | No (one-way transform) | Yes (_convert_weight_packed_for_cpu_inverse) | adding_new_format: name: "How to Add a New Quantization Format" steps: - "1. Define the LUT values (16 entries for 4-bit, 256 for 8-bit)" - "2. Decide zero_point strategy: explicit tensor, or encoded in LUT" - "3. Copy an existing kernel (e.g., quantization-gptq) as template" - "4. Replace the LUT constant and zero_point handling" - "5. Adjust the unpack_B function if packing layout differs" - "6. Keep tinygemm_kernel_nn and brgemm logic unchanged" - "7. Update build.toml with the new kernel name" - "8. Write test with ref_gemm_Xbit() using pure PyTorch" note: | ~90% of the code is reusable. Only the LUT, zero-point, and scale loading code change between formats. The tinygemm micro-kernel, Unroll<> helper, and brgemm fallback are identical. -
runtime_dispatch.yaml 8.5 KB
# Runtime CPU Feature Detection and Multi-Tier Dispatch # # Every CPU kernel MUST implement this pattern. The dispatcher detects # hardware features at runtime via CPUID and routes to the best implementation. # # Source: Extracted from kernels-community/rmsnorm, flash-attn2, quantization-gptq dispatch_pattern: name: "Three-Tier Runtime Dispatch" description: | All CPU kernels share the same dispatch architecture: 1. cpu_features.hpp: CPUID detection (compile-time header, no SIMD flags needed) 2. *_cpu.cpp: dispatcher entry point (compiled without SIMD flags) 3. *_avx2.cpp / *_avx512.cpp: SIMD implementations (compiled WITH flags) The dispatcher is compiled WITHOUT any -mavx* flags so it runs on any CPU. Each SIMD implementation is in a SEPARATE translation unit compiled with its own flags. flow: | CPUFeatures::hasAVX512BF16() ├─ true → avx512::kernel_impl() ├─ false → CPUFeatures::hasAVX2() │ ├─ true → avx2::kernel_impl() │ └─ false → ATen fallback (torch ops) └─ (never crashes: always has a fallback) cpu_features_hpp: name: "cpu_features.hpp: CPUID Detection" description: | Shared header that detects CPU features at runtime. Uses cpuid instruction directly (no external dependency). Results are cached in static variables. CRITICAL: Must also check OS support via XCR0 register (XGETBV): the CPU may support AVX512 but the OS may not have enabled the save/restore of AVX512 state. Without this check, the kernel SEGFAULTS on WSL1 and some VMs. detection_methods: - name: hasAVX2 cpuid_leaf: 7 register: EBX bit: 5 description: "AVX2: 256-bit integer SIMD" - name: hasAVX512 cpuid_leaf: 7 register: EBX bit: 16 os_check: "XCR0 bits 1,2,5,6,7 must be set (SSE + AVX + AVX512 state)" description: "AVX-512 Foundation" - name: hasAVX512BF16 cpuid_leaf: 7 subleaf: 1 register: EAX bit: 5 depends_on: hasAVX512 description: "AVX-512 BF16: _mm512_dpbf16_ps" - name: hasAMX cpuid_leaf: 7 register: EDX bits: [22, 24] os_check: "XCR0 bits 17,18 must be set (XTILEDATA + XTILECFG)" description: "AMX-BF16 (bit 22), AMX-TILE (bit 24)" note: | Explicit AMX detection is ONLY used by specific kernels (like flash-attn2 and megablocks) that strictly gate their execution on full hardware support (`hasAllRequiredFeatures()`). For standard quantized GEMMs (GPTQ, AWQ, BnB), `hasAMX` is omitted: you just check `hasAVX512BF16()` and let `brgemm` silently fall back to AVX512 if AMX is absent. template: | #pragma once #ifdef _MSC_VER #include <intrin.h> #else #include <cpuid.h> #endif namespace my_kernel_cpu { class CPUFeatures { public: static bool hasAVX2() { static bool supported = checkAVX2(); return supported; } static bool hasAVX512BF16() { static bool supported = checkAVX512BF16(); return supported; } private: static bool checkAVX2() { unsigned int eax, ebx, ecx, edx; if (__get_cpuid_max(0, nullptr) < 7) return false; __cpuid_count(7, 0, eax, ebx, ecx, edx); return (ebx & (1 << 5)) != 0; // EBX bit 5 } static bool checkAVX512() { unsigned int eax, ebx, ecx, edx; if (__get_cpuid_max(0, nullptr) < 7) return false; __cpuid_count(7, 0, eax, ebx, ecx, edx); if (!(ebx & (1 << 16))) return false; // AVX-512 Foundation // Check OS support for AVX512 state save/restore if (__get_cpuid(1, &eax, &ebx, &ecx, &edx) == 0) return false; if (!(ecx & (1 << 27))) return false; // OSXSAVE unsigned int xcr0_lo, xcr0_hi; __asm__ volatile("xgetbv" : "=a"(xcr0_lo), "=d"(xcr0_hi) : "c"(0)); unsigned long long xcr0 = ((unsigned long long)xcr0_hi << 32) | xcr0_lo; return (xcr0 & 0xE6ULL) == 0xE6ULL; // bits 1,2,5,6,7 } static bool checkAVX512BF16() { if (!checkAVX512()) return false; unsigned int eax, ebx, ecx, edx; __cpuid_count(7, 1, eax, ebx, ecx, edx); return (eax & (1 << 5)) != 0; // EAX bit 5 } }; } // namespace dispatcher_pattern: name: "Dispatcher Entry Point (*_cpu.cpp)" description: | The main dispatch file includes headers for all SIMD tiers and routes based on runtime detection. This file is compiled WITHOUT -mavx* flags. template: | #include "cpu_features.hpp" #include "my_kernel_avx2.hpp" #include "my_kernel_avx512.hpp" #include <ATen/ATen.h> namespace my_kernel_cpu { void forward(torch::Tensor& out, const torch::Tensor& input, const torch::Tensor& weight, float epsilon) { if (CPUFeatures::hasAVX512BF16()) { avx512::forward_impl(out, input, weight, epsilon); } else if (CPUFeatures::hasAVX2()) { avx2::forward_impl(out, input, weight, epsilon); } else { // Generic ATen fallback auto x = input.to(at::kFloat); auto variance = at::mean(at::pow(x, 2), -1, true); out = at::mul(weight, at::mul(x, at::rsqrt(at::add(variance, epsilon)))) .to(input.scalar_type()); } } } // namespace torch_binding_pattern: name: "torch_binding.cpp: PyTorch C++ Extension Binding" description: | The binding file uses #if defined(CPU_KERNEL) preprocessor guards to conditionally compile CPU-specific code. Multiple backends (CPU, XPU, CUDA) can coexist in the same binding file. template: | #include <torch/all.h> #include "registration.h" #if defined(CPU_KERNEL) torch::Tensor my_kernel_cpu_forward( const torch::Tensor& input, const torch::Tensor& weight, double epsilon); #endif torch::Tensor forward(const torch::Tensor& input, const torch::Tensor& weight, double epsilon) { #if defined(CPU_KERNEL) if (input.device().type() == torch::kCPU) { return my_kernel_cpu_forward(input, weight, epsilon); } #endif TORCH_CHECK(false, "Unsupported device type"); } // Register via TORCH_LIBRARY // See: references/huggingface-kernels-integration.md build_toml_pattern: name: "build.toml: Multi-Target CPU Compilation" description: | Each SIMD tier is a separate [kernel.*] section. The generic dispatcher is compiled without SIMD flags; each optimized tier gets its own flags. CRITICAL: The generic kernel MUST list cpu_features.hpp in its src so the header is available at compile time. example: | [general] name = "my-kernel" license = "Apache-2.0" version = 1 backends = ["cpu"] [general.hub] repo-id = "kernels-community/my-kernel" [torch] src = ["torch-ext/torch_binding.cpp"] [kernel.my_kernel_cpu] backend = "cpu" depends = ["torch"] include = ["my_kernel_cpu"] src = [ "my_kernel_cpu/my_kernel_cpu_torch.cpp", "my_kernel_cpu/my_kernel_cpu.cpp", "my_kernel_cpu/my_kernel_cpu.hpp", "my_kernel_cpu/cpu_features.hpp", ] [kernel.my_kernel_cpu_avx2] backend = "cpu" cxx-flags = ["-mavx2", "-mfma", "-mf16c", "-fopenmp"] depends = ["torch"] include = ["my_kernel_cpu"] src = [ "my_kernel_cpu/my_kernel_avx2.cpp", "my_kernel_cpu/my_kernel_avx2.hpp", "my_kernel_cpu/cpu_types_avx2.hpp", ] [kernel.my_kernel_cpu_avx512] backend = "cpu" cxx-flags = [ "-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl", ] depends = ["torch"] include = ["my_kernel_cpu"] src = [ "my_kernel_cpu/my_kernel_avx512.cpp", "my_kernel_cpu/my_kernel_avx512.hpp", "my_kernel_cpu/cpu_types_avx512.hpp", ] compiler_flags_reference: avx2: required: ["-mavx2", "-mfma"] optional: ["-mf16c", "-fopenmp"] avx512_basic: required: ["-mavx512f"] optional: ["-mavx512vl", "-mavx512dq", "-mavx512bw", "-fopenmp"] avx512_bf16: required: ["-mavx512f", "-mavx512bf16"] optional: ["-mavx512vl", "-mfma", "-mf16c", "-fopenmp"] avx512_full: all: ["-mfma", "-fopenmp", "-mf16c", "-mavx512f", "-mavx512bf16", "-mavx512vl", "-mavx512dq", "-mavx512bw", "-mavx512vbmi"] amx: required: ["-mamx-bf16", "-mamx-int8", "-mamx-tile"] combined_with: "avx512_full" -
simd_optimization_patterns.yaml 7 KB
# SIMD Optimization Patterns for AVX2 and AVX512 # # Vector type abstractions and common SIMD patterns extracted from # kernels-community CPU kernels (rmsnorm, flash-attn2, quantization-*). vector_type_abstractions: name: "Typed Vector Wrappers" description: | Wrap raw __m256/__m512 intrinsics in C++ structs for type safety and readability. Each struct knows its element count, supports arithmetic operators, and provides conversion between precision levels. Source: kernels-community/rmsnorm/rmsnorm_cpu/cpu_types_avx512.hpp avx512_types: - name: FP32Vec16 register: __m512 elements: 16 operations: | struct FP32Vec16 { __m512 reg; constexpr static int VEC_ELEM_NUM = 16; FP32Vec16(float v) : reg(_mm512_set1_ps(v)) {} FP32Vec16(__m512 r) : reg(r) {} FP32Vec16 operator+(const FP32Vec16& o) const { return _mm512_add_ps(reg, o.reg); } FP32Vec16 operator*(const FP32Vec16& o) const { return _mm512_mul_ps(reg, o.reg); } float reduce_sum() const { return _mm512_reduce_add_ps(reg); } }; - name: BF16Vec32 register: __m512i elements: 32 operations: | struct BF16Vec32 { __m512i reg; constexpr static int VEC_ELEM_NUM = 32; BF16Vec32(const at::BFloat16* ptr) : reg(_mm512_loadu_si512(ptr)) {} void save(at::BFloat16* ptr) const { _mm512_storeu_si512(ptr, reg); } }; - name: FP16Vec16 register: __m256i elements: 16 note: "FP16 uses 256-bit registers (16 elements) to match AVX512 fp32 width" avx2_types: - name: FP32Vec8 register: __m256 elements: 8 operations: "Same pattern as FP32Vec16 but with _mm256_* intrinsics" conversion_pattern: | // BF16 → FP32 (widen) FP32Vec16::FP32Vec16(const BF16Vec32& v) { reg = _mm512_castsi512_ps( _mm512_slli_epi32( _mm512_cvtepu16_epi32(_mm512_extracti32x8_epi32(v.reg, 0)), 16)); } // FP32 → BF16 (narrow, two FP32Vec16 → one BF16Vec32) BF16Vec32::BF16Vec32(const FP32Vec16& lo, const FP32Vec16& hi) { reg = (__m512i)_mm512_cvtne2ps_pbh(hi.reg, lo.reg); } dispatch_macro: | // Dispatch over scalar types at runtime #define DISPATCH_FLOATING_TYPES(TYPE, NAME, ...) \ AT_DISPATCH_SWITCH(TYPE, NAME, \ AT_DISPATCH_CASE(at::ScalarType::Float, __VA_ARGS__) \ AT_DISPATCH_CASE(at::ScalarType::BFloat16, __VA_ARGS__) \ AT_DISPATCH_CASE(at::ScalarType::Half, __VA_ARGS__)) // Usage: DISPATCH_FLOATING_TYPES(input.scalar_type(), "my_kernel", [&] { using scalar_vec_t = vec_t<scalar_t>; constexpr int VEC_ELEM_NUM = scalar_vec_t::get_elem_num(); // ... }); common_patterns: - name: "Vectorized Reduction (RMSNorm variance)" description: "Compute sum of squares using SIMD, then scalar reduce" code: | FP32Vec16 variance(0.0f); for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) { scalar_vec_t x(input_p + j); FP32Vec16 fp32_x(x); variance = variance + fp32_x * fp32_x; } float s_variance = 1.0f / sqrtf(variance.reduce_sum() / (float)hidden_size + epsilon); - name: "Vectorized Scale+Store (RMSNorm output)" description: "Multiply normalized input by weight, store as original dtype" code: | FP32Vec16 fp32_s_variance(s_variance); for (int j = 0; j < hidden_size; j += VEC_ELEM_NUM) { scalar_vec_t x(input_p + j); scalar_vec_t w(weight + j); FP32Vec16 fp32_x(x); FP32Vec16 fp32_w(w); FP32Vec16 fp32_out = fp32_x * fp32_s_variance * fp32_w; scalar_vec_t out(fp32_out); out.save(output_p + j); } - name: "BF16 Dot Product Accumulation" description: "Use _mm512_dpbf16_ps for fused bf16 multiply-add into fp32" code: | // Accumulate bf16 dot products into fp32: 16 pairs per call, 32 bf16 per vector __m512 acc = _mm512_setzero_ps(); for (int k = 0; k + 32 <= K; k += 32) { __m512bh a = ...; // load 32 bf16 from a + k __m512bh b = ...; // load 32 bf16 from b + k acc = _mm512_dpbf16_ps(acc, a, b); // acc[i] += a[2i]*b[2i] + a[2i+1]*b[2i+1] } note: "_mm512_dpbf16_ps processes 16 bf16 pairs (32 elements) per call; K must be even" - name: "Compile-Time Loop Unrolling" description: "Template-based unrolling for register-tiled micro-kernels" code: | template <typename T, T count, typename F> constexpr void unroll_loop(F&& f) { [&]<T... I>(std::integer_sequence<T, I...>) { (f(std::integral_constant<T, I>{}), ...); }(std::make_integer_sequence<T, count>{}); } // Usage: template <int N> struct Unroll { template <typename F, typename... Args> void operator()(F&& f, Args&&... args) const { unroll_loop<int, N>([&](auto i) { f(i, std::forward<Args>(args)...); }); } }; // In micro-kernel: Unroll<ROWS * COLS>{}(compute, k); - name: "CVT_BF16_TO_FP32 Macro" description: "Convert 16 bf16 values to 16 fp32 values via bit shift" code: | #define CVT_BF16_TO_FP32(a) \ _mm512_castsi512_ps(_mm512_slli_epi32(_mm512_cvtepu16_epi32(a), 16)) note: "BF16 is the upper 16 bits of FP32, so left-shift by 16 is exact conversion" - name: "Prefetch for Memory-Bound Kernels" description: "Software prefetch to hide memory latency" code: | constexpr int PREFETCH_SIZE_K = 16 * 4; // prefetch 4 iterations ahead // Inside inner loop: if constexpr (PREFETCH_SIZE_K > 0) { _mm_prefetch(B + (k + PREFETCH_SIZE_K) * ldb + n, _MM_HINT_T0); } tuning: | Prefetch distance is a Phase 2 tuning parameter. Too small: data not ready when needed. Too large: evicts useful data from cache. Try: 64, 128, 256, 512 bytes. Measure with perf stat L1-dcache-load-misses. at_vec_portable_wrapper: name: "at::vec::Vectorized: requires CPU_CAPABILITY_AVX512" description: | Some kernels use ATen's portable wrapper at::vec::Vectorized<T> (e.g. convert_from_float, or Vectorized<float>::exp() for silu/gelu) instead of raw _mm512_* intrinsics. Unlike raw intrinsics, at::vec selects its SIMD vs scalar implementation from the preprocessor macro CPU_CAPABILITY_AVX512: NOT from the -mavx512f / -march=native compiler flags. Without that macro, at::vec silently compiles to scalar vec_base.h: in particular exp() degrades from Sleef_expf16_u10 (16-wide) to scalar std::expf, which can make the kernel ~2x slower with no visible error. fix: | // For any TU using at::vec, define the macro before including vec.h: #define CPU_CAPABILITY_AVX512 #include <ATen/cpu/vec/vec.h> // OR add -DCPU_CAPABILITY_AVX512 to that section's cxx-flags in build.toml. see_also: "build_system.md → 'at::vec::Vectorized needs a CPU_CAPABILITY macro' and the objdump/nm build self-check." -
threading_patterns.yaml 3 KB
# OpenMP Threading Patterns for CPU Kernels basic_parallel_for: name: "Token-Level Parallelism" description: | The most common pattern: parallelize over the batch/token dimension. Each token is independent, so no synchronization needed. code: | #pragma omp parallel for for (int i = 0; i < num_tokens; ++i) { // Process token i (independent) auto input_p = input + i * hidden_size; auto output_p = out + i * hidden_size; // ... SIMD inner loop over hidden_size ... } applies_to: "RMSNorm, LayerNorm, activations, most elementwise ops" size_threshold: name: "Small Tensor Guard" description: | OpenMP fork/join costs ~10-50us. For small tensors, serial is faster. code: | if (num_tokens >= OMP_THRESHOLD) { #pragma omp parallel for for (int i = 0; i < num_tokens; ++i) { ... } } else { for (int i = 0; i < num_tokens; ++i) { ... } } tuning: | OMP_THRESHOLD depends on hidden_size: - hidden_size >= 4096: threshold = 4 (parallel pays off quickly) - hidden_size >= 1024: threshold = 8 - hidden_size < 1024: threshold = 16-32 reduction: name: "Parallel Reduction" description: | For operations that accumulate across tokens (e.g., grad_weight in RMSNorm backward). code: | // Thread-local accumulators to avoid false sharing #pragma omp parallel { std::vector<float> local_grad_w(hidden_size, 0.0f); #pragma omp for for (int i = 0; i < num_tokens; ++i) { for (int j = 0; j < hidden_size; ++j) { local_grad_w[j] += grad_output[i * hidden_size + j] * normalized[i * hidden_size + j]; } } #pragma omp critical for (int j = 0; j < hidden_size; ++j) { grad_weight[j] += local_grad_w[j]; } } alternative: | Use #pragma omp parallel for reduction(+:sum) for simple scalar reductions. For array reductions, OpenMP 4.5+ supports array reduction but compiler support varies: thread-local + critical is more portable. nested_parallelism: name: "Avoid Nested Parallelism" description: | Do NOT nest #pragma omp parallel regions. This oversubscribes cores. rule: | If the outer loop is already parallel, the inner loop must be serial. The inner loop uses SIMD instead of threads. environment: name: "Environment Variables" variables: OMP_NUM_THREADS: "Number of threads (default: num cores)" OMP_PROC_BIND: "Thread binding (close, spread, master)" OMP_SCHEDULE: "Loop scheduling (static, dynamic, guided)" recommended: | # For benchmarking: export OMP_NUM_THREADS=$(nproc) export OMP_PROC_BIND=close export OMP_SCHEDULE=static # Pin benchmarks to a single NUMA node (bind by node, so it adapts to any # cores-per-node). Cross-socket traffic otherwise makes timings noisy, and on # big-L3 Xeons a cache-resident working set produces misleading speedups. numactl --cpunodebind=0 --membind=0 python benchmark.py -
workflow_details.md 8.7 KB
# Workflow details (CPU kernels) ## Analysis Given a target PyTorch operation: 1. Parse the operation: input and output shapes and dtypes, the mathematical operations (matmul, activation, reduction), the memory access pattern (row-wise, column-wise, random), and the compute intensity (FLOPs per byte moved). 2. Read the knowledge base in `references/`: | File | Content | |---|---| | `runtime_dispatch.yaml` | `cpu_features.hpp` pattern, dispatch tiers | | `build_system.md` | `build.toml` multi-target compilation | | `implementation_reference.md` | C++ templates, `Unroll<N>`, tinygemm | | `correctness.yaml` | Constraints every kernel must hold | | `simd_optimization_patterns.yaml` | AVX512 vector abstractions | | `brgemm_patterns.yaml` | brgemm API (GEMM kernels only) | 3. Run `python scripts/analyze_op.py --op <op_name> --shapes <shapes>` to classify the kernel type, then pick the matching template in `references/implementation_reference.md`. ## Design 1. Identify the kernel type: element-wise, reduction, GEMM, or attention. 2. Select the strategy: element-wise uses AVX512 vectorization with OpenMP; GEMM uses the tinygemm and brgemm dual path with `parallel_2d`; attention uses tiled attention with brgemm. 3. Apply the constraints from `references/correctness.yaml`: unaligned loads (`_mm512_loadu_*`), one ISA tier per translation unit, tail handling for sizes that are not a multiple of the vector width, `registration.h` in `torch_binding.cpp`, and a per-kernel `cpu_features.hpp` in the kernel's own namespace. ## Trial loop For each trial: ### a. Implement or modify the kernel Start from a template in `references/implementation_reference.md` or modify the previous trial. ### b. Validate ```bash python scripts/validate_cpu_kernel.py . ``` A validation failure is fixed in place and does not count as a trial. ### c. Build ```bash kernel-builder build --release pip install dist/*.whl --force-reinstall ``` ### d. Save the trial ```bash python scripts/trial_manager.py save <kernel_name> <kernel_dir> --parent <parent_id> --strategy "description" ``` Omit `--parent` for the first trial. ### e. Benchmark ```bash # Trial t0 measures both baseline and kernel: python scripts/benchmark_cpu.py baseline.py --kernel-package my_kernel --op my_kernel.forward # Trials t1 and later reuse the cached baseline time: python scripts/trial_manager.py baseline-us <kernel_name> python scripts/benchmark_cpu.py baseline.py --kernel-package my_kernel --op my_kernel.forward --baseline-us <cached_value> ``` ### f. Record the result ```bash python scripts/trial_manager.py result <kernel_name> <trial_id> \ --correctness <pass|fail> --speedup <float> \ --baseline_us <float> --kernel_us <float> ``` ### g. Decide the next trial | Condition | Action | |---|---| | Speedup above `early_stop_speedup` | Stop. This is the only valid early stop. | | Speedup improved | Continue on this branch with the next optimization. | | Speedup regressed | Branch back to the best trial and try a different strategy. | | Correctness failed | Fix on the same branch. Read the leaf path in the benchmark output; the usual causes are alignment and tail handling. | | After t1, when `perf_stat_enabled` is true | Run `cpu_profiler.py` once before choosing the next optimization. | | IPC under 1.0 together with a high L1 or LLC miss rate | Memory bound: add prefetch or reduce the tile size. A pure AVX512 FMA loop has low IPC by design, so IPC alone does not decide. | | L1 miss rate above 10% | Tile too large: reduce it to fit L1 (48 KB per core on the target Xeon). | | LLC miss rate above 20% | Working set too large: add cache blocking within the 1 MB L2 budget. | | Plateau after two or more trials | Switch algorithm (tinygemm or brgemm, different blocking). | | `max_trials` reached | Stop. Every trial in `config.yaml` must run. | ### h. Check status ```bash python scripts/trial_manager.py status <kernel_name> python scripts/trial_manager.py best <kernel_name> ``` ## Trial manager commands ```bash python scripts/trial_manager.py init <kernel_name> <baseline_file> python scripts/trial_manager.py save <kernel_name> <source> [--parent <parent_id>] [--strategy "..."] python scripts/trial_manager.py result <kernel_name> <trial_id> [--correctness pass] [--speedup 3.2] [--baseline_us 150.0] [--kernel_us 47.0] python scripts/trial_manager.py status <kernel_name> python scripts/trial_manager.py best <kernel_name> python scripts/trial_manager.py baseline-us <kernel_name> python scripts/trial_manager.py finalize <kernel_name> <output_path> ``` ## Benchmarking `scripts/benchmark_cpu.py` runs two checks, and both must pass for a trial to count as completed. Correctness compares the kernel output to the baseline output with a structure-aware comparator. Tuples, lists, and dicts are walked element-wise and must agree in type, length, and keys. Every tensor leaf must agree in dtype and shape and is compared in its own dtype against a per-dtype tolerance: half an ulp relative for bf16 and fp16, `atol=1e-6, rtol=1e-5` for fp32, `atol=1e-12, rtol=1e-9` for fp64, and exact equality for integer and bool dtypes. A mismatch names the leaf path, for example `output[1]`. Pass `--atol` and `--rtol` to widen every floating tolerance when the kernel's accumulation order legitimately differs from the reference. The baseline must define `get_inputs()` and either `get_reference_output()` or a `Model` class. Performance times both implementations with `torch.utils.benchmark.Timer.blocked_autorange(min_run_time=2.0)` and reports the median time and the speedup. `python scripts/benchmark_cpu.py --self-check` exercises the comparator on a wrong second tuple element, a bf16 truncation error, a scalar reference paired with a tensor kernel output, and the structure, dtype, and shape checks; it exits 1 if any expectation fails and exits 2 with a stated reason when torch is not installed, so the check fails closed on a host without torch. ## Profiling with perf stat ```bash python scripts/cpu_profiler.py --kernel-package my_kernel --op my_kernel.forward ``` The script runs `perf stat`, collects hardware counters, and maps each finding to a reference file. Run it after the first benchmarked trial (t1), again when the speedup plateaus after two or more further trials, and whenever the next optimization level is unclear. It reports: | Counter | What it tells you | |---|---| | IPC (instructions per cycle) | Compute versus memory bound, read together with the miss rates | | L1 cache miss rate | Tile sizing | | LLC (L3) miss rate | Working set size | | Branch miss rate | SIMD versus scalar branching | The output names a reference file for each recommendation, for example: ``` >> IPC < 1.0: memory-bound or dependency-bound if L1 or LLC miss rates are also high; - Add prefetch instructions (_mm_prefetch with _MM_HINT_T0 or _MM_HINT_T1) - Reduce cache blocking tile size Reference: references/memory_patterns.yaml ``` Read the referenced file and apply the pattern in the next trial. ## Skill layout ``` cpu-kernel-authoring/ ├── SKILL.md # Contract, procedure, failure table ├── references/ # Knowledge base │ ├── implementation_reference.md # C++ templates, Unroll<N>, tinygemm │ ├── optimization_strategies.md # Levels, decision tree, checklist │ ├── workflow_details.md # This file │ ├── build_system.md # build.toml multi-target compilation │ ├── runtime_dispatch.yaml # cpu_features.hpp and dispatch │ ├── correctness.yaml # Constraints │ ├── simd_optimization_patterns.yaml # AVX512 vector abstractions │ ├── quantized_gemm_patterns.yaml # LUT, tinygemm, brgemm │ ├── brgemm_patterns.yaml # brgemm API, VNNI packing │ ├── memory_patterns.yaml # Prefetch, cache blocking │ ├── threading_patterns.yaml # OpenMP patterns │ ├── dtype_optimizations.yaml # bf16, fp8, int8 handling │ ├── optimization_levels.yaml # L1 to L4 checklist │ └── huggingface-kernels-integration.md # Hub integration └── scripts/ # Tools the procedure runs; do not recreate them ├── config.py # Shared config loader ├── config.yaml # Session config ├── analyze_op.py # Op analysis: kernel type and strategy ├── validate_cpu_kernel.py # Static checks on C++ kernel code ├── benchmark_cpu.py # Correctness and performance ├── cpu_profiler.py # perf stat and recommendations └── trial_manager.py # Trial tree management ```
-
-
scripts
-
analyze_op.py 11 KB
#!/usr/bin/env python3 """ Analyze a PyTorch operation to guide CPU kernel development. Extracts compute/memory characteristics, identifies kernel type, and recommends SIMD strategy and optimization approach. Usage: python scripts/analyze_op.py --op "rms_norm" --shapes "1024x4096,2048x8192" python scripts/analyze_op.py --file baseline.py """ import argparse import ast import re import sys from pathlib import Path class OpAnalyzer(ast.NodeVisitor): """AST visitor to analyze PyTorch model operations.""" def __init__(self): self.operations = [] self.shapes = {} self.dtypes = set() self.has_matmul = False self.has_linear = False self.activations = [] self.reductions = [] self.elementwise = [] def visit_Call(self, node): """Visit function calls to identify operations.""" if isinstance(node.func, ast.Attribute): if hasattr(node.func.value, "id") and node.func.value.id == "torch": op_name = node.func.attr self.operations.append(op_name) if op_name in ("matmul", "mm", "bmm"): self.has_matmul = True elif op_name in ("sum", "mean", "max", "min", "norm"): self.reductions.append(op_name) elif op_name in ("sigmoid", "tanh", "relu", "gelu", "silu"): self.activations.append(op_name) elif op_name in ("clamp", "abs", "exp", "rsqrt", "sqrt"): self.elementwise.append(op_name) elif hasattr(node.func.value, "attr"): if node.func.value.attr == "functional": op_name = node.func.attr self.operations.append(f"F.{op_name}") if op_name in ("gelu", "relu", "silu", "softmax", "sigmoid"): self.activations.append(op_name) if op_name in ("linear",): self.has_linear = True self.generic_visit(node) def visit_BinOp(self, node): """Visit binary operations.""" op_map = { ast.Mult: "multiply", ast.Div: "divide", ast.Add: "add", ast.Sub: "subtract", } op_type = type(node.op) if op_type in op_map: self.elementwise.append(op_map[op_type]) self.generic_visit(node) def visit_Assign(self, node): """Visit assignments to track nn.Linear.""" if ( isinstance(node.value, ast.Call) and hasattr(node.value.func, "attr") and node.value.func.attr == "Linear" ): self.has_linear = True self.generic_visit(node) def analyze_from_file(filepath: Path) -> dict[str, object]: """Analyze PyTorch file and extract optimization hints.""" with open(filepath, "r") as f: source = f.read() tree = ast.parse(source) analyzer = OpAnalyzer() analyzer.visit(tree) shapes = {} for line in source.split("\n"): if "=" in line and any( dim in line for dim in [ "batch_size", "in_features", "out_features", "hidden_size", "input_size", "seq_len", "num_heads", ] ): match = re.match(r"(\w+)\s*=\s*(\d+)", line.strip()) if match: shapes[match.group(1)] = int(match.group(2)) return _build_analysis(analyzer, shapes) def analyze_from_args(op_name: str, shapes_str: str) -> dict[str, object]: """Analyze operation from command-line args.""" analyzer = OpAnalyzer() op_lower = op_name.lower().replace("_", "").replace("-", "") if op_lower in ("rmsnorm", "layernorm", "rms_norm", "layer_norm"): analyzer.reductions.append("norm") analyzer.elementwise.extend(["multiply", "rsqrt"]) elif op_lower in ("softmax",): analyzer.reductions.extend(["max", "sum"]) analyzer.elementwise.extend(["exp", "divide"]) elif "gemm" in op_lower or "linear" in op_lower or "matmul" in op_lower: analyzer.has_matmul = True elif "attention" in op_lower or "flashatt" in op_lower or "flash_att" in op_lower: analyzer.has_matmul = True analyzer.reductions.append("softmax") elif "gelu" in op_lower or "silu" in op_lower or "relu" in op_lower: analyzer.activations.append(op_lower) elif "moe" in op_lower or "megablocks" in op_lower: analyzer.has_matmul = True shapes = {} if shapes_str: for i, s in enumerate(shapes_str.split(",")): dims = s.strip().split("x") if len(dims) == 2: shapes[f"shape_{i}"] = f"{dims[0]}x{dims[1]}" elif len(dims) == 3: shapes[f"shape_{i}"] = f"{dims[0]}x{dims[1]}x{dims[2]}" return _build_analysis(analyzer, shapes) def _build_analysis(analyzer: OpAnalyzer, shapes: dict[str, int | str]) -> dict[str, object]: """Build analysis result from analyzer state.""" kernel_type = "unknown" if analyzer.has_matmul or analyzer.has_linear: if analyzer.activations or analyzer.elementwise: kernel_type = "gemm_fused" elif analyzer.reductions: kernel_type = "attention" else: kernel_type = "gemm" elif analyzer.reductions: kernel_type = "reduction" elif analyzer.elementwise or analyzer.activations: kernel_type = "elementwise" return { "kernel_type": kernel_type, "operations": analyzer.operations, "activations": analyzer.activations, "reductions": analyzer.reductions, "elementwise": analyzer.elementwise, "shapes": shapes, "has_gemm": analyzer.has_matmul or analyzer.has_linear, } def print_analysis(analysis: dict[str, object]): """Pretty print the analysis results with CPU-specific recommendations.""" print(f"\n{'=' * 70}") print("CPU Kernel Analysis") print(f"{'=' * 70}\n") print(f"Kernel Type: {analysis['kernel_type'].upper()}") print() if analysis["shapes"]: print("Shapes:") for key, val in analysis["shapes"].items(): print(f" {key}: {val}") print() print("Operations:") if analysis["has_gemm"]: print(" * GEMM/Linear detected") if analysis["activations"]: print(f" * Activations: {', '.join(set(analysis['activations']))}") if analysis["reductions"]: print(f" * Reductions: {', '.join(set(analysis['reductions']))}") if analysis["elementwise"]: print(f" * Elementwise: {', '.join(set(analysis['elementwise']))}") print() print("CPU Optimization Strategy:") if analysis["kernel_type"] == "attention": print(" Kernel Category: Attention (Flash-Attention style)") print(" Architecture: Tiled attention with brgemm for Q@K and S@V") print(" Blocking: BLOCK_M=256, BLOCK_N=768 (attention-specific)") print(" Threading: parallel over batch * heads * M-tiles") print(" Requirements: AVX512 + AMX (via brgemm)") print() print(" Reference files:") print(" - references/brgemm_patterns.yaml") print(" - references/memory_patterns.yaml") elif analysis["has_gemm"]: print(" Kernel Category: GEMM") print(" Architecture:") print(" - tinygemm path (M <= 4): fused dequant + _mm512_dpbf16_ps") print(" - brgemm path (M > 4): pre-dequant B + at::native::cpublas::brgemm()") print(" Threading: parallel_2d(m, n, compute_fn) with 2D factorization") print(" Compiler flags: -mavx512f -mavx512bf16 -mavx512vl -mavx512dq -mavx512bw") print(" -mavx512vbmi -mfma -mf16c -fopenmp") print(" Note: AMX flags NOT needed — brgemm dispatches to AMX via oneDNN internally") print(" Key patterns: Unroll<N>, tinygemm_kernel_nn, cpu_types_avx512.hpp") print() print(" Reference files:") print(" - references/brgemm_patterns.yaml") print(" - references/quantized_gemm_patterns.yaml") print(" - references/threading_patterns.yaml") elif analysis["kernel_type"] in ("reduction", "elementwise"): print(" Kernel Category: Element-wise / Reduction") print(" Architecture: Direct AVX512 intrinsics (no brgemm)") print(" Threading: #pragma omp parallel for over rows") print(" Vectorization: FP32Vec16, BF16Vec32 abstractions") print(" Compiler flags: -mavx512f -mavx512bf16 -mavx512vl -mavx512dq -mavx512bw") print(" -mavx512vbmi -mfma -mf16c -fopenmp") print(" Prefetch: _MM_HINT_T1 (L2)") print() print(" Reference files:") print(" - references/simd_optimization_patterns.yaml") print(" - references/memory_patterns.yaml") print(" - references/threading_patterns.yaml") print() print("File Structure:") print(" my_kernel_cpu/") print(" ├── cpu_features.hpp # CPUID detection (own namespace)") print(" ├── my_kernel_cpu.cpp # Dispatcher (runtime feature check)") print(" ├── my_kernel_cpu.hpp # Shared declarations") print(" ├── my_kernel_cpu_torch.cpp # Python ↔ C++ bridge") if not analysis["has_gemm"]: print(" ├── my_kernel_avx2.cpp # AVX2 implementation (optional)") print(" ├── my_kernel_avx2.hpp") print(" ├── my_kernel_avx512.cpp # AVX512 implementation") print(" └── my_kernel_avx512.hpp") print() print(" torch_binding.cpp # Op registration (registration.h)") print(" build.toml # Multi-target compilation config") print() print("Relevant Reference Files:") print(" - references/runtime_dispatch.yaml # cpu_features.hpp pattern") print(" - references/correctness.yaml # Critical constraints") print(" - references/implementation_reference.md # C++ templates") if analysis["has_gemm"]: print(" - references/brgemm_patterns.yaml # brgemm API") print(" - references/quantized_gemm_patterns.yaml # 4-bit GEMM") print(" - references/optimization_levels.yaml # Progressive optimization") print() def main(): parser = argparse.ArgumentParser(description="Analyze PyTorch op for CPU kernel development") parser.add_argument("--op", type=str, help="Operation name (e.g., rms_norm, flash_attention)") parser.add_argument("--shapes", type=str, default="", help="Shapes as MxN,MxN (e.g., 1024x4096,2048x8192)") parser.add_argument("--file", type=Path, help="PyTorch baseline file to analyze") args = parser.parse_args() if args.file: if not args.file.exists(): print(f"Error: File not found: {args.file}") sys.exit(1) analysis = analyze_from_file(args.file) elif args.op: analysis = analyze_from_args(args.op, args.shapes) else: print("Error: Provide either --op or --file") sys.exit(1) print_analysis(analysis) if __name__ == "__main__": main() -
benchmark_cpu.py 16.7 KB
#!/usr/bin/env python3 """ Benchmark a CPU kernel against its PyTorch baseline. Checks correctness with a structure- and dtype-aware comparator, then measures performance with torch.utils.benchmark. Usage: python scripts/benchmark_cpu.py baseline.py --kernel-package my_kernel --op my_kernel.forward python scripts/benchmark_cpu.py baseline.py --kernel-package my_kernel --op my_kernel.forward --baseline-us 123.45 python scripts/benchmark_cpu.py --self-check The first trial measures both baseline and kernel. Later trials pass the cached baseline time with --baseline-us so only the kernel is timed. """ import argparse import importlib import importlib.util import logging import sys from pathlib import Path # torch is imported at module load inside try/except so that --self-check can # exit 2 with a clear reason on a host without torch instead of dying with an # ImportError traceback at import time. try: import torch except ImportError: # pragma: no cover - exercised only on hosts without torch torch = None def _load_module(filepath: Path, module_name: str): spec = importlib.util.spec_from_file_location(module_name, filepath) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) return module def _load_kernel_func(kernel_package: str, op_path: str): """Resolve 'package.attr[.attr]' to the callable inside the installed package.""" parts = op_path.split(".") if len(parts) < 2: print("Error: --op should be package.function (e.g., my_kernel.forward)") sys.exit(1) func = importlib.import_module(parts[0]) for attr in parts[1:]: func = getattr(func, attr) return func def _default_tolerances(): # Tolerances are keyed by the leaf's own dtype so that a bf16 output is judged # at bf16 resolution. Upcasting both sides to float32 and using one loose # tolerance hid one-ulp bf16 errors (a truncating conversion instead of # round-to-nearest-even is 7.8e-3 relative, under the old 1e-2 atol). # For bf16 and fp16 the rtol is half an ulp of the dtype, so any full-ulp # difference fails. Widen with --rtol when the kernel's accumulation order # legitimately differs from the reference. return { torch.bfloat16: (0.0, 2.0**-9), torch.float16: (0.0, 2.0**-11), torch.float32: (1e-6, 1e-5), torch.float64: (1e-12, 1e-9), } def _leaf_mismatch(ref, out, path, tolerances): """Return a mismatch message for one tensor leaf, or None when it matches.""" if not isinstance(out, torch.Tensor): return f"{path}: reference is a tensor, kernel returned {type(out).__name__}" if ref.dtype != out.dtype: return f"{path}: dtype mismatch ref={ref.dtype} kernel={out.dtype}" if ref.shape != out.shape: return f"{path}: shape mismatch ref={tuple(ref.shape)} kernel={tuple(out.shape)}" if ref.dtype in tolerances: atol, rtol = tolerances[ref.dtype] # allclose is evaluated in the leaf's own dtype: no upcast, so the # tolerance is applied at the resolution the consumer will see. if torch.allclose(ref, out, atol=atol, rtol=rtol, equal_nan=True): return None # The diff is widened to float64 only for reporting, where exactness of # the printed number matters and rounding in-dtype would mislead. diff = (ref.double() - out.double()).abs() flat = diff.argmax().item() idx = [] remaining = flat for dim in reversed(ref.shape): idx.insert(0, remaining % dim) remaining //= dim return ( f"{path}: value mismatch dtype={ref.dtype} atol={atol:g} rtol={rtol:g} " f"max_diff={diff.max().item():.6e} mean_diff={diff.mean().item():.6e} " f"worst at {tuple(idx)}: ref={ref.flatten()[flat].item():.6g} " f"kernel={out.flatten()[flat].item():.6g}" ) # Integer and bool leaves carry no rounding, so any difference is a bug. if torch.equal(ref, out): return None diff_count = (ref != out).sum().item() return f"{path}: {diff_count} element(s) differ in exact dtype {ref.dtype}" def compare_structured(ref, out, tolerances, path="output"): """Walk ref and out together; return a list of mismatch messages (empty means equal). Tuples, lists, and dicts are walked element-wise and must agree in type, length, and key set. Every tensor leaf must agree in dtype and shape and fall within the tolerance for its dtype. Each message names the leaf path so a wrong second output is reported as such instead of being dropped. """ if isinstance(ref, torch.Tensor): msg = _leaf_mismatch(ref, out, path, tolerances) return [msg] if msg else [] if isinstance(ref, (tuple, list)): if type(out) is not type(ref): return [f"{path}: reference is {type(ref).__name__}, kernel returned {type(out).__name__}"] if len(ref) != len(out): return [f"{path}: length mismatch ref={len(ref)} kernel={len(out)}"] mismatches = [] for i, (r, o) in enumerate(zip(ref, out)): mismatches += compare_structured(r, o, tolerances, f"{path}[{i}]") return mismatches if isinstance(ref, dict): if not isinstance(out, dict): return [f"{path}: reference is dict, kernel returned {type(out).__name__}"] if set(ref) != set(out): return [f"{path}: key mismatch ref={sorted(map(str, ref))} kernel={sorted(map(str, out))}"] mismatches = [] for k in ref: mismatches += compare_structured(ref[k], out[k], tolerances, f"{path}[{k!r}]") return mismatches # ref cannot be a tensor here (dispatched above), but `out` can: `!=` # against a tensor is element-wise, and a multielement result raises the # ambiguous truth-value RuntimeError instead of naming the path. if isinstance(out, torch.Tensor): return [f"{path}: reference is {type(ref).__name__}, kernel returned {type(out).__name__}"] if ref != out: return [f"{path}: scalar mismatch ref={ref!r} kernel={out!r}"] return [] def _baseline_pair(baseline_mod, inputs): """Reference output and a timed callable, building any Model exactly once.""" if hasattr(baseline_mod, "get_reference_output"): return baseline_mod.get_reference_output(*inputs), baseline_mod.get_reference_output if hasattr(baseline_mod, "Model"): init_inputs = baseline_mod.get_init_inputs() if hasattr(baseline_mod, "get_init_inputs") else [] model = baseline_mod.Model(*init_inputs) model.eval() with torch.no_grad(): ref_output = model(*inputs) return ref_output, (lambda *args: model(*args)) raise AttributeError("baseline must define get_reference_output() or a Model class") def run_correctness(inputs, ref_output, kernel_func, tolerances): """Compare the kernel output to the baseline output; return True when they match.""" print("\n Correctness Check (per-dtype tolerances)") for dtype, (atol, rtol) in tolerances.items(): print(f" {dtype}: atol={atol:g} rtol={rtol:g}") print(" integer and bool dtypes: exact") try: with torch.no_grad(): kernel_output = kernel_func(*inputs) mismatches = compare_structured(ref_output, kernel_output, tolerances) if mismatches: print(" FAIL:") for m in mismatches: print(f" {m}") return False print(" PASS: structure, dtype, shape, and values match") return True except Exception: logging.getLogger(__name__).exception("Kernel correctness check failed") return False def run_performance(inputs, kernel_func, ref_func, baseline_us=None, warmup=10, iters=100): """Time baseline and kernel with torch.utils.benchmark; return (baseline_us, kernel_us, speedup).""" from torch.utils.benchmark import Timer print(f"\n Performance Benchmark (warmup={warmup}, iters={iters})") if baseline_us is not None: # The baseline does not change between trials, so re-timing it only adds # noise and wall time to the trial loop. print(f" Using cached baseline: {baseline_us:.2f} us") bl_us = baseline_us else: # Autograd state is measurement overhead: no timed iteration builds a graph. with torch.no_grad(): bl_timer = Timer( stmt="ref_func(*inputs)", globals={"ref_func": ref_func, "inputs": inputs}, label="Baseline", description="PyTorch", num_threads=torch.get_num_threads(), ) bl_result = bl_timer.blocked_autorange(min_run_time=2.0) bl_us = bl_result.median * 1e6 print(f" Baseline: {bl_us:.2f} us (median)") with torch.no_grad(): kr_timer = Timer( stmt="kernel_func(*inputs)", globals={"kernel_func": kernel_func, "inputs": inputs}, label="Kernel", description="CPU Kernel", num_threads=torch.get_num_threads(), ) kr_result = kr_timer.blocked_autorange(min_run_time=2.0) kr_us = kr_result.median * 1e6 print(f" Kernel: {kr_us:.2f} us (median)") speedup = bl_us / kr_us if kr_us > 0 else 0 marker = "+" if speedup >= 1.0 else "-" print(f" Speedup: {speedup:.2f}x {marker}") return bl_us, kr_us, speedup def _legacy_first_element_float32_close(ref, out, atol=1e-2, rtol=1e-2): # The comparator this file replaced: first tuple element only, both sides # upcast to float32, one loose tolerance. Kept only so the self-check can # show what it let through. ref = ref[0] if isinstance(ref, tuple) else ref out = out[0] if isinstance(out, tuple) else out return torch.allclose(ref.float(), out.float(), atol=atol, rtol=rtol) def self_check(): """Prove the comparator catches the defects the old one hid. Exit 0 when every expectation holds, 1 when one fails, and 2 when torch is missing so the check fails closed on a host without torch. """ if torch is None: print("FAIL: torch is not installed; the comparator operates on torch tensors and cannot be exercised here.") return 2 torch.manual_seed(0) tolerances = _default_tolerances() failures = [] def expect(name, condition): print(f" {'ok ' if condition else 'FAIL'} {name}") if not condition: failures.append(name) print("Self-check: identical structured output passes") a = torch.randn(4, 8).to(torch.bfloat16) b = torch.randn(4, 8) expect("identical (bf16, fp32) tuple passes", compare_structured((a, b), (a.clone(), b.clone()), tolerances) == []) expect("identical dict of tensors passes", compare_structured({"y": a, "n": 3}, {"y": a.clone(), "n": 3}, tolerances) == []) print("Self-check: wrong second tuple element is caught") wrong_second = (a.clone(), b + 0.5) legacy_pass = _legacy_first_element_float32_close((a, b), wrong_second) mismatches = compare_structured((a, b), wrong_second, tolerances) expect("legacy comparator passed the wrong second element", legacy_pass) expect("structured comparator fails", len(mismatches) == 1) expect("mismatch names path output[1]", bool(mismatches) and mismatches[0].startswith("output[1]:")) for m in mismatches: print(f" {m}") print("Self-check: bf16 truncation instead of round-to-nearest-even is caught") # Values in [1, 2) so that every element has the same exponent and the # rounding-mode error is exactly one bf16 ulp (2^-7 = 7.8e-3) on the # elements whose low 16 bits round up. float32 upcasting with atol=1e-2 # accepts a 7.8e-3 error; half-ulp rtol in bf16 does not. x = torch.rand(64, 64) + 1.0 ref_rne = x.to(torch.bfloat16) truncated = (x.view(torch.int32) & -65536).view(torch.float32).to(torch.bfloat16) differing = (ref_rne != truncated).sum().item() expect(f"fixture differs in {differing} elements (needs > 0)", differing > 0) legacy_pass = _legacy_first_element_float32_close(ref_rne, truncated) mismatches = compare_structured(ref_rne, truncated, tolerances) expect("legacy float32 comparator passed the truncated bf16", legacy_pass) expect("structured comparator fails in bf16", len(mismatches) == 1) for m in mismatches: print(f" {m}") print("Self-check: structure, dtype, and shape are enforced") expect("dtype mismatch fails", compare_structured(a, a.float(), tolerances) != []) expect("shape mismatch fails", compare_structured(b, b.t().contiguous(), tolerances) != []) expect("tuple length mismatch fails", compare_structured((a, b), (a,), tolerances) != []) expect("dict key mismatch fails", compare_structured({"y": a}, {"z": a}, tolerances) != []) expect("list vs tuple fails", compare_structured([a], (a,), tolerances) != []) i = torch.arange(10) expect("integer off-by-one fails", compare_structured(i, i + (i == 3).long(), tolerances) != []) scalar_vs_tensor = compare_structured(2.0, torch.tensor([2.0, 2.0]), tolerances) expect( "scalar reference vs multielement tensor fails and names the path", len(scalar_vs_tensor) == 1 and scalar_vs_tensor[0].startswith("output:"), ) if failures: print(f"\nSelf-check FAILED: {len(failures)} expectation(s) not met") return 1 print("\nSelf-check passed") return 0 def main(): parser = argparse.ArgumentParser(description="Benchmark CPU kernel against PyTorch baseline") parser.add_argument("baseline_file", type=Path, nargs="?", help="PyTorch baseline file") parser.add_argument("--kernel-package", help="Kernel package name (pip-installed)") parser.add_argument("--op", help="Kernel function path (e.g., my_kernel.forward)") parser.add_argument("--baseline-us", type=float, default=None, help="Cached baseline time in microseconds") parser.add_argument("--atol", type=float, default=None, help="Override absolute tolerance for every floating dtype") parser.add_argument("--rtol", type=float, default=None, help="Override relative tolerance for every floating dtype") parser.add_argument("--self-check", action="store_true", help="Run the comparator self-check and exit") args = parser.parse_args() if args.self_check: sys.exit(self_check()) if torch is None: print("Error: torch is not installed") sys.exit(1) if args.baseline_file is None or not args.kernel_package or not args.op: parser.error("baseline_file, --kernel-package, and --op are required") if not args.baseline_file.exists(): print(f"Error: Baseline file not found: {args.baseline_file}") sys.exit(1) tolerances = _default_tolerances() if args.atol is not None or args.rtol is not None: tolerances = { dtype: (args.atol if args.atol is not None else atol, args.rtol if args.rtol is not None else rtol) for dtype, (atol, rtol) in tolerances.items() } print(f"\n{'=' * 70}") print("CPU Kernel Benchmark") print(f"{'=' * 70}") print(f"Baseline: {args.baseline_file}") print(f"Kernel package: {args.kernel_package}") print(f"Op: {args.op}") print(f"Threads: {torch.get_num_threads()}") baseline_mod = _load_module(args.baseline_file, "baseline") try: kernel_func = _load_kernel_func(args.kernel_package, args.op) except (ImportError, AttributeError) as e: print(f"\nError loading kernel: {e}") print(f"Make sure '{args.kernel_package}' is installed: pip install dist/*.whl --force-reinstall") sys.exit(1) print(f"\n{'=' * 70}") print("Correctness") print(f"{'=' * 70}") if not hasattr(baseline_mod, "get_inputs"): print("Error: baseline must define get_inputs()") sys.exit(1) inputs = baseline_mod.get_inputs() ref_output, ref_func = _baseline_pair(baseline_mod, inputs) correct = run_correctness(inputs, ref_output, kernel_func, tolerances) print(f"\n Result: {'PASSED' if correct else 'FAILED'}") print(f"\n{'=' * 70}") print("Performance") print(f"{'=' * 70}") bl_us, kr_us, speedup = run_performance(inputs, kernel_func, ref_func, baseline_us=args.baseline_us) print(f"\n{'=' * 70}") print("Summary") print(f"{'=' * 70}") print(f"Correctness: {'PASSED' if correct else 'FAILED'}") if speedup is not None: print(f"Baseline: {bl_us:.2f} us") print(f"Kernel: {kr_us:.2f} us") print(f"Speedup: {speedup:.2f}x") print() if correct and speedup is not None and speedup >= 1.0: print("All checks passed!") sys.exit(0) print("Some checks FAILED - see output above") sys.exit(1) if __name__ == "__main__": main() -
config.py 1014 B
"""Shared config loader: reads config.yaml from the scripts directory (adjacent to this module).""" from pathlib import Path import yaml _CONFIG_DIR = Path(__file__).resolve().parent _DEFAULTS = { "max_trials": 8, "early_stop_speedup": 3.0, "perf_stat_enabled": True, "vtune_enabled": False, "vtune_bin": "/opt/intel/oneapi/vtune/latest/bin64/vtune", "build_command": "kernel-builder build --release", "install_command": "pip install dist/*.whl --force-reinstall --no-deps", } def load_config() -> dict[str, object]: """Load config.yaml; missing keys fall back to defaults, a missing file is an error.""" config_path = _CONFIG_DIR / "config.yaml" if not config_path.exists(): raise FileNotFoundError( f"config file not found: {config_path}. The skill stops and reports " "rather than assuming trial and profiling settings." ) with open(config_path) as f: cfg = yaml.safe_load(f) or {} return {**_DEFAULTS, **cfg} -
config.yaml 641 B
# CPU kernel optimization session configuration. # Edit these values to control optimization sessions. # Phase 2 optimization max_trials: 8 # Maximum number of Phase 2 optimization trials (4-16) early_stop_speedup: 3.0 # Speedup vs PyTorch baseline to allow early stop # Profiling perf_stat_enabled: true # Use perf stat for hardware counters (default) vtune_enabled: false # Set to true to enable VTune microarchitecture analysis vtune_bin: "/opt/intel/oneapi/vtune/latest/bin64/vtune" # Build build_command: "kernel-builder build --release" install_command: "pip install dist/*.whl --force-reinstall --no-deps" -
cpu_profiler.py 11.5 KB
#!/usr/bin/env python3 """ Profile a CPU kernel using perf stat hardware counters. Collects IPC, cache misses, branch mispredictions and maps bottlenecks to optimization patterns. Use when speedup plateaus or you need guidance on which optimization to try next. Usage: python scripts/cpu_profiler.py --kernel-package my_kernel --op my_kernel.forward python scripts/cpu_profiler.py --kernel-package my_kernel --op my_kernel.forward --warmup 10 --iters 50 """ import argparse import os import re import subprocess import sys import tempfile from functools import cache from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent)) from config import load_config as _load_config _CFG = _load_config() @cache def _find_perf_binary(): """Find a working perf binary, handling kernel version mismatches.""" import glob # The distro wrapper resolves the perf build matching the running kernel; use it when it works. try: result = subprocess.run( ["perf", "stat", "echo", "test"], capture_output=True, timeout=10, check=False, ) if result.returncode == 0: return "perf" except (FileNotFoundError, subprocess.TimeoutExpired): pass # The wrapper fails after a kernel upgrade without matching linux-tools; a versioned binary still runs. perf_candidates = sorted(glob.glob("/usr/lib/linux-tools/*/perf"), reverse=True) for candidate in perf_candidates: try: result = subprocess.run( [candidate, "stat", "echo", "test"], capture_output=True, timeout=10, check=False, ) if result.returncode == 0: return candidate except (FileNotFoundError, subprocess.TimeoutExpired): continue return None def _check_perf_available(): """Check if perf stat is available.""" return _find_perf_binary() is not None _RUNNER_SOURCE = """\ import importlib import importlib.util import sys import torch op_path, baseline_path = sys.argv[1], sys.argv[2] warmup, iters = int(sys.argv[3]), int(sys.argv[4]) parts = op_path.split(".") mod = importlib.import_module(parts[0]) func = mod for attr in parts[1:]: func = getattr(func, attr) torch.set_num_threads(torch.get_num_threads()) # Real baseline inputs keep the counters representative of the shapes the kernel will serve. try: spec = importlib.util.spec_from_file_location("baseline", baseline_path) baseline = importlib.util.module_from_spec(spec) spec.loader.exec_module(baseline) inputs = baseline.get_inputs() except Exception: print("Warning: Could not load baseline for inputs. Using dummy inputs.") inputs = [torch.randn(1024, 4096, dtype=torch.bfloat16)] # Warmup keeps page faults and first-touch allocation out of the measured # counters; no_grad keeps autograd graph construction out of every iteration. with torch.no_grad(): for _ in range(warmup): func(*inputs) for _ in range(iters): func(*inputs) """ def run_perf_stat(op_path: str, warmup: int, iters: int, baseline_file: str | None = None) -> dict[str, float]: """Run perf stat and parse results.""" if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)*", op_path): print(f" --op must be a dotted Python path like my_kernel.forward, got '{op_path}'", file=sys.stderr) sys.exit(2) if iters <= 0 or warmup < 0: print( f" --iters must be positive and --warmup nonnegative, got iters={iters} warmup={warmup}", file=sys.stderr, ) sys.exit(2) with tempfile.NamedTemporaryFile(mode="w", suffix=".py", delete=False) as f: f.write(_RUNNER_SOURCE) script_path = f.name try: counters = [ "task-clock", "cycles", "instructions", "cache-references", "cache-misses", "L1-dcache-loads", "L1-dcache-load-misses", "LLC-loads", "LLC-load-misses", "branch-instructions", "branch-misses", "page-faults", ] perf_bin = _find_perf_binary() if perf_bin is None: print(" perf binary not found") return {} cmd = [ perf_bin, "stat", "-e", ",".join(counters), "--", sys.executable, script_path, op_path, baseline_file or "baseline.py", str(warmup), str(iters), ] print(f" Running: {' '.join(cmd[:6])} ...") try: result = subprocess.run( cmd, capture_output=True, text=True, timeout=300, check=False, ) except subprocess.TimeoutExpired: print(" perf stat timed out after 300s; no stats collected") return {} if result.returncode != 0: print(f" perf stat failed (exit code {result.returncode})") if result.stderr: print(f" stderr: {result.stderr[:500]}") return {} # perf stat writes its report to stderr, not stdout. stats = _parse_perf_output(result.stderr) return stats finally: os.unlink(script_path) def _parse_perf_output(output: str) -> dict: """Parse perf stat stderr output into a dict.""" stats = {} for line in output.split("\n"): line = line.strip() if not line or line.startswith("#") or "Performance counter" in line: continue # Pattern: "1,234,567 instructions" or "1,234,567 instructions:u" match = re.match(r"([\d,\.]+)\s+(\S+)", line) if match: value_str = match.group(1).replace(",", "") name = match.group(2).removesuffix(":u").removesuffix(":k") try: value = float(value_str) stats[name] = value except ValueError: pass # Pattern for ratios: "# 0.95 insn per cycle" if "insn per cycle" in line: match = re.search(r"#\s+([\d\.]+)\s+insn per cycle", line) if match: stats["ipc"] = float(match.group(1)) return stats def print_analysis(stats: dict): """Print analysis and optimization recommendations.""" print(f"\n{'=' * 70}") print("Hardware Counter Analysis") print(f"{'=' * 70}\n") ipc = stats.get("ipc", 0) if ipc == 0 and stats.get("instructions", 0) > 0 and stats.get("cycles", 0) > 0: ipc = stats["instructions"] / stats["cycles"] print(f" IPC: {ipc:.2f}") l1_loads = stats.get("L1-dcache-loads", 0) l1_misses = stats.get("L1-dcache-load-misses", 0) l1_miss_rate = (l1_misses / l1_loads * 100) if l1_loads > 0 else 0 llc_loads = stats.get("LLC-loads", 0) llc_misses = stats.get("LLC-load-misses", 0) llc_miss_rate = (llc_misses / llc_loads * 100) if llc_loads > 0 else 0 print(f" L1 miss rate: {l1_miss_rate:.1f}%") print(f" LLC miss rate: {llc_miss_rate:.1f}%") branches = stats.get("branch-instructions", 0) branch_misses = stats.get("branch-misses", 0) branch_miss_rate = (branch_misses / branches * 100) if branches > 0 else 0 print(f" Branch miss rate: {branch_miss_rate:.1f}%") print() print(f"{'=' * 70}") print("Optimization Recommendations") print(f"{'=' * 70}\n") recommendations = [] # IPC bands are a heuristic for scalar-heavy or mixed code on Xeon cores. A tight # AVX512 FMA loop retires few, wide instructions per cycle by design, so low IPC # on such a kernel is not evidence of a memory bound; read it together with the # miss rates below before acting. if ipc < 1.0: recommendations.append( ">> IPC < 1.0: memory-bound or dependency-bound if L1 or LLC miss rates are also high;\n" " expected for a pure AVX512 FMA loop with low miss rates.\n" " - Add prefetch instructions (_mm_prefetch with _MM_HINT_T0 or _MM_HINT_T1)\n" " - Reduce cache blocking tile size\n" " - Check for false sharing in OpenMP parallel regions\n" " Reference: references/memory_patterns.yaml" ) elif ipc < 2.0: recommendations.append( ">> IPC 1.0-2.0: moderate for scalar or mixed code.\n" " - Try loop unrolling (#pragma GCC unroll 4)\n" " - Ensure FMA instructions are being generated (_mm512_fmadd_ps)\n" " Reference: references/simd_optimization_patterns.yaml" ) else: recommendations.append( ">> IPC >= 2.0: the core is not stalling on this path.\n" " - Focus on algorithmic improvements rather than micro-optimization" ) if l1_miss_rate > 10: recommendations.append( f">> L1 miss rate = {l1_miss_rate:.1f}%: High — cache blocking too large.\n" " - Reduce tile size to fit working set in L1 (48KB per core)\n" " - Add L1 prefetch: _mm_prefetch(ptr, _MM_HINT_T0)\n" " Reference: references/memory_patterns.yaml" ) if llc_miss_rate > 20: recommendations.append( f">> LLC miss rate = {llc_miss_rate:.1f}%: High — working set exceeds L3.\n" " - Add cache blocking to keep working set within L2 budget (1MB, use 50%)\n" " - Consider streaming stores for write-only data\n" " Reference: references/memory_patterns.yaml" ) if branch_miss_rate > 5: recommendations.append( f">> Branch miss rate = {branch_miss_rate:.1f}%: Consider branchless patterns.\n" " - Use SIMD masking instead of scalar branches\n" " - Use __builtin_expect for predictable branches" ) if not recommendations: recommendations.append( ">> All counters look healthy. Focus on algorithmic improvements." ) for rec in recommendations: print(f" {rec}\n") def main(): parser = argparse.ArgumentParser(description="Profile CPU kernel with perf stat") parser.add_argument("--kernel-package", required=True, help="Kernel package name") parser.add_argument("--op", required=True, help="Kernel function path (e.g., my_kernel.forward)") parser.add_argument("--baseline", default=None, help="Baseline file for getting inputs") parser.add_argument("--warmup", type=int, default=10, help="Warmup iterations") parser.add_argument("--iters", type=int, default=100, help="Profiled iterations") args = parser.parse_args() print(f"\n{'=' * 70}") print("CPU Kernel Profiler") print(f"{'=' * 70}") print(f"Kernel: {args.op}") print(f"Warmup: {args.warmup}, Iters: {args.iters}") if not _CFG.get("perf_stat_enabled", True): print("\n perf_stat_enabled=false in config.yaml. Skipping.") sys.exit(0) if not _check_perf_available(): print("\n 'perf' not found. Install linux-tools-common or run with perf_stat_enabled=false.") sys.exit(1) stats = run_perf_stat(args.op, args.warmup, args.iters, args.baseline) if stats and "instructions" in stats and "cycles" in stats: print_analysis(stats) else: print( "\n No usable hardware counters collected." " Check perf permissions (try: echo -1 > /proc/sys/kernel/perf_event_paranoid)" " and that the CPU supports the requested events (a VM without a PMU reports" " every counter as <not supported>)." ) if __name__ == "__main__": main() -
trial_manager.py 21 KB
#!/usr/bin/env python3 """Trial Tree State Manager for iterative CPU kernel optimization. Manages a tree of optimization trials for each kernel, tracking parent-child relationships, strategies, correctness, and speedup results. Supports branching back to the best ancestor when a trial regresses. Usage: python scripts/trial_manager.py init <kernel_name> <baseline_file> python scripts/trial_manager.py save <kernel_name> <trial_dir> --parent <parent_id> --strategy "description" python scripts/trial_manager.py result <kernel_name> <trial_id> --correctness <pass|fail> --speedup <float> --baseline_us <float> --kernel_us <float> python scripts/trial_manager.py status <kernel_name> python scripts/trial_manager.py best <kernel_name> python scripts/trial_manager.py baseline-us <kernel_name> python scripts/trial_manager.py finalize <kernel_name> <output_dir> """ import argparse import json import os import re import shutil import sys import tempfile from pathlib import Path TRIALS_DIR = os.path.join(os.getcwd(), "trials") OUTPUT_DIR = os.path.join(os.getcwd(), "output") def _checked_name(kernel_name): """A kernel name is one directory component inside trials/.""" trials_real = Path(TRIALS_DIR).resolve() candidate = (Path(TRIALS_DIR) / kernel_name).resolve() if ( kernel_name in ("", ".", "..") or os.sep in kernel_name or (os.altsep and os.altsep in kernel_name) or not candidate.is_relative_to(trials_real) ): print( f"Error: Kernel name '{kernel_name}' must be a single directory name under trials/.", file=sys.stderr, ) sys.exit(1) def _state_path(kernel_name): _checked_name(kernel_name) return os.path.join(TRIALS_DIR, kernel_name, "state.json") def _trial_dir(kernel_name): _checked_name(kernel_name) return os.path.join(TRIALS_DIR, kernel_name) def _overlaps_trial_store(source): """True when copying source would copy the trial store into itself.""" src = Path(source).resolve() store = Path(TRIALS_DIR).resolve() return src == store or src in store.parents or store in src.parents def _escaping_symlinks(source): """Symlinks under source whose text stops being valid once the tree is copied elsewhere. Only a relative link that stays inside source by path arithmetic alone survives relocation; an absolute link keeps pointing at the original tree, and a relative link that walks out of source (even if it resolves back in) depends on it. """ root = os.path.abspath(source) real_root = os.path.realpath(source) def _stays_inside(dirpath, text): """True if the symlink text stays inside source by path arithmetic alone.""" if os.path.isabs(text): return False # Resolution can leave through a nested symlink even when the text alone does not. resolved = os.path.realpath(os.path.join(dirpath, text)) if os.path.commonpath([real_root, resolved]) != real_root: return False rel = os.path.relpath(os.path.abspath(dirpath), root) depth = 0 if rel == os.curdir else len(rel.split(os.sep)) for part in text.split("/"): if part == "..": depth -= 1 if depth < 0: return False elif part not in ("", "."): depth += 1 return True escaping = [] for dirpath, dirnames, filenames in os.walk(source): for name in dirnames + filenames: link = os.path.join(dirpath, name) if os.path.islink(link) and not _stays_inside(dirpath, os.readlink(link)): escaping.append(link) return escaping def _validate_state(state, kernel_name): """Trial ids are t<number> and a trial's dir is its own id; nothing else is valid state.""" for tid, trial in state.get("trials", {}).items(): if not re.fullmatch(r"t\d+", tid): print( f"Error: Corrupt trial state for '{kernel_name}':" f" trial id '{tid}' is not a t<number> id.", file=sys.stderr, ) sys.exit(1) if trial.get("dir") != tid: print( f"Error: Corrupt trial state for '{kernel_name}':" f" trial '{tid}' dir '{trial.get('dir')}' does not match its id.", file=sys.stderr, ) sys.exit(1) for field in ("speedup", "baseline_us", "kernel_us"): v = trial.get(field) if v is not None and not isinstance(v, (int, float)): print( f"Error: Corrupt trial state for '{kernel_name}':" f" trial '{tid}' field '{field}' is not a number.", file=sys.stderr, ) sys.exit(1) baseline_us = state.get("baseline_us") if baseline_us is not None and ( not isinstance(baseline_us, list) or not all(isinstance(v, (int, float)) for v in baseline_us) ): print( f"Error: Corrupt trial state for '{kernel_name}':" " 'baseline_us' is not a list of numbers.", file=sys.stderr, ) sys.exit(1) best = state.get("best_trial") if best is not None and (not isinstance(best, str) or best not in state.get("trials", {})): print( f"Error: Corrupt trial state for '{kernel_name}':" f" best_trial '{best}' is not a known trial.", file=sys.stderr, ) sys.exit(1) def _load_state(kernel_name): path = _state_path(kernel_name) if not os.path.exists(path): print(f"Error: No trial tree found for '{kernel_name}'. Run 'init' first.", file=sys.stderr) sys.exit(1) with open(path) as f: state = json.load(f) _validate_state(state, kernel_name) return state def _save_state(kernel_name, state): path = _state_path(kernel_name) with open(path, "w") as f: json.dump(state, f, indent=2) # ============================================================================ # Commands # ============================================================================ def cmd_init(args): """Initialize a new trial tree for a kernel.""" kernel_name = args.kernel_name baseline_file = args.baseline_file trial_dir = _trial_dir(kernel_name) if os.path.exists(_state_path(kernel_name)): print( f"Warning: Trial tree for '{kernel_name}' already exists. " f"Use a different name or delete trials/{kernel_name}/." ) else: os.makedirs(trial_dir, exist_ok=True) state = { "kernel_name": kernel_name, "baseline_file": baseline_file, "trials": {}, "best_trial": None, "next_id": 0, "baseline_us": None, } _save_state(kernel_name, state) print(f"Initialized trial tree for '{kernel_name}' in trials/{kernel_name}/") print(f" Baseline: {baseline_file}") def cmd_save(args): """Save a trial by copying kernel files into the trial directory.""" kernel_name = args.kernel_name trial_source = args.trial_source parent = args.parent strategy = args.strategy or "" if not os.path.exists(trial_source): print(f"Error: Trial source '{trial_source}' not found.", file=sys.stderr) sys.exit(1) if _overlaps_trial_store(trial_source): print( f"Error: Trial source '{trial_source}' overlaps 'trials/'." " Save a source outside the trial store.", file=sys.stderr, ) sys.exit(1) if os.path.isdir(trial_source): escaping = _escaping_symlinks(trial_source) if escaping: print( f"Error: Trial source '{trial_source}' has symlinks that point outside it," " so they would dangle once copied: " + ", ".join(escaping), file=sys.stderr, ) sys.exit(1) state = _load_state(kernel_name) if parent is not None and parent not in state["trials"]: if state["next_id"] == 0: print( f"Warning: Ignoring --parent '{parent}' for first trial.", file=sys.stderr, ) parent = None else: print( f"Error: Parent trial '{parent}' not found. Available: {list(state['trials'].keys())}", file=sys.stderr, ) sys.exit(1) trial_id = f"t{state['next_id']}" state["next_id"] += 1 dest = os.path.join(_trial_dir(kernel_name), trial_id) if os.path.lexists(dest): print( f"Error: Trial directory '{dest}' already exists but is not in state.json." " Remove it or repair the trial store before saving.", file=sys.stderr, ) sys.exit(1) # Copy into a fresh staging directory and rename it into place, so a failed # copy never leaves a partial trial and cleanup only removes what this save # created. symlinks=True copies links as links, so a nested symlink back into # the trial store can never make the copy recurse into itself. staging = None try: staging = tempfile.mkdtemp(prefix=f".{trial_id}-", dir=_trial_dir(kernel_name)) staged = os.path.join(staging, trial_id) if os.path.isdir(trial_source): shutil.copytree(trial_source, staged, symlinks=True) else: os.makedirs(staged) shutil.copy2(trial_source, staged) os.rename(staged, dest) except (OSError, RecursionError) as e: if staging is not None: shutil.rmtree(staging, ignore_errors=True) print(f"Error: Failed to copy trial source into '{dest}': {e}", file=sys.stderr) sys.exit(1) shutil.rmtree(staging, ignore_errors=True) state["trials"][trial_id] = { "parent": parent, "dir": trial_id, "strategy": strategy, "correctness": None, "speedup": None, "baseline_us": None, "kernel_us": None, "status": "saved", } _save_state(kernel_name, state) print(f"Saved trial {trial_id}: {strategy}") print(f" Parent: {parent or 'root'}") print(f" Dir: trials/{kernel_name}/{trial_id}/") def cmd_result(args): """Record results for a trial.""" kernel_name = args.kernel_name trial_id = args.trial_id state = _load_state(kernel_name) if trial_id not in state["trials"]: print(f"Error: Trial '{trial_id}' not found. Available: {list(state['trials'].keys())}", file=sys.stderr) sys.exit(1) trial = state["trials"][trial_id] if args.correctness: trial["correctness"] = args.correctness if args.speedup is not None: trial["speedup"] = args.speedup if args.baseline_us is not None: trial["baseline_us"] = args.baseline_us if args.kernel_us is not None: trial["kernel_us"] = args.kernel_us # The baseline time is per kernel, not per trial; caching it lets later trials skip re-timing it. if args.baseline_us is not None and state.get("baseline_us") is None: state["baseline_us"] = [args.baseline_us] if trial["correctness"] == "fail": trial["status"] = "failed" elif trial["correctness"] == "pass" and trial["speedup"] is not None: trial["status"] = "completed" else: trial["status"] = "partial" best_speedup = -1.0 best_id = None for tid, t in state["trials"].items(): if ( t.get("correctness") == "pass" and t.get("speedup") is not None and t["speedup"] > best_speedup ): best_speedup = t["speedup"] best_id = tid state["best_trial"] = best_id _save_state(kernel_name, state) status_icon = {"completed": "+", "failed": "X", "partial": "~", "saved": "?"} icon = status_icon.get(trial["status"], "?") runtime_str = "" if trial.get("baseline_us") is not None and trial.get("kernel_us") is not None: runtime_str = f", baseline={trial['baseline_us']:.2f}us, kernel={trial['kernel_us']:.2f}us" print( f"[{icon}] {trial_id}: correctness={trial['correctness']}, speedup={trial['speedup']}{runtime_str}" ) if state["best_trial"]: best = state["trials"][state["best_trial"]] best_runtime = "" if best.get("baseline_us") is not None and best.get("kernel_us") is not None: best_runtime = ( f", baseline={best['baseline_us']:.2f}us, kernel={best['kernel_us']:.2f}us" ) print(f" Best trial: {state['best_trial']} ({best['speedup']}x{best_runtime})") def cmd_status(args): """Show trial tree status as ASCII tree.""" kernel_name = args.kernel_name state = _load_state(kernel_name) print(f"Trial tree: {state['kernel_name']}") print(f" Baseline: {state['baseline_file']}") print(f" Best: {state['best_trial'] or 'none'}") print(f" Trials: {len(state['trials'])}") print() if not state["trials"]: print(" (no trials yet)") return children = {} roots = [] for tid, t in state["trials"].items(): parent = t["parent"] if parent is None: roots.append(tid) else: children.setdefault(parent, []).append(tid) def sort_key(tid): return int(tid[1:]) roots.sort(key=sort_key) for kids in children.values(): kids.sort(key=sort_key) def print_node(tid, prefix="", is_last=True): trial = state["trials"][tid] connector = "└── " if is_last else "├── " is_best = tid == state["best_trial"] status_icon = {"completed": "+", "failed": "X", "partial": "~", "saved": "?"} icon = status_icon.get(trial["status"], "?") speedup_str = f"{trial['speedup']:.2f}x" if trial["speedup"] is not None else "---" runtime_str = "" if trial.get("baseline_us") is not None and trial.get("kernel_us") is not None: runtime_str = f" (bl={trial['baseline_us']:.0f}us, kr={trial['kernel_us']:.0f}us)" best_marker = " <<<< BEST" if is_best else "" strategy_short = trial["strategy"][:60] if trial["strategy"] else "" print( f"{prefix}{connector}[{icon}] {tid}: {speedup_str}{runtime_str} | {strategy_short}{best_marker}" ) child_prefix = prefix + (" " if is_last else "│ ") kids = children.get(tid, []) for i, child in enumerate(kids): print_node(child, child_prefix, i == len(kids) - 1) for i, root in enumerate(roots): print_node(root, " ", i == len(roots) - 1) def cmd_best(args): """Get the best trial info.""" kernel_name = args.kernel_name state = _load_state(kernel_name) if state["best_trial"] is None: print("No correct trials yet.") sys.exit(1) best_id = state["best_trial"] best = state["trials"][best_id] best_dir = os.path.join(_trial_dir(kernel_name), best["dir"]) print(f"best_trial: {best_id}") print(f"speedup: {best['speedup']}") if best.get("baseline_us") is not None: print(f"baseline_us: {best['baseline_us']}") if best.get("kernel_us") is not None: print(f"kernel_us: {best['kernel_us']}") print(f"strategy: {best['strategy']}") print(f"dir: {best_dir}") print(f"parent: {best['parent'] or 'root'}") def cmd_baseline_us(args): """Print cached baseline time(s).""" kernel_name = args.kernel_name state = _load_state(kernel_name) baseline_us = state.get("baseline_us") if baseline_us is None: print( "No baseline_us cached yet. Run benchmark and record result for t0 first.", file=sys.stderr, ) sys.exit(1) print(",".join(f"{v:.2f}" for v in baseline_us)) def cmd_finalize(args): """Copy the best correct trial to the output path.""" kernel_name = args.kernel_name output_path = args.output_path state = _load_state(kernel_name) if state["best_trial"] is None: print("Error: No correct trials to finalize.", file=sys.stderr) sys.exit(1) best_id = state["best_trial"] best = state["trials"][best_id] src = os.path.join(_trial_dir(kernel_name), best["dir"]) if os.path.islink(src): print( f"Error: Trial directory '{src}' is a symlink; a trial must be a real directory" " inside the trial store.", file=sys.stderr, ) sys.exit(1) # A bare name is a label, not a path; keep every finalized kernel under one output root. if os.path.dirname(output_path) == "": os.makedirs(OUTPUT_DIR, exist_ok=True) output_path = os.path.join(OUTPUT_DIR, output_path) # Stage beside the destination, then swap, so an existing output never # keeps files the chosen trial no longer has. Resolve a symlinked # destination first so the swap replaces the target, not the link. dest = Path(output_path) staging = None backup = None try: if dest.is_symlink(): dest = dest.resolve() dest.parent.mkdir(parents=True, exist_ok=True) staging = Path(tempfile.mkdtemp(prefix=".finalize-", dir=dest.parent)) staged = staging / dest.name backup = Path(f"{staged}.old") if Path(src).is_dir(): shutil.copytree(src, staged, symlinks=True) else: shutil.copy2(src, staged) if dest.exists() or dest.is_symlink(): os.rename(dest, backup) os.rename(staged, dest) except OSError as e: if backup is not None and (backup.exists() or backup.is_symlink()) and not (dest.exists() or dest.is_symlink()): os.rename(backup, dest) if staging is not None: shutil.rmtree(staging, ignore_errors=True) print(f"Error: Failed to finalize into '{output_path}': {e}", file=sys.stderr) sys.exit(1) try: if backup is not None and (backup.exists() or backup.is_symlink()): if backup.is_symlink() or not backup.is_dir(): backup.unlink() else: shutil.rmtree(backup, ignore_errors=True) except OSError as e: print(f"Warning: finalized into '{output_path}' but could not remove the displaced output '{backup}': {e}", file=sys.stderr) shutil.rmtree(staging, ignore_errors=True) runtime_str = "" if best.get("baseline_us") is not None and best.get("kernel_us") is not None: runtime_str = f", baseline={best['baseline_us']:.2f}us, kernel={best['kernel_us']:.2f}us" print(f"Finalized {best_id} ({best['speedup']}x{runtime_str}) -> {output_path}") print(f" Strategy: {best['strategy']}") # ============================================================================ # CLI # ============================================================================ def main(): parser = argparse.ArgumentParser(description="Trial Tree State Manager (CPU Kernels)") subparsers = parser.add_subparsers(dest="command", required=True) p_init = subparsers.add_parser("init", help="Initialize trial tree") p_init.add_argument("kernel_name", help="Kernel identifier (e.g., rmsnorm)") p_init.add_argument("baseline_file", help="Path to PyTorch baseline file") p_save = subparsers.add_parser("save", help="Save a trial") p_save.add_argument("kernel_name", help="Kernel identifier") p_save.add_argument("trial_source", help="Path to the trial kernel dir or file") p_save.add_argument("--parent", default=None, help="Parent trial ID (e.g., t0)") p_save.add_argument("--strategy", default="", help="Description of optimization strategy") p_result = subparsers.add_parser("result", help="Record trial results") p_result.add_argument("kernel_name", help="Kernel identifier") p_result.add_argument("trial_id", help="Trial ID (e.g., t0)") p_result.add_argument("--correctness", choices=["pass", "fail"], help="Correctness result") p_result.add_argument("--speedup", type=float, help="Speedup over baseline") p_result.add_argument("--baseline_us", type=float, help="Baseline runtime in microseconds") p_result.add_argument("--kernel_us", type=float, help="Kernel runtime in microseconds") p_status = subparsers.add_parser("status", help="Show trial tree status") p_status.add_argument("kernel_name", help="Kernel identifier") p_best = subparsers.add_parser("best", help="Get best trial info") p_best.add_argument("kernel_name", help="Kernel identifier") p_baseline_us = subparsers.add_parser("baseline-us", help="Print cached baseline time(s)") p_baseline_us.add_argument("kernel_name", help="Kernel identifier") p_finalize = subparsers.add_parser("finalize", help="Copy best trial to output") p_finalize.add_argument("kernel_name", help="Kernel identifier") p_finalize.add_argument("output_path", help="Output path (bare name defaults to output/)") args = parser.parse_args() commands = { "init": cmd_init, "save": cmd_save, "result": cmd_result, "status": cmd_status, "best": cmd_best, "baseline-us": cmd_baseline_us, "finalize": cmd_finalize, } commands[args.command](args) if __name__ == "__main__": main() -
validate_cpu_kernel.py 13.1 KB
#!/usr/bin/env python3 """ Validate C++ CPU kernel for common issues. Checks build.toml, C++ source files, and kernel structure for correctness and common pitfalls. Usage: python scripts/validate_cpu_kernel.py <kernel_dir> python scripts/validate_cpu_kernel.py . """ import re import sys from pathlib import Path if sys.version_info < (3, 11): sys.exit("validate_cpu_kernel.py needs Python 3.11+ (tomllib in the standard library)") import tomllib # noqa: E402 class ValidationError: def __init__(self, level: str, message: str, file: str | None = None, line_num: int | None = None): self.level = level # 'ERROR', 'WARNING', 'INFO' self.message = message self.file = file self.line_num = line_num def __str__(self): prefix = {"ERROR": "X", "WARNING": "!", "INFO": "i"}[self.level] loc = "" if self.file: loc += f" [{self.file}" if self.line_num: loc += f":{self.line_num}" loc += "]" return f"[{prefix}] {self.level}: {self.message}{loc}" def validate_build_toml(kernel_dir: Path) -> list[ValidationError]: """Validate build.toml configuration against the parsed kernel sections.""" errors = [] build_toml = kernel_dir / "build.toml" if not build_toml.exists(): errors.append(ValidationError("ERROR", "build.toml not found")) return errors try: with open(build_toml, "rb") as f: data = tomllib.load(f) except tomllib.TOMLDecodeError as e: errors.append(ValidationError("ERROR", f"build.toml is not valid TOML: {e}", "build.toml")) return errors kernel_table = data.get("kernel", {}) if not isinstance(kernel_table, dict): errors.append(ValidationError( "ERROR", f"build.toml 'kernel' must be a table of [kernel.<name>] sections, got {type(kernel_table).__name__}", "build.toml", )) return errors sections = { f"kernel.{name}": body for name, body in kernel_table.items() if isinstance(body, dict) } cpu_sections = {name: body for name, body in sections.items() if body.get("backend") == "cpu"} if not cpu_sections: errors.append(ValidationError("ERROR", "No CPU backend sections found in build.toml", "build.toml")) for name, body in cpu_sections.items(): # kernel-builder does not add the kernel directory to the include path; # without `include` headers fail to resolve. include = body.get("include", []) if isinstance(include, str): include = [include] if not include: errors.append(ValidationError( "WARNING", f"Section [{name}] missing 'include' directive for header resolution", "build.toml", )) def _flags(body: dict) -> list[str]: """Extract and normalize compiler flags from a kernel section body.""" flags = body.get("cxx-flags", body.get("flags", [])) if isinstance(flags, list): return [str(f) for f in flags] return str(flags).split() # dq/bw/vbmi are needed only by GEMM byte-shuffle paths; requiring them elsewhere is noise. gemm_indicators = ["gemm", "gptq", "quantiz", "bnb", "bitsandbytes", "megablocks", "moe"] # Each [kernel.*] section is its own translation unit, so a flag in one # tier never reaches another; check every AVX512 section on its own. for name, body in sections.items(): flags = _flags(body) if "-mavx512f" not in flags: continue is_gemm_kernel = any(ind in name.lower() for ind in gemm_indicators) for flag in ("-mavx512bf16", "-mavx512vl"): if flag not in flags: errors.append(ValidationError( "WARNING", f"AVX512 section [{name}] missing core flag: {flag}", "build.toml", )) if is_gemm_kernel: for flag in ("-mavx512dq", "-mavx512bw", "-mavx512vbmi"): if flag not in flags: errors.append(ValidationError( "INFO", f"GEMM kernel section [{name}] may benefit from flag: {flag}", "build.toml", )) if "-fopenmp" not in flags: errors.append(ValidationError( "WARNING", f"AVX512 section [{name}] missing -fopenmp flag", "build.toml", )) return errors def validate_cpp_file(filepath: Path) -> list[ValidationError]: """Validate a C++ source file for common CPU kernel issues.""" errors = [] fname = filepath.name with open(filepath) as f: lines = f.readlines() source = "".join(lines) # Aligned loads fault on tensors whose storage offset is not 64-byte aligned, which PyTorch does not guarantee. for i, line in enumerate(lines): if "_mm512_load_" in line and "_mm512_loadu_" not in line: stripped = line.lstrip() if not stripped.startswith("//") and not stripped.startswith("*"): errors.append(ValidationError( "WARNING", "Using aligned load (_mm512_load_*). Prefer _mm512_loadu_* for safety.", fname, i + 1, )) if "_mm256_load_" in line and "_mm256_loadu_" not in line: stripped = line.lstrip() if not stripped.startswith("//") and not stripped.startswith("*"): errors.append(ValidationError( "WARNING", "Using aligned load (_mm256_load_*). Prefer _mm256_loadu_* for safety.", fname, i + 1, )) # A hidden size not divisible by the vector width leaves a tail; missing tail code reads past the row. if "avx512" in fname.lower(): has_vector_ops = "_mm512_" in source has_tail = "remainder" in source.lower() or "tail" in source.lower() or "mask" in source.lower() if has_vector_ops and not has_tail: errors.append(ValidationError( "INFO", "No tail/remainder handling detected. Ensure hidden_size is always divisible by vector width, " "or add masked/scalar fallback.", fname, )) if fname.endswith("_cpu.cpp") and "avx" not in fname: has_cpu_features = "cpu_features.hpp" in source or "cpu_features" in source if not has_cpu_features: # The dispatcher may reach cpu_features.hpp through its own header; a direct-include check alone false-positives. included_headers = re.findall(r'#include\s+"([^"]+\.hpp)"', source) parent_dir = filepath.parent for header in included_headers: header_path = parent_dir / header if header_path.exists(): with open(header_path) as hf: if "cpu_features.hpp" in hf.read(): has_cpu_features = True break if not has_cpu_features: errors.append(ValidationError( "WARNING", "Dispatcher file should include cpu_features.hpp for runtime feature detection.", fname, )) # Mixing ISA tiers in one translation unit forces one flag set on both, so the fallback tier inherits AVX512 codegen. has_avx2 = "_mm256_" in source has_avx512 = "_mm512_" in source if has_avx2 and has_avx512: # AVX512 files use _mm256 for partial-width loads (zero points); only an AVX2-only file mixing in _mm512 is a defect. is_avx512_file = "avx512" in fname.lower() if not is_avx512_file: errors.append(ValidationError( "WARNING", "Mixing AVX2 (_mm256_*) and AVX512 (_mm512_*) intrinsics in same file. " "Each ISA tier should be in a separate translation unit.", fname, )) if "#pragma omp" in source and "#include <omp.h>" not in source: # Pragmas compile without the header; only omp_* calls need it, so this is not an error. pass # float64 halves SIMD lane count; hot-path doubles are almost always accidental. for i, line in enumerate(lines): if "double " in line and "epsilon" not in line.lower() and "eps" not in line.lower(): stripped = line.lstrip() if not stripped.startswith("//") and not stripped.startswith("*"): errors.append(ValidationError( "INFO", "double (float64) detected. Consider float/bf16 for better SIMD throughput.", fname, i + 1, )) return errors def validate_torch_binding(kernel_dir: Path) -> list[ValidationError]: """Validate torch_binding.cpp (located at torch-ext/torch_binding.cpp).""" errors = [] binding = kernel_dir / "torch-ext" / "torch_binding.cpp" if not binding.exists(): binding = kernel_dir / "torch_binding.cpp" if not binding.exists(): errors.append(ValidationError("ERROR", "torch_binding.cpp not found (expected at torch-ext/torch_binding.cpp)")) return errors # kernel-builder expects torch-ext/; a root-level binding builds locally but breaks the Hub layout. if binding.parent.name != "torch-ext": errors.append(ValidationError( "WARNING", "torch_binding.cpp should be in torch-ext/ directory (torch-ext/torch_binding.cpp)", str(binding.relative_to(kernel_dir)), )) with open(binding) as f: source = f.read() if "registration.h" not in source: errors.append(ValidationError( "ERROR", "torch_binding.cpp should include registration.h", "torch_binding.cpp", )) if "TORCH_LIBRARY_EXPAND" not in source: errors.append(ValidationError( "ERROR", "torch_binding.cpp should use TORCH_LIBRARY_EXPAND macro for op registration", "torch_binding.cpp", )) if "REGISTER_EXTENSION" not in source: errors.append(ValidationError( "WARNING", "torch_binding.cpp should use REGISTER_EXTENSION macro", "torch_binding.cpp", )) return errors def validate_kernel_structure(kernel_dir: Path) -> list[ValidationError]: """Validate overall kernel directory structure.""" errors = [] cpu_dirs = [d for d in kernel_dir.iterdir() if d.is_dir() and "cpu" in d.name.lower()] if not cpu_dirs: errors.append(ValidationError( "WARNING", "No *_cpu/ subdirectory found. CPU kernel code should be in <kernel>_cpu/.", )) return errors for cpu_dir in cpu_dirs: if not (cpu_dir / "cpu_features.hpp").exists(): errors.append(ValidationError( "ERROR", f"Missing cpu_features.hpp in {cpu_dir.name}/", )) avx512_files = list(cpu_dir.glob("*avx512*")) if not avx512_files: errors.append(ValidationError( "WARNING", f"No AVX512 implementation found in {cpu_dir.name}/", )) return errors def validate_kernel(kernel_dir: Path) -> list[ValidationError]: """Run all validation checks.""" errors = [] errors.extend(validate_build_toml(kernel_dir)) errors.extend(validate_torch_binding(kernel_dir)) errors.extend(validate_kernel_structure(kernel_dir)) for ext in ("*.cpp", "*.hpp"): for cpp_file in kernel_dir.rglob(ext): errors.extend(validate_cpp_file(cpp_file)) return errors def print_results(errors: list[ValidationError], kernel_dir: Path) -> int: """Pretty print validation results.""" print(f"\n{'=' * 70}") print(f"CPU Kernel Validation: {kernel_dir}") print(f"{'=' * 70}\n") error_list = [e for e in errors if e.level == "ERROR"] warning_list = [e for e in errors if e.level == "WARNING"] info_list = [e for e in errors if e.level == "INFO"] if error_list: print("ERRORS (must fix):") for err in error_list: print(f" {err}") print() if warning_list: print("WARNINGS (should review):") for err in warning_list: print(f" {err}") print() if info_list: print("INFO:") for err in info_list: print(f" {err}") print() if error_list: print(f"Status: FAILED ({len(error_list)} errors)") return 1 elif warning_list: print(f"Status: PASSED with warnings ({len(warning_list)} warnings)") return 0 else: print("Status: PASSED") return 0 def main(): if len(sys.argv) != 2: print("Usage: python scripts/validate_cpu_kernel.py <kernel_dir>") sys.exit(1) kernel_dir = Path(sys.argv[1]).resolve() if not kernel_dir.is_dir(): print(f"Error: Kernel directory not found: {kernel_dir}") sys.exit(1) errors = validate_kernel(kernel_dir) exit_code = print_results(errors, kernel_dir) sys.exit(exit_code) if __name__ == "__main__": main()
-
-
SKILL.md 10.7 KB
--- name: cpu-kernel-authoring description: 'Use when writing, optimizing, or benchmarking a C++ CPU kernel with AVX2 or AVX512 intrinsics for the Hugging Face kernels ecosystem. Not for CUDA kernels: use cuda.' disable-model-invocation: true --- # CPU kernel authoring ## Contract | Field | Bound contract | |---|---| | Trigger | A C++ CPU kernel for the Hugging Face kernels ecosystem must be written, optimized, or benchmarked with AVX2 or AVX512 intrinsics against a PyTorch baseline. | | Authority | Reversible local. Writes C++ kernel sources, `build.toml`, and `torch_binding.cpp` under the kernel directory, a wheel under `dist/`, the installed kernel package in the active Python environment, and trial state under `trials/<kernel_name>/` and `output/`. Rollback is version control for the sources, `pip uninstall <package>` for the package, and removal of `dist/`, `trials/<kernel_name>/`, and `output/`. No remote mutation. | | Side effect | Kernel sources and build files change; a wheel is built and installed; trial directories and result records accumulate. | | Done | The kernel passes the correctness check in `scripts/benchmark_cpu.py`, every trial up to `max_trials` has run or the speedup exceeded `early_stop_speedup`, and the best trial is finalized into `output/` with its final measurement; or a failure class from the table below is reported with the recovery step taken. | ## Inputs - Kernel name (required): the trial-tree label, for example `my_rmsnorm`. Used only by `trial_manager.py`, which accepts it as a single directory name under `trials/`, never a path. - Baseline file (required): a `baseline.py` that defines `get_inputs()` and either `get_reference_output()` or a `Model` class (with optional `get_init_inputs()`). It is the ground truth for correctness and the speed reference. - Operation name (required): the plain name `analyze_op.py --op` looks up, for example `rms_norm`. - Input shapes (required): comma-separated shape strings for `analyze_op.py --shapes`, for example `"1024x4096,2048x8192"`. - Package and function path (required from step 5): the installed package name, for example `my_kernel`, and its callable as `package.function`, for example `my_kernel.rms_norm`. `benchmark_cpu.py` and `cpu_profiler.py` take this path as their `--op`; it is not the operation name above. - Toolchain (required): Python 3.11+ (`validate_cpu_kernel.py` parses `build.toml` with the standard-library `tomllib`), `kernel-builder`, `pip`, PyYAML (imported by `scripts/config.py`), `numactl` (used by the pinned benchmark in step 8), a C++ compiler with AVX512 support, and PyTorch. `perf` is required only when `perf_stat_enabled` is true. The work has two phases. The correctness phase builds the tiers in order (generic ATen fallback, optional AVX2, AVX512) and each tier must pass correctness before the next starts. The performance phase iterates on the AVX512 tier through the trial tree until `max_trials` is exhausted or `early_stop_speedup` is exceeded. ## Procedure 1. Read `scripts/config.yaml` and note `max_trials`, `early_stop_speedup`, `perf_stat_enabled`, `vtune_enabled`, `build_command`, and `install_command`. Use those two commands wherever this procedure builds or installs. Done when: every value is known. 2. Run `python scripts/analyze_op.py --op <op_name> --shapes <shapes>` and read the compute and memory characteristics and the suggested SIMD strategy. Read `references/workflow_details.md` for the analysis and design steps. Done when: the kernel type is fixed as element-wise, reduction, GEMM, or attention. 3. Run `python scripts/trial_manager.py init <kernel_name> <baseline_file>`. Done when: `trials/<kernel_name>/` exists and records the baseline. 4. Write the generic tier: `<kernel>_cpu/cpu_features.hpp` in the kernel's own namespace, the dispatcher `<kernel>_cpu/<kernel>_cpu.cpp` with an ATen-only fallback, the bridge `<kernel>_cpu/<kernel>_cpu_torch.cpp`, `torch-ext/torch_binding.cpp` using the `registration.h` macros, and `build.toml` with one `[kernel.*]` section per tier and `include = ["<kernel>_cpu"]` in every section. Read `references/runtime_dispatch.yaml`, `references/build_system.md`, `references/implementation_reference.md`, and `references/correctness.yaml` while writing. Run `python scripts/validate_cpu_kernel.py <kernel_dir>`. Done when: validation reports no error. 5. Build and install with the configured commands, by default `kernel-builder build --release` then `pip install dist/*.whl --force-reinstall --no-deps`. Done when: `python -c "import <package>"` succeeds. 6. Run `python scripts/benchmark_cpu.py <baseline_file> --kernel-package <package> --op <package>.<function>`. The correctness check walks tuples, lists, and dicts element-wise, requires equal structure, dtype, and shape, and compares each tensor leaf in its own dtype: half an ulp relative for bf16 and fp16, `atol=1e-6, rtol=1e-5` for fp32, `atol=1e-12, rtol=1e-9` for fp64, exact for integer and bool. Widen with `--atol` and `--rtol` only when the kernel's accumulation order legitimately differs from the reference, and record the reason in the trial's `--strategy`. Done when: correctness passes and the baseline and kernel times are recorded; on failure, go to the failure table. 7. Add the AVX512 tier in its own translation unit `<kernel>_cpu/<kernel>_avx512.cpp` with its own `cxx-flags` section (`-mavx512f -mavx512bf16 -mavx512vl` for element-wise kernels; GEMM kernels add `-mavx512dq -mavx512bw -mavx512vbmi -mamx-tile -mamx-bf16 -mamx-int8`), and `-fopenmp` in every SIMD section. Add an AVX2 tier only when it gives an element-wise kernel a measurable benefit; GEMM kernels dispatch AVX512 to fallback. Repeat steps 4 to 6, then run `python scripts/trial_manager.py save <kernel_name> <kernel_dir> --strategy "<description>"` and record the numbers with `python scripts/trial_manager.py result <kernel_name> <trial_id> --correctness pass --speedup <x> --baseline_us <us> --kernel_us <us>`. Done when: the AVX512 tier is correct and trial t0 is recorded. This ends the correctness phase. 8. Pin the benchmark to one NUMA node for every later measurement: `numactl --cpunodebind=0 --membind=0 python scripts/benchmark_cpu.py ... --baseline-us <cached>`, where the cached value comes from `python scripts/trial_manager.py baseline-us <kernel_name>`. Done when: the pinned command is the one used from here on. 9. When `perf_stat_enabled` is true, run `python scripts/cpu_profiler.py --kernel-package <package> --op <package>.<function>` once after the first benchmarked trial. Read IPC together with the L1 and LLC miss rates: a pure AVX512 FMA loop has low IPC by design, so a memory bound is claimed only when a miss rate is also high. Done when: the profile is read and the next change is chosen from `references/optimization_strategies.md`. 10. For each remaining trial up to `max_trials`: change one thing in the AVX512 tier (blocking, prefetch, unrolling, threading, or a different algorithm from `references/simd_optimization_patterns.yaml`, `references/memory_patterns.yaml`, `references/threading_patterns.yaml`, `references/dtype_optimizations.yaml`, `references/brgemm_patterns.yaml`, `references/quantized_gemm_patterns.yaml`, and `references/optimization_levels.yaml`), validate, build, benchmark, then `save` with `--parent <best_or_current_id>` and `result`. A regression branches back to the best trial; a plateau after two trials changes the algorithm, data layout, or fusion instead of sweeping the same knobs. Stop early only when the speedup exceeds `early_stop_speedup`. Done when: `max_trials` trials are recorded or the early stop fired. 11. Run `python scripts/trial_manager.py finalize <kernel_name> output/`, then re-run the pinned `benchmark_cpu.py` without `--baseline-us` for the final measurement. Read `references/huggingface-kernels-integration.md` if the kernel is to be published to the Hub. Done when: `output/` holds the best trial's sources and its final correctness and speedup are recorded. Modify only `.cpp` and `.hpp` files, `torch_binding.cpp`, and `build.toml`. Do not write new benchmark or timing scripts; `scripts/benchmark_cpu.py` is the only timing source. When a script fails, report the error rather than working around it. ## Failure and recovery | Failure class | Behavior | |---|---| | `scripts/config.yaml` missing | Stop and report. Do not assume trial counts. | | `analyze_op.py` reports no matmul, reduction, or activation for the op | The script recognizes norm, softmax, gemm, linear, matmul, attention, gelu, silu, relu, moe, and megablocks by name and classifies anything else as plain element-wise. Classify the kernel type by hand from the baseline and continue at step 3. | | `validate_cpu_kernel.py` reports an error | Fix the named file or `build.toml` section; re-run validation. A validation fix does not count as a trial. | | `kernel-builder build` fails | Read the compiler output; fix the source or the section's `cxx-flags`; rebuild. | | Correctness fails | Read the leaf path, dtype, and worst-element values in the benchmark output. A wrong second output or a dtype or shape change is a binding or dispatcher bug; a one-ulp bf16 difference on many elements is a rounding-mode or conversion bug; a tail-only difference is missing tail handling; a large scattered difference is an alignment bug. Fix on the same branch, rebuild, re-benchmark. Do not enter the performance phase with a failing kernel. | | Kernel slower than baseline on small tensors | Add a `num_tokens` threshold below which the dispatcher calls the ATen fallback; see `references/threading_patterns.yaml`. | | `perf` unavailable or `perf stat` returns no counters | Continue without profiling and report it; choose the next trial from `references/optimization_levels.yaml`. | | Speedup regressed | `save` the next trial with `--parent` set to the best trial id from `python scripts/trial_manager.py best <kernel_name>`. | | Plateau after two or more trials | Change algorithm, data layout, or fusion strategy. Do not sweep the same parameters. | | `max_trials` reached below `early_stop_speedup` | Finalize the best trial and report the speedup reached and the trial tree from `trial_manager.py status`. | ## Output - Kernel sources: `<kernel>_cpu/` with `cpu_features.hpp`, the dispatcher, the bridge, the AVX512 implementation, and any AVX2 implementation, plus `torch-ext/torch_binding.cpp` and `build.toml`. - Installed package: the wheel under `dist/` and the installed `<package>`. - Trial tree: `trials/<kernel_name>/` with each saved trial, its parent, strategy, correctness, and timing. - Correctness report: the `benchmark_cpu.py` output naming per-dtype tolerances and, on failure, the leaf path of each mismatch. - Performance report: baseline and kernel microseconds and speedup from the NUMA-pinned run. - Final kernel: `output/` holding the best trial's sources and its final measurement.
Comments (0)
Sign in to join the conversation.
Reviews (0)
No reviews yet.
No comments yet.