Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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())
Original file line number Diff line number Diff line change
@@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -1023,6 +1024,64 @@ std::vector<std::string> 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<std::string>& 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<int>(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<std::string> 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<std::string> 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<std::string>& store)
Expand Down Expand Up @@ -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");

Expand Down
Original file line number Diff line number Diff line change
@@ -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 <stdbool.h>
#include <stddef.h>

#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 */
Loading
Loading