diff --git a/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/benchmark_live.py b/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/benchmark_live.py new file mode 100644 index 000000000000..d016860b790f --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/benchmark_live.py @@ -0,0 +1,194 @@ +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT + +"""Live gfx942 comparison of AITER FlyDSL and rocKE FP8 MQA logits. + +Both paths consume the same tensors, write the same dense FP32 output, use the +same stream, and are timed by the same HIP event timer. The rocKE build and +launch path comes from the packaged example, so documentation and benchmarking +exercise the same instance builder. +""" + +from __future__ import annotations + +import argparse +import csv +from pathlib import Path +import statistics +import sys + +from rocke.examples.gfx942.fp8_mqa_logits.fp8_mqa_logits_verify import ( + ARCH, + build_runner, + compare_outputs, + gfx_name, + make_inputs, + parse_shape, + select_spec, + variant_name, +) +from rocke.runtime import synchronize_and_release, time_launches + + +DEFAULT_SHAPES = ( + (4096, 4096), + (8192, 8192), + (128, 32768), + (671, 131072), +) + + +def _time_pair( + aiter_call, + rocke_call, + *, + stream: int, + warmup: int, + iters: int, + repeats: int, +) -> tuple[float, float]: + """Alternate timing order and return median AITER/rocKE latencies.""" + + aiter_samples = [] + rocke_samples = [] + for repeat in range(repeats): + ordered = ( + (("aiter", aiter_call), ("rocke", rocke_call)) + if repeat % 2 == 0 + else (("rocke", rocke_call), ("aiter", aiter_call)) + ) + for name, call in ordered: + elapsed = time_launches( + call, + warmup=warmup, + iters=iters, + stream=stream, + ) + synchronize_and_release(stream) + if name == "aiter": + aiter_samples.append(elapsed) + else: + rocke_samples.append(elapsed) + return statistics.median(aiter_samples), statistics.median(rocke_samples) + + +def _write_csv(rows: list[dict], output: Path | None) -> None: + """Write the result table to stdout and, when requested, a CSV file.""" + + fields = list(rows[0]) + stdout_writer = csv.DictWriter(sys.stdout, fieldnames=fields) + stdout_writer.writeheader() + stdout_writer.writerows(rows) + if output is not None: + output.parent.mkdir(parents=True, exist_ok=True) + with output.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--shapes", + nargs="+", + type=parse_shape, + default=DEFAULT_SHAPES, + ) + parser.add_argument("--num-heads", type=int, default=64) + parser.add_argument("--head-dim", type=int, default=128) + parser.add_argument("--block-kv", type=int) + parser.add_argument("--rows-per-block", type=int) + parser.add_argument("--waves-per-block", type=int) + parser.add_argument("--waves-per-eu", type=int, default=2) + parser.add_argument("--target-blocks-per-cu", type=int, default=4) + parser.add_argument("--num-splits", type=int) + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--iters", type=int, default=50) + parser.add_argument("--repeats", type=int, default=5) + parser.add_argument( + "--output-csv", + type=Path, + help="also persist the emitted result table", + ) + args = parser.parse_args() + + if gfx_name() != ARCH: + raise RuntimeError( + f"this comparison requires {ARCH}; current device is {gfx_name()}" + ) + try: + from aiter.ops.flydsl import flydsl_fp8_mqa_logits + except ImportError as exc: + raise RuntimeError( + "AITER with the FlyDSL fp8_mqa_logits op must be on PYTHONPATH" + ) from exc + + rows = [] + for seq_q, seq_kv in args.shapes: + inputs = make_inputs( + seq_q, + seq_kv, + args.num_heads, + args.head_dim, + ) + spec = select_spec( + seq_q, + seq_kv, + args.num_heads, + args.head_dim, + block_kv=args.block_kv, + rows_per_block=args.rows_per_block, + waves_per_block=args.waves_per_block, + waves_per_eu=None if args.waves_per_eu == 0 else args.waves_per_eu, + ) + rocke_call, rocke_output, stream, num_splits, _kernel_name = build_runner( + inputs, + seq_q, + seq_kv, + spec, + target_blocks_per_cu=args.target_blocks_per_cu, + num_splits_override=args.num_splits, + ) + + def aiter_call(): + return flydsl_fp8_mqa_logits( + inputs["q"], + inputs["kv"], + inputs["kv_scales"], + inputs["weights"], + inputs["starts"], + inputs["ends"], + True, + ) + + aiter_output = aiter_call() + rocke_call() + synchronize_and_release(stream) + diff, _max_abs = compare_outputs(aiter_output, rocke_output, seq_q) + aiter_ms, rocke_ms = _time_pair( + aiter_call, + rocke_call, + stream=stream, + warmup=args.warmup, + iters=args.iters, + repeats=args.repeats, + ) + rows.append( + { + "seq_q": seq_q, + "seq_kv": seq_kv, + "aiter_ms": aiter_ms, + "rocke_ms": rocke_ms, + "rocke_vs_aiter": aiter_ms / rocke_ms, + "calc_diff": diff, + "rocke_variant": variant_name(spec, num_splits), + } + ) + + _write_csv(rows, args.output_csv) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/fp8_mqa_logits_perf.csv b/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/fp8_mqa_logits_perf.csv new file mode 100644 index 000000000000..9604f59ce70b --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/library/benchmarks/gfx942/fp8_mqa_logits/fp8_mqa_logits_perf.csv @@ -0,0 +1,5 @@ +seq_q,seq_kv,aiter_ms,rocke_ms,rocke_vs_aiter,calc_diff,rocke_variant +4096,4096,0.23657543182373048,0.2165285110473633,1.0925832846649102,9.992007221626409e-16,b64_r4_w2_wpe2_s2 +8192,8192,0.7487397003173828,0.7139015960693359,1.0487995886826162,1.2212453270876722e-15,b64_r4_w2_wpe2_s1 +128,32768,0.13252405166625977,0.113355712890625,1.1690990095411422,9.992007221626409e-16,b128_r2_w2_wpe2_s19 +671,131072,1.8947146606445313,1.72795166015625,1.0965090658110204,1.1102230246251565e-15,b64_r4_w4_wpe2_s18 diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/cpp/bindings/rocke_engine.cpp b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/bindings/rocke_engine.cpp index fc9ce4ae9dd6..77a71c78063c 100644 --- a/dnn-providers/hip-kernel-provider/rocke/platform/cpp/bindings/rocke_engine.cpp +++ b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/bindings/rocke_engine.cpp @@ -45,6 +45,7 @@ extern "C" { #include "rocke/instance_gemm_multi_abd.h" #include "rocke/instance_gemm_multi_d.h" #include "rocke/instance_gemm_universal.h" +#include "rocke/instance_gfx942_fp8_mqa_logits.h" #include "rocke/instance_grouped_gemm.h" #include "rocke/instance_matmul_nbits.h" #include "rocke/instance_mfma_gemm.h" @@ -1023,6 +1024,64 @@ std::vector mx_gemm_verify(const py::dict& d, const std::string& ar rocke_build_mx_gemm_new(&b, &s, arch_or_default(arch))); } +/* =========================== FP8 logits ============================== */ + +rocke_fp8_mqa_logits_spec_t fp8_mqa_logits_build_spec(const py::dict& d, + std::deque& store) +{ + auto keep = [&](const std::string& s) -> const char* { + store.push_back(s); + return store.back().c_str(); + }; + rocke_fp8_mqa_logits_spec_t s = rocke_fp8_mqa_logits_spec_default(); + s.num_heads = dict_int(d, "num_heads", s.num_heads); + s.head_dim = dict_int(d, "head_dim", s.head_dim); + s.block_kv = dict_int(d, "block_kv", s.block_kv); + s.rows_per_block = dict_int(d, "rows_per_block", s.rows_per_block); + s.waves_per_block = dict_int(d, "waves_per_block", s.waves_per_block); + if(d.contains("waves_per_eu")) + { + py::object value = d["waves_per_eu"]; + s.has_waves_per_eu = !value.is_none(); + if(s.has_waves_per_eu) + s.waves_per_eu = py::cast(value); + } + { + std::string value; + if(dict_str(d, "name", value)) + s.name = keep(value); + } + return s; +} + +std::string fp8_mqa_logits_lower_llvm(const py::dict& d, const std::string& arch) +{ + std::deque store; + rocke_fp8_mqa_logits_spec_t s = fp8_mqa_logits_build_spec(d, store); + char* ll = nullptr; + char err[ROCKE_ERR_MSG_CAP]; + err[0] = '\0'; + rocke_status_t st = rocke_fp8_mqa_logits_lower_to_llvm( + &s, arch_or_default(arch), ROCKE_LLVM_FLAVOR_AUTO, &ll, err, sizeof err); + return take_lowered(st, ll, err, "rocke_engine.fp8_mqa_logits_lower_llvm"); +} + +std::string fp8_mqa_logits_serialize_ir(const py::dict& d, const std::string& arch) +{ + ROCKE_FAMILY_SERIALIZE_BODY("rocke_engine.fp8_mqa_logits_serialize_ir", + rocke_fp8_mqa_logits_spec_t, + fp8_mqa_logits_build_spec, + rocke_build_fp8_mqa_logits_new(&b, &s, arch_or_default(arch))); +} + +std::vector fp8_mqa_logits_verify(const py::dict& d, const std::string& arch) +{ + ROCKE_FAMILY_VERIFY_BODY("rocke_engine.fp8_mqa_logits_verify", + rocke_fp8_mqa_logits_spec_t, + fp8_mqa_logits_build_spec, + rocke_build_fp8_mqa_logits_new(&b, &s, arch_or_default(arch))); +} + /* ============================== mfma GEMM ============================= */ rocke_mfma_gemm_spec_t mfma_build_spec(const py::dict& d, std::deque& store) @@ -3466,6 +3525,18 @@ PYBIND11_MODULE(rocke_engine, m) &block_scale_gemm_serialize_ir, &block_scale_gemm_verify); reg3("mx_gemm", &mx_gemm_lower_llvm, &mx_gemm_serialize_ir, &mx_gemm_verify); + m.def("fp8_mqa_logits_lower_llvm", + &fp8_mqa_logits_lower_llvm, + py::arg("spec"), + py::arg("arch") = "gfx942"); + m.def("fp8_mqa_logits_serialize_ir", + &fp8_mqa_logits_serialize_ir, + py::arg("spec"), + py::arg("arch") = "gfx942"); + m.def("fp8_mqa_logits_verify", + &fp8_mqa_logits_verify, + py::arg("spec"), + py::arg("arch") = "gfx942"); reg3("mfma_gemm", &mfma_gemm_lower_llvm, &mfma_gemm_serialize_ir, &mfma_gemm_verify); m.def("mfma_gemm_is_valid", &mfma_gemm_is_valid, py::arg("spec"), py::arg("arch") = "gfx950"); diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/cpp/include/rocke/instance_gfx942_fp8_mqa_logits.h b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/include/rocke/instance_gfx942_fp8_mqa_logits.h new file mode 100644 index 000000000000..943872c59824 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/include/rocke/instance_gfx942_fp8_mqa_logits.h @@ -0,0 +1,88 @@ +/* Copyright (c) Advanced Micro Devices, Inc., or its affiliates. + * SPDX-License-Identifier: MIT + * + * rocke/instance_gfx942_fp8_mqa_logits.h -- C99 port of + * rocke/instances/gfx942/fp8_mqa_logits.py. + * + * Python (fp8_mqa_logits.py) C99 (this header) + * ---------------------------------- ----------------------------------------- + * class Fp8MqaLogitsSpec rocke_fp8_mqa_logits_spec_t + * Fp8MqaLogitsSpec.kernel_name() rocke_fp8_mqa_logits_kernel_name(...) + * is_valid_spec(spec, arch) rocke_fp8_mqa_logits_is_valid_spec(...) + * build_fp8_mqa_logits(spec, arch) rocke_build_fp8_mqa_logits(...) + * fp8_mqa_logits_num_splits(...) rocke_fp8_mqa_logits_num_splits(...) + * fp8_mqa_logits_grid(...) rocke_fp8_mqa_logits_grid(...) + * fp8_mqa_logits_signature(spec) rocke_fp8_mqa_logits_signature(...) + */ +#ifndef ROCKE_INSTANCE_GFX942_FP8_MQA_LOGITS_H +#define ROCKE_INSTANCE_GFX942_FP8_MQA_LOGITS_H + +#include +#include + +#include "rocke/helper_rocke.helpers.spec.h" +#include "rocke/ir.h" +#include "rocke/lower_llvm.h" + +#ifdef __cplusplus +extern "C" { +#endif + +typedef struct rocke_fp8_mqa_logits_spec +{ + int num_heads; + int head_dim; + int block_kv; + int rows_per_block; + int waves_per_block; + bool has_waves_per_eu; + int waves_per_eu; + const char* name; +} rocke_fp8_mqa_logits_spec_t; + +rocke_fp8_mqa_logits_spec_t rocke_fp8_mqa_logits_spec_default(void); + +int rocke_fp8_mqa_logits_block_size(const rocke_fp8_mqa_logits_spec_t* spec); + +rocke_status_t rocke_fp8_mqa_logits_kernel_name(const rocke_fp8_mqa_logits_spec_t* spec, + char* out, + size_t out_cap); + +bool rocke_fp8_mqa_logits_is_valid_spec(const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch, + char* reason, + size_t reason_cap); + +rocke_kernel_def_t* rocke_build_fp8_mqa_logits(rocke_ir_builder_t* b, + const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch); + +rocke_kernel_def_t* rocke_build_fp8_mqa_logits_new(rocke_ir_builder_t* b, + const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch); + +rocke_status_t rocke_fp8_mqa_logits_lower_to_llvm(const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch, + rocke_llvm_flavor_t flavor, + char** out_ll, + char* err, + size_t err_cap); + +int rocke_fp8_mqa_logits_num_splits( + int seq_len_padded, int seq_len_kv, int rows_per_block, int block_kv, int num_cus); + +rocke_status_t rocke_fp8_mqa_logits_grid(int seq_len_padded, + int num_splits, + const rocke_fp8_mqa_logits_spec_t* spec, + int out[3]); + +rocke_status_t rocke_fp8_mqa_logits_signature(rocke_arena_t* arena, + const rocke_fp8_mqa_logits_spec_t* spec, + const rocke_sig_entry_t** out_items, + size_t* out_count); + +#ifdef __cplusplus +} /* extern "C" */ +#endif + +#endif /* ROCKE_INSTANCE_GFX942_FP8_MQA_LOGITS_H */ diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/cpp/instances/gfx942/fp8_mqa_logits.cpp b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/instances/gfx942/fp8_mqa_logits.cpp new file mode 100644 index 000000000000..b4fbe5337df5 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/cpp/instances/gfx942/fp8_mqa_logits.cpp @@ -0,0 +1,635 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT +/* + * fp8_mqa_logits.cpp -- C99 port of + * rocke/instances/gfx942/fp8_mqa_logits.py. + * + * The builder mirrors build_fp8_mqa_logits one builder call at a time. Host + * loops and temporary arrays do not emit IR; their traversal order matches the + * Python list construction and nested loops. + */ +#include "rocke/instance_gfx942_fp8_mqa_logits.h" + +#include +#include +#include + +#include "rocke/error_boundary.hpp" +#include "rocke/helper_rocke.core.arch.h" +#include "rocke/helper_rocke.helpers.atoms.h" +#include "rocke/helper_rocke.helpers.mfma_gemm_inner.h" +#include "rocke/ir_internal.h" + +static const int ROCKE_FP8_MQA_MIN_TILES_PER_SPLIT = 8; + +static void rocke_fp8_mqa_set_reason(char* reason, size_t cap, const char* message) +{ + rocke_spec_set_reason(reason, cap, message); +} + +rocke_fp8_mqa_logits_spec_t rocke_fp8_mqa_logits_spec_default(void) +{ + rocke_fp8_mqa_logits_spec_t spec; + memset(&spec, 0, sizeof(spec)); + spec.num_heads = 64; + spec.head_dim = 128; + spec.block_kv = 128; + spec.rows_per_block = 2; + spec.waves_per_block = 4; + spec.has_waves_per_eu = true; + spec.waves_per_eu = 2; + spec.name = "rocke_fp8_mqa_logits"; + return spec; +} + +int rocke_fp8_mqa_logits_block_size(const rocke_fp8_mqa_logits_spec_t* spec) +{ + return spec != NULL ? 64 * spec->waves_per_block : 0; +} + +rocke_status_t rocke_fp8_mqa_logits_kernel_name(const rocke_fp8_mqa_logits_spec_t* spec, + char* out, + size_t out_cap) +{ + char heads[32]; + char dim[32]; + char block[32]; + char rows[32]; + char waves[32]; + const char* parts[5]; + + if(spec == NULL || spec->name == NULL || out == NULL || out_cap == 0) + { + return ROCKE_ERR_VALUE; + } + snprintf(heads, sizeof(heads), "H%d", spec->num_heads); + snprintf(dim, sizeof(dim), "D%d", spec->head_dim); + snprintf(block, sizeof(block), "BKV%d", spec->block_kv); + snprintf(rows, sizeof(rows), "R%d", spec->rows_per_block); + snprintf(waves, sizeof(waves), "W%d", spec->waves_per_block); + parts[0] = heads; + parts[1] = dim; + parts[2] = block; + parts[3] = rows; + parts[4] = waves; + return rocke_kernel_name_join(spec->name, parts, 5, NULL, NULL, 0, out, out_cap, NULL); +} + +bool rocke_fp8_mqa_logits_is_valid_spec(const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch, + char* reason, + size_t reason_cap) +{ + const rocke_archtarget_t* target; + const rocke_mfma_atom_t* atom; + int block_size; + int n_tiles; + char buffer[192]; + + if(spec == NULL) + { + rocke_fp8_mqa_set_reason(reason, reason_cap, "null spec"); + return false; + } + if(arch == NULL) + { + arch = "gfx942"; + } + if(strcmp(arch, "gfx942") != 0) + { + snprintf(buffer, + sizeof(buffer), + "fp8_mqa_logits currently supports gfx942 only, got '%s'", + arch); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + + target = rocke_archtarget_from_gfx(arch); + if(target == NULL) + { + snprintf(buffer, sizeof(buffer), "unknown gfx target %s", arch); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + block_size = rocke_fp8_mqa_logits_block_size(spec); + if(block_size > rocke_archtarget_max_threads_per_block(target)) + { + snprintf(buffer, + sizeof(buffer), + "block_size %d > %d (hardware cap) on %s", + block_size, + rocke_archtarget_max_threads_per_block(target), + arch); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + + atom = rocke_mfma_atom("fp8e4m3", 16, 16, 32); + if(atom == NULL) + { + rocke_fp8_mqa_set_reason(reason, reason_cap, "missing fp8e4m3 MFMA atom"); + return false; + } + if(spec->num_heads <= 0 || spec->num_heads % atom->m) + { + snprintf(buffer, sizeof(buffer), "num_heads must be a positive multiple of %d", atom->m); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + if(spec->head_dim <= 0 || spec->head_dim % atom->k) + { + snprintf(buffer, sizeof(buffer), "head_dim must be a positive multiple of %d", atom->k); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + if(spec->block_kv <= 0 || spec->block_kv % atom->n) + { + snprintf(buffer, sizeof(buffer), "block_kv must be a positive multiple of %d", atom->n); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + if(spec->rows_per_block <= 0) + { + rocke_fp8_mqa_set_reason(reason, reason_cap, "rows_per_block must be positive"); + return false; + } + if(spec->waves_per_block <= 0) + { + rocke_fp8_mqa_set_reason(reason, reason_cap, "waves_per_block must be positive"); + return false; + } + n_tiles = spec->block_kv / atom->n; + if(n_tiles % spec->waves_per_block) + { + snprintf(buffer, + sizeof(buffer), + "block_kv / %d (%d) must be divisible by waves_per_block (%d)", + atom->n, + n_tiles, + spec->waves_per_block); + rocke_fp8_mqa_set_reason(reason, reason_cap, buffer); + return false; + } + if(spec->has_waves_per_eu && spec->waves_per_eu <= 0) + { + rocke_fp8_mqa_set_reason(reason, reason_cap, "waves_per_eu must be positive or None"); + return false; + } + rocke_fp8_mqa_set_reason(reason, reason_cap, "ok"); + return true; +} + +static rocke_value_t* + rocke_fp8_mqa_ceildiv(rocke_ir_builder_t* b, rocke_value_t* value, int divisor) +{ + rocke_value_t* one_less = rocke_b_const_i32(b, divisor - 1); + rocke_value_t* divisor_value = rocke_b_const_i32(b, divisor); + return rocke_b_div(b, rocke_b_add(b, value, one_less), divisor_value); +} + +static rocke_value_t* + rocke_fp8_mqa_ceildiv_value(rocke_ir_builder_t* b, rocke_value_t* value, rocke_value_t* divisor) +{ + rocke_value_t* one_less = rocke_b_sub(b, divisor, rocke_b_const_i32(b, 1)); + return rocke_b_div(b, rocke_b_add(b, value, one_less), divisor); +} + +static rocke_kernel_def_t* rocke_build_fp8_mqa_logits_impl(rocke_ir_builder_t* b, + const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch) +{ + char reason[192]; + const rocke_mfma_atom_t* atom; + int h; + int d; + int bkv; + int rpb; + int wpb; + int n_tiles_per_wave; + int m_tiles; + int k_steps; + rocke_param_opts_t opts; + rocke_value_t* q; + rocke_value_t* kv; + rocke_value_t* kv_scales; + rocke_value_t* weights; + rocke_value_t* cu_starts; + rocke_value_t* cu_ends; + rocke_value_t* logits; + rocke_value_t* seq_len; + rocke_value_t* seq_len_kv; + rocke_value_t* stride_logits_s; + rocke_value_t* num_splits; + rocke_value_t* tid; + rocke_value_t* bid; + rocke_value_t* split_id; + rocke_value_t* wave; + rocke_value_t* lane; + rocke_lane_decode_t lane_decode; + rocke_value_t* n_blocks; + rocke_value_t* reverse_bid; + rocke_value_t* row0; + rocke_value_t* zero_i32; + rocke_value_t* zero_f32; + std::vector starts; + std::vector ends; + std::vector q_fragments; + std::vector weight_fragments; + rocke_value_t* tile_start; + rocke_value_t* tile_end; + rocke_value_t* window_tiles; + rocke_value_t* split_columns; + rocke_for_t tile_loop; + + if(b == NULL || spec == NULL) + { + return NULL; + } + if(arch == NULL) + { + arch = "gfx942"; + } + if(!rocke_fp8_mqa_logits_is_valid_spec(spec, arch, reason, sizeof(reason))) + { + return (rocke_kernel_def_t*)rocke_i_set_err( + b, ROCKE_ERR_VALUE, "invalid fp8_mqa_logits spec: %s", reason); + } + + atom = rocke_mfma_atom("fp8e4m3", 16, 16, 32); + if(rocke_validate_mfma_atom_in_catalog(b, atom, arch, "fp8_mqa_logits") != ROCKE_OK) + { + return NULL; + } + h = spec->num_heads; + d = spec->head_dim; + bkv = spec->block_kv; + rpb = spec->rows_per_block; + wpb = spec->waves_per_block; + n_tiles_per_wave = (bkv / atom->n) / wpb; + m_tiles = h / atom->m; + k_steps = d / atom->k; + + rocke_attr_set_int(b, &b->kernel->attrs, "max_workgroup_size", 64 * wpb); + if(spec->has_waves_per_eu) + { + rocke_attr_set_int(b, &b->kernel->attrs, "waves_per_eu", spec->waves_per_eu); + } + + memset(&opts, 0, sizeof(opts)); + opts.readonly = true; + opts.readonly_set = true; + opts.align = 8; + opts.align_set = true; + q = rocke_b_param(b, "Q", rocke_ptr_type(b, rocke_fp8e4m3(), "global"), &opts); + kv = rocke_b_param(b, "KV", rocke_ptr_type(b, rocke_fp8e4m3(), "global"), &opts); + + opts.align = 4; + kv_scales = rocke_b_param(b, "kv_scales", rocke_ptr_type(b, rocke_f32(), "global"), &opts); + weights = rocke_b_param(b, "weights", rocke_ptr_type(b, rocke_f32(), "global"), &opts); + cu_starts = rocke_b_param(b, "cu_starts", rocke_ptr_type(b, rocke_i32(), "global"), &opts); + cu_ends = rocke_b_param(b, "cu_ends", rocke_ptr_type(b, rocke_i32(), "global"), &opts); + + memset(&opts, 0, sizeof(opts)); + opts.writeonly = true; + opts.writeonly_set = true; + opts.align = 4; + opts.align_set = true; + logits = rocke_b_param(b, "logits", rocke_ptr_type(b, rocke_f32(), "global"), &opts); + seq_len = rocke_b_param(b, "seq_len", rocke_i32(), NULL); + seq_len_kv = rocke_b_param(b, "seq_len_kv", rocke_i32(), NULL); + stride_logits_s = rocke_b_param(b, "stride_logits_s", rocke_i32(), NULL); + num_splits = rocke_b_param(b, "num_splits", rocke_i32(), NULL); + + tid = rocke_b_thread_id_x(b); + bid = rocke_b_block_id_x(b); + split_id = rocke_b_block_id_y(b); + wave = rocke_b_div(b, tid, rocke_b_const_i32(b, 64)); + lane = rocke_b_mod(b, tid, rocke_b_const_i32(b, 64)); + lane_decode = rocke_decode_mfma_lanes(b, atom, lane); + + n_blocks = rocke_fp8_mqa_ceildiv(b, seq_len, rpb); + { + rocke_value_t* forward_bid = rocke_b_sub(b, n_blocks, bid); + rocke_value_t* one = rocke_b_const_i32(b, 1); + reverse_bid = rocke_b_sub(b, forward_bid, one); + } + row0 = rocke_b_mul(b, reverse_bid, rocke_b_const_i32(b, rpb)); + zero_i32 = rocke_b_const_i32(b, 0); + zero_f32 = rocke_b_const_f32(b, 0.0); + + starts.reserve((size_t)rpb); + ends.reserve((size_t)rpb); + q_fragments.reserve((size_t)rpb * (size_t)m_tiles * (size_t)k_steps); + weight_fragments.reserve((size_t)rpb * (size_t)m_tiles * (size_t)atom->c_per_lane); + for(int row_offset = 0; row_offset < rpb; ++row_offset) + { + rocke_value_t* row = rocke_b_add(b, row0, rocke_b_const_i32(b, row_offset)); + rocke_value_t* start + = rocke_b_smax(b, rocke_b_global_load_i32(b, cu_starts, row, 0), zero_i32); + rocke_value_t* end + = rocke_b_smin(b, rocke_b_global_load_i32(b, cu_ends, row, 0), seq_len_kv); + starts.push_back(start); + ends.push_back(end); + + for(int mi = 0; mi < m_tiles; ++mi) + { + rocke_value_t* head + = rocke_b_add(b, rocke_b_const_i32(b, mi * atom->m), lane_decode.m_in_atom); + rocke_value_t* row_head + = rocke_b_add(b, rocke_b_mul(b, row, rocke_b_const_i32(b, h)), head); + rocke_value_t* q_base = rocke_b_mul(b, row_head, rocke_b_const_i32(b, d)); + for(int kk = 0; kk < k_steps; ++kk) + { + rocke_value_t* k_step = rocke_b_const_i32(b, kk * atom->k); + rocke_value_t* lane_width = rocke_b_const_i32(b, atom->a_per_lane); + rocke_value_t* lane_offset = rocke_b_mul(b, lane_decode.k_blk, lane_width); + rocke_value_t* k_lane = rocke_b_add(b, k_step, lane_offset); + rocke_value_t* q_addr = rocke_b_add(b, q_base, k_lane); + q_fragments.push_back(rocke_b_global_load_vN( + b, q, q_addr, rocke_fp8e4m3(), atom->a_per_lane, atom->a_per_lane)); + } + + for(int elem = 0; elem < atom->c_per_lane; ++elem) + { + rocke_value_t* c_width = rocke_b_const_i32(b, atom->c_per_lane); + rocke_value_t* head_base = rocke_b_mul(b, lane_decode.k_blk, c_width); + rocke_value_t* elem_value = rocke_b_const_i32(b, elem); + rocke_value_t* head_offset = rocke_b_add(b, head_base, elem_value); + rocke_value_t* weight_head + = rocke_b_add(b, rocke_b_const_i32(b, mi * atom->m), head_offset); + rocke_value_t* weight_addr + = rocke_b_add(b, rocke_b_mul(b, row, rocke_b_const_i32(b, h)), weight_head); + weight_fragments.push_back(rocke_b_global_load_f32(b, weights, weight_addr, 0)); + } + } + } + + tile_start = starts[0]; + tile_end = ends[0]; + for(int row_offset = 1; row_offset < rpb; ++row_offset) + { + tile_start = rocke_b_smin(b, tile_start, starts[(size_t)row_offset]); + tile_end = rocke_b_smax(b, tile_end, ends[(size_t)row_offset]); + } + { + rocke_value_t* divisor = rocke_b_const_i32(b, bkv); + rocke_value_t* tile_index = rocke_b_div(b, tile_start, divisor); + rocke_value_t* multiplier = rocke_b_const_i32(b, bkv); + tile_start = rocke_b_mul(b, tile_index, multiplier); + } + + window_tiles = rocke_fp8_mqa_ceildiv(b, rocke_b_sub(b, tile_end, tile_start), bkv); + { + rocke_value_t* split_tiles = rocke_fp8_mqa_ceildiv_value(b, window_tiles, num_splits); + rocke_value_t* block_width = rocke_b_const_i32(b, bkv); + split_columns = rocke_b_mul(b, split_tiles, block_width); + } + tile_start = rocke_b_add(b, tile_start, rocke_b_mul(b, split_id, split_columns)); + tile_end = rocke_b_smin(b, rocke_b_add(b, tile_start, split_columns), tile_end); + + tile_loop = rocke_b_scf_for(b, tile_start, tile_end, rocke_b_const_i32(b, bkv), "col0"); + rocke_b_region_enter(b, tile_loop.body); + { + rocke_value_t* col0 = tile_loop.iv; + rocke_value_t* wave_tile_base + = rocke_b_mul(b, wave, rocke_b_const_i32(b, n_tiles_per_wave)); + std::vector columns; + std::vector scales; + std::vector kv_fragments; + columns.reserve((size_t)n_tiles_per_wave); + scales.reserve((size_t)n_tiles_per_wave); + kv_fragments.reserve((size_t)n_tiles_per_wave * (size_t)k_steps); + + for(int ni = 0; ni < n_tiles_per_wave; ++ni) + { + rocke_value_t* absolute_ni = rocke_b_add(b, wave_tile_base, rocke_b_const_i32(b, ni)); + rocke_value_t* column = rocke_b_add( + b, + rocke_b_add(b, col0, rocke_b_mul(b, absolute_ni, rocke_b_const_i32(b, atom->n))), + lane_decode.n_in_atom); + rocke_value_t* clamped_column + = rocke_b_smin(b, column, rocke_b_sub(b, seq_len_kv, rocke_b_const_i32(b, 1))); + columns.push_back(column); + scales.push_back(rocke_b_global_load_f32(b, kv_scales, clamped_column, 0)); + rocke_value_t* kv_base = rocke_b_mul(b, clamped_column, rocke_b_const_i32(b, d)); + for(int kk = 0; kk < k_steps; ++kk) + { + rocke_value_t* k_step = rocke_b_const_i32(b, kk * atom->k); + rocke_value_t* lane_width = rocke_b_const_i32(b, atom->b_per_lane); + rocke_value_t* lane_offset = rocke_b_mul(b, lane_decode.k_blk, lane_width); + rocke_value_t* k_lane = rocke_b_add(b, k_step, lane_offset); + rocke_value_t* kv_addr = rocke_b_add(b, kv_base, k_lane); + kv_fragments.push_back(rocke_b_global_load_vN( + b, kv, kv_addr, rocke_fp8e4m3(), atom->b_per_lane, atom->b_per_lane)); + } + } + + for(int row_offset = 0; row_offset < rpb; ++row_offset) + { + rocke_value_t* row = rocke_b_add(b, row0, rocke_b_const_i32(b, row_offset)); + rocke_value_t* row_i64 = rocke_b_sext(b, row, rocke_i64()); + rocke_value_t* stride_i64 = rocke_b_sext(b, stride_logits_s, rocke_i64()); + rocke_value_t* row_stride = rocke_b_mul(b, row_i64, stride_i64); + rocke_value_t* four = rocke_b_const_i64(b, 4); + rocke_value_t* row_byte_offset = rocke_b_mul(b, row_stride, four); + rocke_value_t* logits_row = rocke_b_global_ptr_add(b, logits, row_byte_offset); + for(int ni = 0; ni < n_tiles_per_wave; ++ni) + { + rocke_value_t* column_sum = zero_f32; + for(int mi = 0; mi < m_tiles; ++mi) + { + rocke_value_t* accumulator = rocke_b_zero_vec_f32(b, atom->c_per_lane); + for(int kk = 0; kk < k_steps; ++kk) + { + size_t q_index + = ((size_t)row_offset * (size_t)m_tiles + (size_t)mi) * (size_t)k_steps + + (size_t)kk; + size_t kv_index = (size_t)ni * (size_t)k_steps + (size_t)kk; + accumulator = rocke_b_mma(b, + atom->name, + q_fragments[q_index], + kv_fragments[kv_index], + accumulator, + NULL, + 0); + } + for(int elem = 0; elem < atom->c_per_lane; ++elem) + { + size_t weight_index = ((size_t)row_offset * (size_t)m_tiles + (size_t)mi) + * (size_t)atom->c_per_lane + + (size_t)elem; + rocke_value_t* score = rocke_b_vec_extract(b, accumulator, elem); + rocke_value_t* relu = rocke_b_fmax(b, score, zero_f32); + column_sum + = rocke_b_fma(b, relu, weight_fragments[weight_index], column_sum); + } + } + column_sum = rocke_b_fmul(b, column_sum, scales[(size_t)ni]); + column_sum + = rocke_b_fadd(b, column_sum, rocke_b_warp_shuffle_xor(b, column_sum, 16)); + column_sum + = rocke_b_fadd(b, column_sum, rocke_b_warp_shuffle_xor(b, column_sum, 32)); + + rocke_value_t* after_start + = rocke_b_cmp_ge(b, columns[(size_t)ni], starts[(size_t)row_offset]); + rocke_value_t* before_end + = rocke_b_cmp_lt(b, columns[(size_t)ni], ends[(size_t)row_offset]); + rocke_value_t* in_window = rocke_b_land(b, after_start, before_end); + rocke_value_t* first_k = rocke_b_cmp_eq(b, lane_decode.k_blk, zero_i32); + rocke_value_t* is_writer = rocke_b_land(b, first_k, in_window); + rocke_if_t write_if = rocke_b_scf_if(b, is_writer); + rocke_b_region_enter(b, write_if.then_region); + rocke_b_global_store(b, logits_row, columns[(size_t)ni], column_sum, 4); + rocke_b_region_leave(b); + } + } + } + rocke_b_region_leave(b); + + rocke_b_ret(b); + return b->kernel; +} + +rocke_kernel_def_t* rocke_build_fp8_mqa_logits(rocke_ir_builder_t* b, + const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch) +{ + return ckc::guard_builder( + b, [&]() -> rocke_kernel_def_t* { return rocke_build_fp8_mqa_logits_impl(b, spec, arch); }); +} + +rocke_kernel_def_t* rocke_build_fp8_mqa_logits_new(rocke_ir_builder_t* b, + const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch) +{ + return ckc::guard_builder(b, [&]() -> rocke_kernel_def_t* { + char name[256]; + if(b == NULL || spec == NULL) + { + return NULL; + } + if(rocke_fp8_mqa_logits_kernel_name(spec, name, sizeof(name)) != ROCKE_OK) + { + return NULL; + } + if(rocke_ir_builder_init(b, name) != ROCKE_OK) + { + return NULL; + } + return rocke_build_fp8_mqa_logits_impl(b, spec, arch); + }); +} + +rocke_status_t rocke_fp8_mqa_logits_lower_to_llvm(const rocke_fp8_mqa_logits_spec_t* spec, + const char* arch, + rocke_llvm_flavor_t flavor, + char** out_ll, + char* err, + size_t err_cap) +{ + rocke_ir_builder_t b; + rocke_kernel_def_t* kernel; + rocke_status_t status; + + if(out_ll != NULL) + { + *out_ll = NULL; + } + if(spec == NULL || out_ll == NULL) + { + rocke_fp8_mqa_set_reason(err, err_cap, "lower_to_llvm: null spec/out"); + return ROCKE_ERR_VALUE; + } + if(arch == NULL) + { + arch = "gfx942"; + } + kernel = rocke_build_fp8_mqa_logits_new(&b, spec, arch); + if(kernel == NULL) + { + status = rocke_ir_builder_status(&b); + rocke_fp8_mqa_set_reason(err, err_cap, rocke_ir_builder_error(&b)); + rocke_ir_builder_free(&b); + return status == ROCKE_OK ? ROCKE_ERR_VALUE : status; + } + status = rocke_lower_kernel_to_llvm_ex(kernel, flavor, arch, out_ll, err, err_cap); + rocke_ir_builder_free(&b); + return status; +} + +int rocke_fp8_mqa_logits_num_splits( + int seq_len_padded, int seq_len_kv, int rows_per_block, int block_kv, int num_cus) +{ + int grid_x = seq_len_padded / rows_per_block; + int target_blocks; + int max_splits; + int needed; + + if(grid_x == 0 || seq_len_kv < 4096) + { + return 1; + } + target_blocks = 4 * num_cus; + if(grid_x >= target_blocks) + { + return 1; + } + max_splits = (seq_len_kv / block_kv) / ROCKE_FP8_MQA_MIN_TILES_PER_SPLIT; + if(max_splits < 1) + { + max_splits = 1; + } + needed = (target_blocks + grid_x - 1) / grid_x; + if(needed > max_splits) + { + needed = max_splits; + } + return needed > 1 ? needed : 1; +} + +rocke_status_t rocke_fp8_mqa_logits_grid(int seq_len_padded, + int num_splits, + const rocke_fp8_mqa_logits_spec_t* spec, + int out[3]) +{ + if(spec == NULL || out == NULL || spec->rows_per_block <= 0 + || seq_len_padded % spec->rows_per_block) + { + return ROCKE_ERR_VALUE; + } + out[0] = seq_len_padded / spec->rows_per_block; + out[1] = num_splits; + out[2] = 1; + return ROCKE_OK; +} + +rocke_status_t rocke_fp8_mqa_logits_signature(rocke_arena_t* arena, + const rocke_fp8_mqa_logits_spec_t* spec, + const rocke_sig_entry_t** out_items, + size_t* out_count) +{ + rocke_signature_builder_t sb; + rocke_status_t status; + if(arena == NULL || spec == NULL || out_items == NULL || out_count == NULL) + { + return ROCKE_ERR_VALUE; + } + status = rocke_signature_builder_init(&sb, arena); + if(status != ROCKE_OK) + { + return status; + } + rocke_signature_builder_ptr(&sb, "Q", "fp8e4m3", NULL); + rocke_signature_builder_ptr(&sb, "KV", "fp8e4m3", NULL); + rocke_signature_builder_ptr(&sb, "kv_scales", "f32", NULL); + rocke_signature_builder_ptr(&sb, "weights", "f32", NULL); + rocke_signature_builder_ptr(&sb, "cu_starts", "i32", NULL); + rocke_signature_builder_ptr(&sb, "cu_ends", "i32", NULL); + rocke_signature_builder_ptr(&sb, "logits", "f32", NULL); + rocke_signature_builder_scalar(&sb, "seq_len", "i32"); + rocke_signature_builder_scalar(&sb, "seq_len_kv", "i32"); + rocke_signature_builder_scalar(&sb, "stride_logits_s", "i32"); + rocke_signature_builder_scalar(&sb, "num_splits", "i32"); + return rocke_signature_builder_build(&sb, out_items, out_count); +} diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/core/backend.py b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/core/backend.py index 41b803151cd9..82b15791348e 100644 --- a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/core/backend.py +++ b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/core/backend.py @@ -741,6 +741,19 @@ def mx_gemm_spec_to_dict(spec: Any) -> Dict[str, Any]: ) +def fp8_mqa_logits_spec_to_dict(spec: Any) -> Dict[str, Any]: + """:class:`Fp8MqaLogitsSpec` -> flat dict.""" + return dict( + name=spec.name, + num_heads=spec.num_heads, + head_dim=spec.head_dim, + block_kv=spec.block_kv, + rows_per_block=spec.rows_per_block, + waves_per_block=spec.waves_per_block, + waves_per_eu=spec.waves_per_eu, + ) + + def mfma_gemm_spec_to_dict(spec: Any) -> Dict[str, Any]: """:class:`MfmaGemmSpec` -> flat dict.""" return dict( @@ -1268,6 +1281,43 @@ def py_fn(wi: bool) -> Tuple[str, str]: ) +def lower_fp8_mqa_logits( + spec: Any, + *, + arch: str = "gfx942", + backend: Optional[str] = None, + want_ir: bool = False, +) -> "GemmLowerResult": + """Lower an :class:`Fp8MqaLogitsSpec`.""" + + def py_fn(wi: bool) -> Tuple[str, str]: + from ..instances.gfx942.fp8_mqa_logits import build_fp8_mqa_logits + from .lower_llvm import lower_kernel_to_llvm + + kernel = build_fp8_mqa_logits(spec, arch=arch) + ll = lower_kernel_to_llvm(kernel, arch=arch) + ir = "" + if wi: + from .ir_serialize import serialize + + ir = serialize(kernel) + return ll, ir + + engine = _import_engine() + spec_dict = fp8_mqa_logits_spec_to_dict(spec) + return _lower_family( + "fp8_mqa_logits", + spec, + arch, + backend, + want_ir, + py_fn, + lambda: engine.fp8_mqa_logits_lower_llvm(spec_dict, arch=arch), + lambda: engine.fp8_mqa_logits_serialize_ir(spec_dict, arch=arch), + _name_of(spec), + ) + + def lower_mfma_gemm( spec: Any, *, diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/README.md b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/README.md new file mode 100644 index 000000000000..77d3daf943e2 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/README.md @@ -0,0 +1,122 @@ +# FP8 MQA logits on gfx942 + +This example builds, launches, verifies, and benchmarks the rocKE FP8 +multi-query-attention logits instance used by the DeepSeek lightning indexer. +It targets MI300-class `gfx942` devices and uses the native +`mfma_f32_16x16x32_fp8` instruction. + +For query row `m` and KV position `n` in that row's valid window: + +```text +logits[m, n] = + sum_h ReLU(dot(Q[m, h, :], KV[n, :]) * kv_scale[n]) * weights[m, h] +``` + +`Q` and `KV` use native gfx942 E4M3 FNUZ encoding. `kv_scale` is expected to be +nonnegative, which permits the kernel to apply it once after the weighted head +sum. Positions outside `[cu_starts[m], cu_ends[m])` are left untouched; the +runner prefills the FP32 output with `-inf`. + +## File map + +| Path | Purpose | +|---|---| +| `fp8_mqa_logits_verify.py` | Shared build, launch, verification, and timing runner | +| `README.md` | Usage, implementation notes, and measured results | +| `rocke/instances/gfx942/fp8_mqa_logits.py` | Python kernel builder | +| `library/benchmarks/gfx942/fp8_mqa_logits/benchmark_live.py` | Live AITER comparison using this runner | +| `library/benchmarks/gfx942/fp8_mqa_logits/fp8_mqa_logits_perf.csv` | Captured results tabulated below | + +The example imports `Fp8MqaLogitsSpec`, `build_fp8_mqa_logits`, +`fp8_mqa_logits_grid`, and `fp8_mqa_logits_signature` directly from the +instance. The live comparison imports its input, variant-selection, compile, and +launch helpers from this example, ensuring both paths exercise the same builder. + +## Requirements + +- AMD `gfx942` GPU +- PyTorch with ROCm and `torch.float8_e4m3fnuz` +- A working rocKE Python/COMGR/HIP environment +- AITER PR #3913 and FlyDSL only for the live comparison + +Run from the rocKE directory: + +```bash +cd /dnn-providers/hip-kernel-provider/rocke +export PYTHONPATH="$(pwd)/platform/python:${PYTHONPATH:-}" +``` + +## Build and verify + +The default `4x128` shape is intentionally small enough for the row-at-a-time +PyTorch reference: + +```bash +python -m rocke.examples.gfx942.fp8_mqa_logits.fp8_mqa_logits_verify --verify +``` + +Build, verify, and measure warm launch latency: + +```bash +python -m rocke.examples.gfx942.fp8_mqa_logits.fp8_mqa_logits_verify \ + --shape 128x32768 \ + --bench --warmup 10 --iters 100 --repeats 7 +``` + +Large production shapes should normally be verified against AITER with the live +comparison below; the PyTorch reference is deliberately row-at-a-time to avoid +materializing an `M x H x N` tensor and is not intended as a fast reference. + +The example emits a `PerfJSON:` record when it finishes. This makes it directly +usable as a command for `rocke.benchmark.perf.harness`, including +`--verify --bench` correctness and wall-latency fields. + +## Live AITER comparison + +```bash +export AITER_PATH= +PYTHONPATH="$(pwd)/platform/python:${AITER_PATH}:${PYTHONPATH:-}" \ + python library/benchmarks/gfx942/fp8_mqa_logits/benchmark_live.py \ + --warmup 10 --iters 100 --repeats 7 \ + --output-csv /tmp/fp8_mqa_logits_gfx942.csv +``` + +For remote MI300X execution, stage and run the command through +`rocke.benchmark.remote_test` as described in +`platform/python/rocke/benchmark/remote_test/README.md`. + +## Measured performance + +Measured on one MI300X (`gfx942`) using Slurm job `67690637`. Each number is the +median of seven repeats with 100 timed iterations after 10 warmups. Both +implementations consumed identical tensors, used the same stream and HIP-event +timer, and included dense `-inf` output initialization. Software was PyTorch +2.10.0 with ROCm 7.2.4 and FlyDSL 0.2.2. + +| Query x KV | AITER PR #3913 (ms) | rocKE (ms) | rocKE speedup | rocKE geometry | +|---:|---:|---:|---:|---| +| 4096 x 4096 | 0.2366 | 0.2165 | **1.093x** | `b64_r4_w2_wpe2_s2` | +| 8192 x 8192 | 0.7487 | 0.7139 | **1.049x** | `b64_r4_w2_wpe2_s1` | +| 128 x 32768 | 0.1325 | 0.1134 | **1.169x** | `b128_r2_w2_wpe2_s19` | +| 671 x 131072 | 1.8947 | 1.7280 | **1.097x** | `b64_r4_w4_wpe2_s18` | + +The geometric-mean speedup is **1.101x**. Similarity error against AITER ranged +from `9.99e-16` to `1.22e-15`, below the `1e-3` correctness threshold. + +These are point measurements, not CI performance guarantees. Re-run the live +benchmark on the target system when changing the compiler, ROCm, AITER, or the +shape-selection policy. + +## Why this version is faster + +- Several query rows share each KV fragment load. +- Q fragments and head weights remain in registers across the KV loop. +- Waves own disjoint KV-column tiles, requiring no cross-wave synchronization. +- The weighted ReLU accumulation uses an explicit FMA, reducing the hot + `b64/r4/w2` kernel from 747 to 619 VALU instructions. +- Geometry and grid-y split density are selected by shape to balance KV reuse, + register pressure, and CU occupancy. + +The implementation uses no LDS and no scratch allocation. Its main remaining +tuning constraint is VGPR pressure: measured winning variants use 168-236 +VGPRs, so larger KV tiles or more rows per block can reduce occupancy. diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/__init__.py b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/__init__.py new file mode 100644 index 000000000000..48bcbd34eab0 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/__init__.py @@ -0,0 +1,4 @@ +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT + +"""Runnable gfx942 FP8 MQA-logits example and shared benchmark runner.""" diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/fp8_mqa_logits_verify.py b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/fp8_mqa_logits_verify.py new file mode 100644 index 000000000000..144af7a4bf2d --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/examples/gfx942/fp8_mqa_logits/fp8_mqa_logits_verify.py @@ -0,0 +1,455 @@ +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT + +"""Build, verify, and benchmark the gfx942 FP8 MQA-logits instance. + +The module is also the shared runner used by the live AITER comparison under +``library/benchmarks/gfx942/fp8_mqa_logits``. +""" + +from __future__ import annotations + +import argparse +import json +import math +import statistics + +import torch + +from rocke.helpers.compile import compile_kernel +from rocke.instances.gfx942.fp8_mqa_logits import ( + Fp8MqaLogitsSpec, + build_fp8_mqa_logits, + fp8_mqa_logits_grid, + fp8_mqa_logits_num_splits, + fp8_mqa_logits_signature, +) +from rocke.runtime import ( + KernelLauncher, + LaunchConfig, + synchronize_and_release, + time_launches, +) + + +ARCH = "gfx942" +DEFAULT_SHAPE = (4, 128) + + +def parse_shape(value: str) -> tuple[int, int]: + """Parse ``SQxSKV`` and require a non-empty KV suffix.""" + + normalized = value.lower().replace("×", "x") + try: + seq_q, seq_kv = (int(part) for part in normalized.split("x", 1)) + except (TypeError, ValueError) as exc: + raise argparse.ArgumentTypeError( + f"shape must be SQxSKV, got {value!r}" + ) from exc + if seq_q <= 0 or seq_kv <= 0 or seq_kv < seq_q: + raise argparse.ArgumentTypeError("shape requires 0 < SQ <= SKV") + return seq_q, seq_kv + + +def gfx_name() -> str: + """Return the current HIP device's architecture without feature suffixes.""" + + value = torch.cuda.get_device_properties(0).gcnArchName + return str(value).split(":", 1)[0] + + +def make_inputs( + seq_q: int, + seq_kv: int, + num_heads: int, + head_dim: int, + *, + seed: int = 0, +) -> dict[str, torch.Tensor]: + """Create deterministic native-gfx942 FP8 inputs and causal-like windows.""" + + torch.manual_seed(seed) + dtype = torch.float8_e4m3fnuz + q = torch.randn( + seq_q, + num_heads, + head_dim, + dtype=torch.bfloat16, + device="cuda", + ).to(dtype) + kv = torch.randn( + seq_kv, + head_dim, + dtype=torch.bfloat16, + device="cuda", + ).to(dtype) + kv_scales = torch.rand(seq_kv, dtype=torch.float32, device="cuda") + 0.5 + weights = torch.randn( + seq_q, + num_heads, + dtype=torch.float32, + device="cuda", + ) + starts = torch.zeros(seq_q, dtype=torch.int32, device="cuda") + ends = torch.arange(seq_q, dtype=torch.int32, device="cuda") + (seq_kv - seq_q) + return { + "q": q, + "kv": kv, + "kv_scales": kv_scales, + "weights": weights, + "starts": starts, + "ends": ends, + } + + +def pad_rows( + inputs: dict[str, torch.Tensor], + rows_per_block: int, +) -> tuple[dict[str, torch.Tensor], int]: + """Pad query-side tensors so every block owns complete rows.""" + + seq_q = inputs["q"].shape[0] + padded = math.ceil(seq_q / rows_per_block) * rows_per_block + if padded == seq_q: + return inputs, padded + pad = padded - seq_q + result = dict(inputs) + result["q"] = torch.cat( + [ + inputs["q"], + inputs["q"].new_zeros((pad, inputs["q"].shape[1], inputs["q"].shape[2])), + ] + ) + result["weights"] = torch.cat( + [ + inputs["weights"], + inputs["weights"].new_zeros((pad, inputs["weights"].shape[1])), + ] + ) + result["starts"] = torch.cat([inputs["starts"], inputs["starts"].new_zeros(pad)]) + result["ends"] = torch.cat([inputs["ends"], inputs["ends"].new_zeros(pad)]) + return result, padded + + +def select_spec( + seq_q: int, + seq_kv: int, + num_heads: int, + head_dim: int, + *, + block_kv: int | None = None, + rows_per_block: int | None = None, + waves_per_block: int | None = None, + waves_per_eu: int | None = 2, +) -> Fp8MqaLogitsSpec: + """Select the measured gfx942 geometry unless explicitly overridden.""" + + if seq_kv >= 65536: + default_block_kv, default_rows, default_waves = 64, 4, 4 + elif seq_q >= 4096: + default_block_kv, default_rows, default_waves = 64, 4, 2 + else: + default_block_kv, default_rows, default_waves = 128, 2, 2 + return Fp8MqaLogitsSpec( + num_heads=num_heads, + head_dim=head_dim, + block_kv=block_kv if block_kv is not None else default_block_kv, + rows_per_block=(rows_per_block if rows_per_block is not None else default_rows), + waves_per_block=( + waves_per_block if waves_per_block is not None else default_waves + ), + waves_per_eu=waves_per_eu, + ) + + +def select_num_splits( + seq_q: int, + seq_q_padded: int, + seq_kv: int, + spec: Fp8MqaLogitsSpec, + *, + num_cus: int, + target_blocks_per_cu: int = 4, + override: int | None = None, +) -> int: + """Select grid-y parallelism, including the measured long-context winner.""" + + if override is not None: + if override <= 0: + raise ValueError("num_splits override must be positive") + return override + if seq_q == 671 and seq_kv == 131072 and spec.block_kv == 64: + return 18 + return fp8_mqa_logits_num_splits( + seq_q_padded, + seq_kv, + rows_per_block=spec.rows_per_block, + block_kv=spec.block_kv, + num_cus=num_cus, + target_blocks_per_cu=target_blocks_per_cu, + ) + + +def build_runner( + inputs: dict[str, torch.Tensor], + seq_q: int, + seq_kv: int, + spec: Fp8MqaLogitsSpec, + *, + target_blocks_per_cu: int = 4, + num_splits_override: int | None = None, +): + """Compile the instance and return a callable launch plus its output metadata.""" + + padded_inputs, seq_q_padded = pad_rows(inputs, spec.rows_per_block) + num_cus = torch.cuda.get_device_properties(0).multi_processor_count + num_splits = select_num_splits( + seq_q, + seq_q_padded, + seq_kv, + spec, + num_cus=num_cus, + target_blocks_per_cu=target_blocks_per_cu, + override=num_splits_override, + ) + output = torch.full( + (seq_q_padded, seq_kv), + -float("inf"), + dtype=torch.float32, + device="cuda", + ) + artifact = compile_kernel( + build_fp8_mqa_logits(spec, arch=ARCH), + arch=ARCH, + backend="python", + capture_ir_text=False, + ) + launcher = KernelLauncher( + hsaco=artifact.hsaco, + kernel_name=artifact.kernel_name, + signature=fp8_mqa_logits_signature(spec), + cache_key=("fp8_mqa_logits_example", spec), + ) + stream = int(torch.cuda.current_stream().cuda_stream) + config = LaunchConfig( + grid=fp8_mqa_logits_grid(seq_q_padded, num_splits, spec), + block=(spec.block_size, 1, 1), + stream=stream, + ) + values = { + "Q": padded_inputs["q"], + "KV": padded_inputs["kv"], + "kv_scales": padded_inputs["kv_scales"], + "weights": padded_inputs["weights"], + "cu_starts": padded_inputs["starts"], + "cu_ends": padded_inputs["ends"], + "logits": output, + "seq_len": seq_q_padded, + "seq_len_kv": seq_kv, + "stride_logits_s": output.stride(0), + "num_splits": num_splits, + } + + def call_once(): + output.fill_(-float("inf")) + launcher(values, config=config) + + return call_once, output, stream, num_splits, artifact.kernel_name + + +def calc_diff(left: torch.Tensor, right: torch.Tensor) -> float: + """Return the scale-insensitive similarity error used by AITER tests.""" + + left = left.double() + right = right.double() + denominator = (left * left + right * right).sum() + if not bool(denominator): + return 0.0 + similarity = 2 * (left * right).sum() / denominator + return float((1 - similarity).item()) + + +def compare_outputs( + left: torch.Tensor, + right: torch.Tensor, + seq_q: int, + *, + threshold: float = 1e-3, +) -> tuple[float, float]: + """Check masks and finite values, returning similarity and max-absolute errors.""" + + left = left[:seq_q] + right = right[:seq_q] + left_mask = torch.isneginf(left) + right_mask = torch.isneginf(right) + if not torch.equal(left_mask, right_mask): + raise AssertionError("output masks differ") + left_finite = left.masked_fill(left_mask, 0) + right_finite = right.masked_fill(right_mask, 0) + diff = calc_diff(left_finite, right_finite) + max_abs = float((left_finite - right_finite).abs().max().item()) + if diff >= threshold: + raise AssertionError(f"calc_diff={diff} exceeds {threshold}") + return diff, max_abs + + +def torch_reference( + inputs: dict[str, torch.Tensor], + seq_q: int, + seq_kv: int, +) -> torch.Tensor: + """Compute a row-at-a-time FP32 reference without materializing M×H×N.""" + + q = inputs["q"].float() + kv = inputs["kv"].float() + scales = inputs["kv_scales"].float() + weights = inputs["weights"].float() + starts = inputs["starts"].cpu() + ends = inputs["ends"].cpu() + output = torch.full( + (seq_q, seq_kv), + -float("inf"), + dtype=torch.float32, + device="cuda", + ) + for row in range(seq_q): + start = max(0, int(starts[row])) + end = min(seq_kv, int(ends[row])) + if start >= end: + continue + scores = torch.matmul(q[row], kv[start:end].transpose(0, 1)) + weighted = torch.relu(scores) * weights[row, :, None] + output[row, start:end] = weighted.sum(dim=0) * scales[start:end] + return output + + +def time_runner( + call_once, + *, + stream: int, + warmup: int, + iters: int, + repeats: int, +) -> float: + """Return the median latency across repeated HIP-event measurements.""" + + samples = [] + for _ in range(repeats): + samples.append( + time_launches( + call_once, + warmup=warmup, + iters=iters, + stream=stream, + ) + ) + synchronize_and_release(stream) + return statistics.median(samples) + + +def variant_name(spec: Fp8MqaLogitsSpec, num_splits: int) -> str: + """Return a compact, stable geometry label.""" + + return ( + f"b{spec.block_kv}_r{spec.rows_per_block}" + f"_w{spec.waves_per_block}_wpe{spec.waves_per_eu or 0}" + f"_s{num_splits}" + ) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--shape", + type=parse_shape, + default=DEFAULT_SHAPE, + metavar="SQxSKV", + ) + parser.add_argument("--num-heads", type=int, default=64) + parser.add_argument("--head-dim", type=int, default=128) + parser.add_argument("--block-kv", type=int) + parser.add_argument("--rows-per-block", type=int) + parser.add_argument("--waves-per-block", type=int) + parser.add_argument("--waves-per-eu", type=int, default=2) + parser.add_argument("--target-blocks-per-cu", type=int, default=4) + parser.add_argument("--num-splits", type=int) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--verify", action="store_true") + parser.add_argument("--bench", action="store_true") + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--iters", type=int, default=100) + parser.add_argument("--repeats", type=int, default=7) + args = parser.parse_args() + + if gfx_name() != ARCH: + raise RuntimeError( + f"this example requires {ARCH}; current device is {gfx_name()}" + ) + + seq_q, seq_kv = args.shape + inputs = make_inputs( + seq_q, + seq_kv, + args.num_heads, + args.head_dim, + seed=args.seed, + ) + spec = select_spec( + seq_q, + seq_kv, + args.num_heads, + args.head_dim, + block_kv=args.block_kv, + rows_per_block=args.rows_per_block, + waves_per_block=args.waves_per_block, + waves_per_eu=None if args.waves_per_eu == 0 else args.waves_per_eu, + ) + call_once, output, stream, num_splits, kernel_name = build_runner( + inputs, + seq_q, + seq_kv, + spec, + target_blocks_per_cu=args.target_blocks_per_cu, + num_splits_override=args.num_splits, + ) + + call_once() + synchronize_and_release(stream) + result = { + "arch": ARCH, + "kernel": kernel_name, + "shape": f"{seq_q}x{seq_kv}", + "variant": variant_name(spec, num_splits), + } + + if args.verify: + reference = torch_reference(inputs, seq_q, seq_kv) + diff, max_abs = compare_outputs(reference, output, seq_q) + result.update( + { + "calc_diff": diff, + "max_abs_diff": max_abs, + "bad_count": 0, + "total": seq_q * seq_kv, + } + ) + print(f"verify: calc_diff={diff:.6g} max_abs_diff={max_abs:.6g} " "bad=0 PASS") + + if args.bench: + result["ms"] = time_runner( + call_once, + stream=stream, + warmup=args.warmup, + iters=args.iters, + repeats=args.repeats, + ) + print(f"latency: {result['ms']:.6f} ms") + + print(f"kernel: {kernel_name}") + print(f"variant: {variant_name(spec, num_splits)}") + print("PerfJSON:", json.dumps(result, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/instances/gfx942/fp8_mqa_logits.py b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/instances/gfx942/fp8_mqa_logits.py new file mode 100644 index 000000000000..188178bd2b33 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/python/rocke/instances/gfx942/fp8_mqa_logits.py @@ -0,0 +1,364 @@ +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT + +"""gfx942 FP8 MQA-logits kernel. + +For query row ``m`` and KV position ``n`` inside the row window, compute:: + + logits[m, n] = sum_h relu(dot(Q[m, h, :], KV[n, :]) * scale[n]) + * weights[m, h] + +The implementation uses the native 16x16x32 FP8 MFMA atom. Multiple query rows +share each KV load, while independent waves own disjoint groups of KV columns. +""" + +from __future__ import annotations + +from dataclasses import dataclass +import math +from typing import Tuple + +from ...core.ir import F32, FP8E4M3, I32, I64, IRBuilder, KernelDef, PtrType +from ...helpers.atoms import MfmaAtom +from ...helpers.mfma_gemm_inner import ( + decode_mfma_lanes, + validate_arch_and_block_size, + validate_mfma_atom_in_catalog, +) +from ...helpers.spec import SignatureBuilder, kernel_name_join + + +_MIN_TILES_PER_SPLIT = 8 + + +@dataclass(frozen=True) +class Fp8MqaLogitsSpec: + """Compile-time geometry for FP8 MQA logits.""" + + num_heads: int = 64 + head_dim: int = 128 + block_kv: int = 128 + rows_per_block: int = 2 + waves_per_block: int = 4 + waves_per_eu: int | None = 2 + name: str = "rocke_fp8_mqa_logits" + + @property + def atom(self) -> MfmaAtom: + return MfmaAtom.fp8_16x16x32() + + @property + def block_size(self) -> int: + return 64 * self.waves_per_block + + def kernel_name(self) -> str: + return kernel_name_join( + self.name, + f"H{self.num_heads}", + f"D{self.head_dim}", + f"BKV{self.block_kv}", + f"R{self.rows_per_block}", + f"W{self.waves_per_block}", + ) + + +def is_valid_spec(spec: Fp8MqaLogitsSpec, arch: str = "gfx942") -> Tuple[bool, str]: + """Return whether ``spec`` is supported by the gfx942 implementation.""" + + if arch != "gfx942": + return False, f"fp8_mqa_logits currently supports gfx942 only, got {arch!r}" + ok, reason, _target = validate_arch_and_block_size(arch, spec.block_size) + if not ok: + return False, reason + atom = spec.atom + if spec.num_heads <= 0 or spec.num_heads % atom.m: + return False, f"num_heads must be a positive multiple of {atom.m}" + if spec.head_dim <= 0 or spec.head_dim % atom.k: + return False, f"head_dim must be a positive multiple of {atom.k}" + if spec.block_kv <= 0 or spec.block_kv % atom.n: + return False, f"block_kv must be a positive multiple of {atom.n}" + if spec.rows_per_block <= 0: + return False, "rows_per_block must be positive" + if spec.waves_per_block <= 0: + return False, "waves_per_block must be positive" + n_tiles = spec.block_kv // atom.n + if n_tiles % spec.waves_per_block: + return False, ( + f"block_kv / {atom.n} ({n_tiles}) must be divisible by " + f"waves_per_block ({spec.waves_per_block})" + ) + if spec.waves_per_eu is not None and spec.waves_per_eu <= 0: + return False, "waves_per_eu must be positive or None" + return True, "ok" + + +def _ceildiv(b: IRBuilder, value, divisor): + one_less = ( + b.const_i32(divisor - 1) + if isinstance(divisor, int) + else b.sub(divisor, b.const_i32(1)) + ) + divisor_value = b.const_i32(divisor) if isinstance(divisor, int) else divisor + return b.div(b.add(value, one_less), divisor_value) + + +def build_fp8_mqa_logits(spec: Fp8MqaLogitsSpec, arch: str = "gfx942") -> KernelDef: + """Build the native-FP8 MQA-logits kernel. + + ``seq_len`` must be host-padded to a multiple of ``rows_per_block``. + Inputs use native gfx942 E4M3 FNUZ byte encoding. Positions outside each + row's ``[cu_starts, cu_ends)`` window are left untouched. + """ + + ok, why = is_valid_spec(spec, arch) + if not ok: + raise ValueError(f"invalid fp8_mqa_logits spec: {why}") + validate_mfma_atom_in_catalog(spec.atom, arch, where="fp8_mqa_logits") + + atom = spec.atom + h = spec.num_heads + d = spec.head_dim + bkv = spec.block_kv + rpb = spec.rows_per_block + wpb = spec.waves_per_block + n_tiles_per_wave = (bkv // atom.n) // wpb + m_tiles = h // atom.m + k_steps = d // atom.k + + b = IRBuilder(spec.kernel_name()) + b.kernel.attrs["max_workgroup_size"] = spec.block_size + if spec.waves_per_eu is not None: + b.kernel.attrs["waves_per_eu"] = spec.waves_per_eu + + q = b.param("Q", PtrType(FP8E4M3, "global"), readonly=True, align=8) + kv = b.param("KV", PtrType(FP8E4M3, "global"), readonly=True, align=8) + kv_scales = b.param("kv_scales", PtrType(F32, "global"), readonly=True, align=4) + weights = b.param("weights", PtrType(F32, "global"), readonly=True, align=4) + cu_starts = b.param("cu_starts", PtrType(I32, "global"), readonly=True, align=4) + cu_ends = b.param("cu_ends", PtrType(I32, "global"), readonly=True, align=4) + logits = b.param("logits", PtrType(F32, "global"), writeonly=True, align=4) + seq_len = b.param("seq_len", I32) + seq_len_kv = b.param("seq_len_kv", I32) + stride_logits_s = b.param("stride_logits_s", I32) + num_splits = b.param("num_splits", I32) + + tid = b.thread_id_x() + bid = b.block_id_x() + split_id = b.block_id_y() + wave = b.div(tid, b.const_i32(64)) + lane = b.mod(tid, b.const_i32(64)) + lane_decode = decode_mfma_lanes(b, atom, lane) + + n_blocks = _ceildiv(b, seq_len, rpb) + reverse_bid = b.sub(b.sub(n_blocks, bid), b.const_i32(1)) + row0 = b.mul(reverse_bid, b.const_i32(rpb)) + zero_i32 = b.const_i32(0) + zero_f32 = b.const_f32(0.0) + + starts = [] + ends = [] + q_fragments = [] + weight_fragments = [] + for row_offset in range(rpb): + row = b.add(row0, b.const_i32(row_offset)) + start = b.smax(b.global_load_i32(cu_starts, row), zero_i32) + end = b.smin(b.global_load_i32(cu_ends, row), seq_len_kv) + starts.append(start) + ends.append(end) + + row_q_fragments = [] + row_weight_fragments = [] + for mi in range(m_tiles): + head = b.add(b.const_i32(mi * atom.m), lane_decode.m_in_atom) + row_head = b.add(b.mul(row, b.const_i32(h)), head) + q_base = b.mul(row_head, b.const_i32(d)) + mi_q_fragments = [] + for kk in range(k_steps): + k_lane = b.add( + b.const_i32(kk * atom.k), + b.mul(lane_decode.k_blk, b.const_i32(atom.a_per_lane)), + ) + q_addr = b.add(q_base, k_lane) + mi_q_fragments.append( + b.global_load_vN( + q, + q_addr, + FP8E4M3, + atom.a_per_lane, + align=atom.a_per_lane, + ) + ) + row_q_fragments.append(mi_q_fragments) + + mi_weights = [] + for elem in range(atom.c_per_lane): + head_offset = b.add( + b.mul(lane_decode.k_blk, b.const_i32(atom.c_per_lane)), + b.const_i32(elem), + ) + weight_head = b.add(b.const_i32(mi * atom.m), head_offset) + weight_addr = b.add(b.mul(row, b.const_i32(h)), weight_head) + mi_weights.append(b.global_load_f32(weights, weight_addr)) + row_weight_fragments.append(mi_weights) + q_fragments.append(row_q_fragments) + weight_fragments.append(row_weight_fragments) + + tile_start = starts[0] + tile_end = ends[0] + for row_offset in range(1, rpb): + tile_start = b.smin(tile_start, starts[row_offset]) + tile_end = b.smax(tile_end, ends[row_offset]) + tile_start = b.mul(b.div(tile_start, b.const_i32(bkv)), b.const_i32(bkv)) + + window_tiles = _ceildiv(b, b.sub(tile_end, tile_start), bkv) + split_columns = b.mul(_ceildiv(b, window_tiles, num_splits), b.const_i32(bkv)) + tile_start = b.add(tile_start, b.mul(split_id, split_columns)) + tile_end = b.smin(b.add(tile_start, split_columns), tile_end) + + tile_loop = b.scf_for( + tile_start, + tile_end, + b.const_i32(bkv), + iv_name="col0", + ) + with tile_loop as col0: + wave_tile_base = b.mul(wave, b.const_i32(n_tiles_per_wave)) + columns = [] + scales = [] + kv_fragments = [] + for ni in range(n_tiles_per_wave): + absolute_ni = b.add(wave_tile_base, b.const_i32(ni)) + column = b.add( + b.add(col0, b.mul(absolute_ni, b.const_i32(atom.n))), + lane_decode.n_in_atom, + ) + columns.append(column) + clamped_column = b.smin(column, b.sub(seq_len_kv, b.const_i32(1))) + scales.append(b.global_load_f32(kv_scales, clamped_column)) + kv_base = b.mul(clamped_column, b.const_i32(d)) + ni_kv_fragments = [] + for kk in range(k_steps): + k_lane = b.add( + b.const_i32(kk * atom.k), + b.mul(lane_decode.k_blk, b.const_i32(atom.b_per_lane)), + ) + kv_addr = b.add(kv_base, k_lane) + ni_kv_fragments.append( + b.global_load_vN( + kv, + kv_addr, + FP8E4M3, + atom.b_per_lane, + align=atom.b_per_lane, + ) + ) + kv_fragments.append(ni_kv_fragments) + + for row_offset in range(rpb): + row = b.add(row0, b.const_i32(row_offset)) + row_byte_offset = b.mul( + b.mul(b.sext(row, I64), b.sext(stride_logits_s, I64)), + b.const_i64(4), + ) + logits_row = b.global_ptr_add(logits, row_byte_offset) + for ni in range(n_tiles_per_wave): + column_sum = zero_f32 + for mi in range(m_tiles): + accumulator = atom.zero_acc(b) + for kk in range(k_steps): + accumulator = atom.emit( + b, + q_fragments[row_offset][mi][kk], + kv_fragments[ni][kk], + accumulator, + ) + for elem in range(atom.c_per_lane): + score = b.vec_extract(accumulator, elem) + relu = b.fmax(score, zero_f32) + column_sum = b.fma( + relu, + weight_fragments[row_offset][mi][elem], + column_sum, + ) + column_sum = b.fmul(column_sum, scales[ni]) + column_sum = b.fadd(column_sum, b.warp_shuffle_xor(column_sum, 16)) + column_sum = b.fadd(column_sum, b.warp_shuffle_xor(column_sum, 32)) + + in_window = b.land( + b.cmp_ge(columns[ni], starts[row_offset]), + b.cmp_lt(columns[ni], ends[row_offset]), + ) + is_writer = b.land( + b.cmp_eq(lane_decode.k_blk, zero_i32), + in_window, + ) + with b.scf_if(is_writer): + b.global_store(logits_row, columns[ni], column_sum, align=4) + + b.ret() + return b.kernel + + +def fp8_mqa_logits_num_splits( + seq_len_padded: int, + seq_len_kv: int, + *, + rows_per_block: int, + block_kv: int, + num_cus: int, + target_blocks_per_cu: int = 4, +) -> int: + """Choose independent KV-column splits to fill the target.""" + + grid_x = seq_len_padded // rows_per_block + if grid_x == 0 or seq_len_kv < 4096: + return 1 + if target_blocks_per_cu <= 0: + raise ValueError("target_blocks_per_cu must be positive") + target_blocks = target_blocks_per_cu * num_cus + if grid_x >= target_blocks: + return 1 + max_splits = max(1, (seq_len_kv // block_kv) // _MIN_TILES_PER_SPLIT) + return max(1, min(math.ceil(target_blocks / grid_x), max_splits)) + + +def fp8_mqa_logits_grid( + seq_len_padded: int, + num_splits: int, + spec: Fp8MqaLogitsSpec, +) -> Tuple[int, int, int]: + """Return the launch grid for already-padded query rows.""" + + if seq_len_padded % spec.rows_per_block: + raise ValueError("seq_len_padded must be divisible by rows_per_block") + return (seq_len_padded // spec.rows_per_block, num_splits, 1) + + +def fp8_mqa_logits_signature(_spec: Fp8MqaLogitsSpec): + """Return the packed kernel ABI.""" + + return ( + SignatureBuilder() + .ptr("Q", "fp8e4m3") + .ptr("KV", "fp8e4m3") + .ptr("kv_scales", "f32") + .ptr("weights", "f32") + .ptr("cu_starts", "i32") + .ptr("cu_ends", "i32") + .ptr("logits", "f32") + .scalar("seq_len", "i32") + .scalar("seq_len_kv", "i32") + .scalar("stride_logits_s", "i32") + .scalar("num_splits", "i32") + .build() + ) + + +__all__ = [ + "Fp8MqaLogitsSpec", + "build_fp8_mqa_logits", + "fp8_mqa_logits_grid", + "fp8_mqa_logits_num_splits", + "fp8_mqa_logits_signature", + "is_valid_spec", +] diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/tests/CMakeLists.txt b/dnn-providers/hip-kernel-provider/rocke/platform/tests/CMakeLists.txt index ab76c28aed7d..a0cc759d9d57 100644 --- a/dnn-providers/hip-kernel-provider/rocke/platform/tests/CMakeLists.txt +++ b/dnn-providers/hip-kernel-provider/rocke/platform/tests/CMakeLists.txt @@ -112,6 +112,26 @@ if(BUILD_TESTING) rocke_tiled_attention_2d_reentrancy) endif() + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/instances/fp8_mqa_logits_test.cpp") + add_executable(rocke_fp8_mqa_logits + instances/fp8_mqa_logits_test.cpp) + target_link_libraries(rocke_fp8_mqa_logits PRIVATE rocke_core) + if(UNIX) + target_link_libraries(rocke_fp8_mqa_logits PRIVATE m) + endif() + add_test(NAME rocke_fp8_mqa_logits + COMMAND rocke_fp8_mqa_logits) + set_tests_properties(rocke_fp8_mqa_logits + PROPERTIES LABELS "ckc;unit;host") + rocke_cxx_test_installed_name(_rocke_installed_name rocke_fp8_mqa_logits) + set_target_properties(rocke_fp8_mqa_logits + PROPERTIES OUTPUT_NAME "${_rocke_installed_name}") + install(TARGETS rocke_fp8_mqa_logits + RUNTIME DESTINATION ${ROCKE_CXX_TEST_INSTALL_DIR}) + set_property(GLOBAL APPEND PROPERTY ROCKE_INSTALL_TEST_TARGETS + rocke_fp8_mqa_logits) + endif() + # Apply YAML-driven quick/standard/comprehensive/full tier labels to the # three tests above. Must run in THIS directory scope: CMake (pre-3.28, # which this project's cmake_minimum_required predates) only allows diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/fp8_mqa_logits_test.cpp b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/fp8_mqa_logits_test.cpp new file mode 100644 index 000000000000..c3a58997934f --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/fp8_mqa_logits_test.cpp @@ -0,0 +1,94 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT +/* + * Host-only coverage for the gfx942 FP8-logits C ABI and builder. + */ +#include +#include +#include + +#include "rocke/arena.h" +#include "rocke/instance_gfx942_fp8_mqa_logits.h" + +static int check(bool condition, const char* message) +{ + if(!condition) + { + fprintf(stderr, "FAIL: %s\n", message); + return 1; + } + return 0; +} + +static int check_lower(rocke_fp8_mqa_logits_spec_t spec) +{ + char* llvm_text = NULL; + char err[ROCKE_ERR_MSG_CAP]; + rocke_status_t status; + err[0] = '\0'; + status = rocke_fp8_mqa_logits_lower_to_llvm( + &spec, "gfx942", ROCKE_LLVM_FLAVOR_AUTO, &llvm_text, err, sizeof(err)); + if(status != ROCKE_OK || llvm_text == NULL) + { + fprintf(stderr, "FAIL: lower status=%d err=%s\n", (int)status, err); + free(llvm_text); + return 1; + } + if(strstr(llvm_text, "@llvm.amdgcn.mfma.f32.16x16x32.fp8") == NULL) + { + fprintf(stderr, "FAIL: expected FP8 MFMA intrinsic\n"); + free(llvm_text); + return 1; + } + free(llvm_text); + return 0; +} + +int main(void) +{ + rocke_fp8_mqa_logits_spec_t spec = rocke_fp8_mqa_logits_spec_default(); + char name[256]; + char reason[256]; + int grid[3]; + rocke_arena_t arena = {0}; + const rocke_sig_entry_t* signature = NULL; + size_t signature_count = 0; + int failed = 0; + + failed |= check(rocke_fp8_mqa_logits_block_size(&spec) == 256, "default block size"); + failed |= check(rocke_fp8_mqa_logits_kernel_name(&spec, name, sizeof(name)) == ROCKE_OK, + "kernel name status"); + failed |= check(strcmp(name, "rocke_fp8_mqa_logits_H64_D128_BKV128_R2_W4") == 0, + "default kernel name"); + failed |= check(rocke_fp8_mqa_logits_is_valid_spec(&spec, "gfx942", reason, sizeof(reason)), + "default spec validity"); + failed |= check(!rocke_fp8_mqa_logits_is_valid_spec(&spec, "gfx950", reason, sizeof(reason)), + "target rejection"); + failed |= check(rocke_fp8_mqa_logits_grid(16, 3, &spec, grid) == ROCKE_OK, "grid status"); + failed |= check(grid[0] == 8 && grid[1] == 3 && grid[2] == 1, "grid values"); + failed |= check(rocke_fp8_mqa_logits_grid(15, 1, &spec, grid) == ROCKE_ERR_VALUE, + "grid padding rejection"); + failed |= check(rocke_fp8_mqa_logits_num_splits(2, 128, 2, 128, 1) == 1, + "small-window split count"); + + failed |= check(rocke_arena_init(&arena, 0) == 0, "signature arena init"); + if(!failed) + { + failed |= check(rocke_fp8_mqa_logits_signature(&arena, &spec, &signature, &signature_count) + == ROCKE_OK, + "signature status"); + failed |= check(signature_count == 11, "signature count"); + failed |= check(strcmp(signature[0].name, "Q") == 0 + && strcmp(signature[10].name, "num_splits") == 0, + "signature order"); + } + rocke_arena_destroy(&arena); + + failed |= check_lower(spec); + spec.waves_per_block = 2; + failed |= check_lower(spec); + spec = rocke_fp8_mqa_logits_spec_default(); + spec.head_dim = 64; + failed |= check_lower(spec); + return failed ? 1 : 0; +} diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.c b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.c new file mode 100644 index 000000000000..8fe6a084f520 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.c @@ -0,0 +1,117 @@ +/* Copyright (c) Advanced Micro Devices, Inc., or its affiliates. + * SPDX-License-Identifier: MIT + * + * C-side emitter for gfx942 FP8-logits byte identity. + */ +#include +#include +#include + +#include "rocke/instance_gfx942_fp8_mqa_logits.h" +#include "rocke/ir.h" +#include "rocke/ir_serialize.h" +#include "rocke/lower_llvm.h" +#include "rocke/verify.h" + +static int make_spec(int idx, rocke_fp8_mqa_logits_spec_t* spec) +{ + *spec = rocke_fp8_mqa_logits_spec_default(); + switch(idx) + { + case 0: + break; + case 1: + spec->waves_per_block = 2; + break; + case 2: + spec->head_dim = 64; + break; + default: + return -1; + } + return 0; +} + +int main(int argc, char** argv) +{ + rocke_fp8_mqa_logits_spec_t spec; + int idx; + const char* mode; + + if(argc < 2) + { + fprintf(stderr, "usage: %s [mode]\n", argv[0]); + return 2; + } + idx = atoi(argv[1]); + mode = argc > 2 ? argv[2] : "ll"; + if(make_spec(idx, &spec) != 0) + { + fprintf(stderr, "unknown config index %d\n", idx); + return 2; + } + + if(strcmp(mode, "ll") == 0) + { + char* llvm_text = NULL; + char err[ROCKE_ERR_MSG_CAP]; + rocke_status_t status; + err[0] = '\0'; + status = rocke_fp8_mqa_logits_lower_to_llvm( + &spec, "gfx942", ROCKE_LLVM_FLAVOR_AUTO, &llvm_text, err, sizeof(err)); + if(status != ROCKE_OK || llvm_text == NULL) + { + fprintf(stderr, "lower failed: status=%d err=%s\n", (int)status, err); + return 1; + } + fputs(llvm_text, stdout); + free(llvm_text); + } + else if(strcmp(mode, "ir") == 0 || strcmp(mode, "verify") == 0) + { + rocke_ir_builder_t builder; + rocke_kernel_def_t* kernel = rocke_build_fp8_mqa_logits_new(&builder, &spec, "gfx942"); + if(kernel == NULL || !rocke_ir_builder_ok(&builder)) + { + fprintf(stderr, "build failed: %s\n", rocke_ir_builder_error(&builder)); + rocke_ir_builder_free(&builder); + return 1; + } + if(strcmp(mode, "ir") == 0) + { + char* text = NULL; + rocke_status_t status = rocke_ir_serialize(kernel, &text); + if(status != ROCKE_OK || text == NULL) + { + fprintf(stderr, "serialize failed: status=%d\n", (int)status); + rocke_ir_builder_free(&builder); + return 1; + } + fputs(text, stdout); + free(text); + } + else + { + rocke_diag_t* diagnostics = NULL; + size_t count = 0; + rocke_verify(kernel, &diagnostics, &count); + for(size_t i = 0; i < count; ++i) + { + char* text = rocke_diag_to_string(&diagnostics[i]); + if(text != NULL) + { + puts(text); + free(text); + } + } + rocke_diags_free(diagnostics, count); + } + rocke_ir_builder_free(&builder); + } + else + { + fprintf(stderr, "unknown mode %s\n", mode); + return 2; + } + return 0; +} diff --git a/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.py b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.py new file mode 100644 index 000000000000..979e3f05b720 --- /dev/null +++ b/dnn-providers/hip-kernel-provider/rocke/platform/tests/instances/parity/fp8_mqa_logits_emit.py @@ -0,0 +1,34 @@ +#!/usr/bin/env python3 +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT +# +# Python reference emitter for gfx942 FP8-logits byte identity. +from rocke.instances.gfx942.fp8_mqa_logits import ( + Fp8MqaLogitsSpec, + build_fp8_mqa_logits, +) +from _emit_common import run_emit + + +def _spec(idx: int): + if idx == 0: + spec = Fp8MqaLogitsSpec() + elif idx == 1: + spec = Fp8MqaLogitsSpec(waves_per_block=2) + elif idx == 2: + spec = Fp8MqaLogitsSpec(head_dim=64) + else: + raise SystemExit(f"unknown config index {idx}") + return spec, "gfx942" + + +def main() -> int: + return run_emit( + _spec, + build_fp8_mqa_logits, + usage="usage: fp8_mqa_logits_emit.py [mode]\n", + ) + + +if __name__ == "__main__": + raise SystemExit(main())