Skip to content
Open
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
91 changes: 75 additions & 16 deletions rocm/scripts/hipblaslt_patch_setup.sh
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,12 @@
# * 7.13.0 -- family D (per-arch lib subdir gfx942/; CPX
# WGMXCC=4 defect already fixed upstream so the
# CPX rewrite step is a sentinel-only no-op there)
# * 7.14.0 -- family D, but the Tensile leaves are now stored
# zlib-compressed (*.dat.zlib) in gfx942/. fp16
# FORWARD leaf ONLY: route the 120 small-N shapes to
# an in-pool kernel (solution 243113). fp32/bf16 pass
# stock, so no SPX backward/CU228 rows and no CPX
# WGMXCC rewrite here (proven on MI300A, 3 Sep 2026).
#

set -euo pipefail
Expand Down Expand Up @@ -143,7 +149,7 @@ done

# ── version gate ────────────────────────────────────────────────────
case "${ROCM_VERSION}" in
7.1.0|7.1.1|7.2.0|7.2.2|7.2.3|7.2.4|7.13.0) ;;
7.1.0|7.1.1|7.2.0|7.2.2|7.2.3|7.2.4|7.13.0|7.14.0) ;;
*)
echo "[hipblaslt_patch] no fix vendored for rocm/${ROCM_VERSION}; exiting NOOP (rc=${NOOP_RC})"
exit ${NOOP_RC}
Expand All @@ -168,7 +174,11 @@ MODULE_FILE="${MODULE_PATH}/rocmplus-${ROCM_VERSION}/hipblaslt/patched.lua"
# layout the SDK uses so HIPBLASLT_TENSILE_LIBPATH=${DST_LIBDIR}
# resolves the same way after the overlay is selected.
ARCH_SUBDIR=""
if [ -d "${SRC_LIBDIR}/gfx942" ] && compgen -G "${SRC_LIBDIR}/gfx942/*.dat" >/dev/null; then
if [ -d "${SRC_LIBDIR}/gfx942" ] \
&& { compgen -G "${SRC_LIBDIR}/gfx942/*.dat" >/dev/null \
|| compgen -G "${SRC_LIBDIR}/gfx942/*.dat.zlib" >/dev/null; }; then
# 7.13.0 ships gfx942/*.dat (plain); 7.14.0 ships gfx942/*.dat.zlib
# (zlib-compressed). Either one selects the per-arch layout.
ARCH_SUBDIR="gfx942"
fi
EFFECTIVE_SRC_LIBDIR="${SRC_LIBDIR}${ARCH_SUBDIR:+/${ARCH_SUBDIR}}"
Expand Down Expand Up @@ -235,6 +245,7 @@ import argparse
import os
import pathlib
import sys
import zlib

import msgpack

Expand Down Expand Up @@ -342,9 +353,37 @@ SOLUTION_INDEX = {
"7.13.0":{("HH", "Alik_Bljk_Cijk_Dijk_gfx942.dat"): 236740,
("HH", "Ailk_Bljk_Cijk_Dijk_CU228_gfx942.dat"): 225772,
("HH", "Ailk_Bjlk_Cijk_Dijk_CU228_gfx942.dat"): 222586},
# rocm/7.14.0: leaves are zlib-compressed (*.dat.zlib). Only the fp16
# FORWARD leaf regressed on MI300A (120 small-N shapes returned 0
# solutions); fp32 and bf16 pass stock. Route those 120 shapes to an
# existing in-pool kernel -- solution 243113,
# MT16x16x256/MI16x16x1/WGMXCC1/GSU1 (family-D forward signature) --
# so no kernel binary is built. No CU228 backward rows (fp16 backward
# was not among the misses) and no CPX WGMXCC rewrite (defect fixed
# upstream). Empirically proven FAIL->PASS, 120->0 misses, with no
# fp32/bf16 regression.
"7.14.0":{("HH", "Alik_Bljk_Cijk_Dijk_gfx942.dat"): 243113},
}
CPX_SENTINEL = "__cpx_patch_v1__"

# Versions whose Tensile leaves are stored zlib-compressed (*.dat.zlib).
# The patcher appends ".zlib" to the on-disk name and transparently
# decompresses on read / recompresses (level 9) on write for these.
ZLIB_VERSIONS = {"7.14.0"}

# Versions that get NO CPX WorkGroupMappingXCC rewrite (forward-leaf-only
# overlays, or versions where the CU38 WGMXCC=4 defect is fixed upstream).
NO_CPX_VERSIONS = {"7.14.0"}


def _read_lib(path):
"""Read a Tensile library leaf, transparently zlib-decompressing when the
on-disk name ends in .zlib (7.14.0+). Returns the unpacked msgpack tree."""
raw = open(path, "rb").read()
if str(path).endswith(".zlib"):
raw = zlib.decompress(raw)
return msgpack.unpackb(raw, raw=False, strict_map_key=False)


def equality_matching(data):
for r in data["library"]["rows"]:
Expand All @@ -366,20 +405,26 @@ def _patch_one_row(table, key, idx):
return "added"


def patch_spx(src, dst, shapes_for_file, idx):
"""Patch one SPX .dat with all `shapes_for_file` rows in one pass.
def patch_spx(src, dst, shapes_for_file, idx, zlib_mode=False):
"""Patch one SPX .dat[.zlib] with all `shapes_for_file` rows in one pass.

A single read+write per file is required when multiple rows share a
filename, otherwise the second write would overwrite the first.
"""
data = msgpack.unpack(open(src, "rb"), raw=False, strict_map_key=False)
data = _read_lib(src)
table = equality_matching(data)["table"]
actions = []
for M, N, B, K in shapes_for_file:
actions.append((M, N, B, K, _patch_one_row(table, [M, N, B, K], idx)))
dst.parent.mkdir(parents=True, exist_ok=True)
with open(dst, "wb") as f:
msgpack.pack(data, f)
if zlib_mode:
# 7.14.0: recompress the leaf to *.dat.zlib. use_bin_type=True and
# zlib level 9 match the proven overlay encoding that hipBLASLt loads.
with open(dst, "wb") as f:
f.write(zlib.compress(msgpack.packb(data, use_bin_type=True), 9))
else:
with open(dst, "wb") as f:
msgpack.pack(data, f)
print(f" [spx] {dst.name}: {len(actions)} row(s) -> solution {idx}")
for M, N, B, K, what in actions:
print(f" {what:7s} [{M}, {N}, {B}, {K}]")
Expand Down Expand Up @@ -446,27 +491,39 @@ def main():
sdir = pathlib.Path(args.src_libdir)
ddir = pathlib.Path(args.dst_libdir)
idx_map = SOLUTION_INDEX[args.rocm_version]
# 7.14.0+ store the leaves zlib-compressed: the on-disk name gains a
# ".zlib" suffix while the SOLUTION_INDEX keys stay on the plain ".dat"
# name. patch_spx recompresses; _read_lib decompresses transparently.
zlib_mode = args.rocm_version in ZLIB_VERSIONS
disk = (lambda s: s + ".zlib") if zlib_mode else (lambda s: s)
# bf16 (BB / BB_HA) is opt-in per version: it runs only when this
# version carries a BB solution index. Versions without one get the
# fp16-only (HH) behaviour, unchanged.
bf16_enabled = any(tag == "BB" for (tag, _s) in idx_map)
# Group shapes by (dtype_tag, suffix): each file is read+written
# once with all matching rows applied. Skip any (tag, suffix) this
# version has no index for (auto-scopes BB rows to bf16 versions).
# version has no index for (auto-scopes BB rows to bf16 versions, and
# scopes 7.14.0 to the single forward leaf it lists).
by_file = {}
for tag, suffix, M, N, B, K in SPX_SHAPES:
by_file.setdefault((tag, suffix), []).append((M, N, B, K))
for (tag, suffix), shapes in by_file.items():
if (tag, suffix) not in idx_map:
continue
prefix = DAT_PREFIXES[tag]
patch_spx(sdir / (prefix + suffix), ddir / (prefix + suffix),
shapes, idx_map[(tag, suffix)])
for tag, suffix in CPX_DATS:
if tag.startswith("BB") and not bf16_enabled:
continue
prefix = DAT_PREFIXES[tag]
patch_cpx(sdir / (prefix + suffix), ddir / (prefix + suffix))
patch_spx(sdir / (prefix + disk(suffix)), ddir / (prefix + disk(suffix)),
shapes, idx_map[(tag, suffix)], zlib_mode=zlib_mode)
# CPX WorkGroupMappingXCC rewrite: skipped for NO_CPX_VERSIONS (7.14.0's
# CU38 defect is fixed upstream and the overlay is forward-leaf-only).
if args.rocm_version in NO_CPX_VERSIONS:
print(f" [cpx] skipped for rocm/{args.rocm_version} "
f"(forward-leaf-only overlay; CU38 WGMXCC defect fixed upstream)")
else:
for tag, suffix in CPX_DATS:
if tag.startswith("BB") and not bf16_enabled:
continue
prefix = DAT_PREFIXES[tag]
patch_cpx(sdir / (prefix + disk(suffix)), ddir / (prefix + disk(suffix)))


if __name__ == "__main__":
Expand All @@ -481,7 +538,9 @@ ${SUDO} mkdir -p "${EFFECTIVE_DST_LIBDIR}"
# section correct as coverage grows: fp16-only versions yield the 6 HH
# files; bf16-enabled versions (7.2.4) additionally yield the BB SPX
# leaf + 3 BB_HA CU38 leaves.
mapfile -t PATCHED_FILES < <(cd "${WORKDIR}" && ls -1 *.dat 2>/dev/null)
# WORKDIR holds only what the inline patcher wrote: plain *.dat (<=7.13.0)
# or zlib-compressed *.dat.zlib (7.14.0+). Match either.
mapfile -t PATCHED_FILES < <(cd "${WORKDIR}" && ls -1 2>/dev/null | grep -E '\.dat(\.zlib)?$')
[ "${#PATCHED_FILES[@]}" -gt 0 ] \
|| send-error "inline patcher produced no .dat files in ${WORKDIR}"
for fname in "${PATCHED_FILES[@]}"; do
Expand Down