Skip to content

[blas][rocblas] Support int8 inputs with float output in gemm_batch - #763

Open
zjin-lcf wants to merge 6 commits into
uxlfoundation:developfrom
zjin-lcf:feature/rocblas-int8-gemm-batch
Open

zjin-lcf wants to merge 6 commits into
uxlfoundation:developfrom
zjin-lcf:feature/rocblas-int8-gemm-batch

Conversation

@zjin-lcf

@zjin-lcf zjin-lcf commented Aug 15, 2026 •

Copy link
Copy Markdown
Contributor

Summary

rocBLAS reaches int8 inputs only with an int32 output and compute type, which takes int32
alpha and beta, whereas oneMath specifies a float output and float scalars for this
combination. This accumulates the products exactly in an int32 workspace and applies the
float scalars in a scaling kernel afterwards, covering the buffer strided, USM strided and
USM grouped entry points in both layouts.

  • k is bounded so the int32 accumulator cannot wrap. Every int8 magnitude is at most 128,
    so 128 * 128 bounds a product and k above INT32_MAX / (128 * 128) (k > 131071)
    reports unimplemented. The cuBLAS backend accumulates this combination in float and has
    no such ceiling, so the same call succeeds on NVIDIA and throws on AMD past that limit.
    This is documented in docs/domains/blas.rst.
  • Sizes whose workspace or kernel range would overflow are rejected before anything is
    allocated, using checked multiplication and addition rather than a check on the product.
  • An ldc below the row count is rejected. rocBLAS rejects such a call itself, but only
    from its host task, which does not hold back the separately submitted scaling kernel, and
    the kernel would fold one column of C onto the next.
  • Grouped scaling stays a single kernel submission regardless of group or batch count, by
    locating the group that owns each entry with a binary search over per-group metadata.
  • The scaling kernels launch over a flat range, since HIP cannot map every large
    multi-dimensional range onto its grid. The added Int8Int8SinglePrecisionLargePrimeRange
    regression holds a large prime in n to cover that.

#761 is already on develop; this branch is merged onto current develop rather than
rebased, so the remaining diff is the rocBLAS backend, the large-prime USM regression, the
BLAS backend-limitation note, and a tighter integer-accumulation error model on AMD.

Verification

Built with DPC++ against ROCm 7.1.1 and run on an AMD Instinct MI210 (gfx90a).

The GEMM batch suite excluding int8 (the DPC++ ProgramManager kernel-map assertion when
launching SYCL kernels from a backend .so) passes in both dispatch modes:

$ ./bin/test_main_blas_ct --gtest_filter='*GemmBatch*-*Int8*:*GemmBatchUsmTests.Complex*'
[  PASSED  ] 32 tests.
$ ./bin/test_main_blas_rt --gtest_filter='*GemmBatch*-*Int8*:*GemmBatchUsmTests.Complex*'
[  PASSED  ] 32 tests.

Int8 coverage is from a standalone program compiled together with rocblas_batch.cpp so
the scaling kernels live in the same image:

strided alpha and beta set   PASS
strided beta zero            PASS
strided alpha zero           PASS
strided both zero            PASS
k above 131071 unimplemented PASS
large prime n flat range     PASS
bench 512^3 batch=4          7.52 ms  143 GFLOP/s

Made with Cursor

zjin-lcf and others added 3 commits August 14, 2026 14:44
cuBLAS reaches this combination through cublasGemmStridedBatchedEx and
cublasGemmBatchedEx, which already accept the datatypes the existing
launchers forward, so the column-major buffer, USM strided and USM group
entry points only needed routing to the implementation instead of
throwing unimplemented. The int32 output combination stays unimplemented
because cuBLAS produces it only under CUBLAS_COMPUTE_32I, which takes
int32 alpha and beta, whereas oneMath specifies float scalars.

A float output accumulated from int8 inputs is rounded at the magnitude
of the terms summed rather than at the magnitude of the output entry, so
an entry whose sum cancels cannot meet any relative bound. The shared
checker takes an optional absolute tolerance, defaulted to zero so that
existing callers are unaffected, and the int8-to-float gemm_batch tests
pass eps times k * 128 * 128, an upper bound on the accumulated
magnitude sum|a*b|. Int8Int8SinglePrecisionErrorModel covers that path
with fixed data whose leading rows and columns cancel exactly.

Co-authored-by: Cursor <cursoragent@cursor.com>
…del bound

An entry whose terms and stored C value are all zero gives a zero model
bound, which the reported usage ratio would divide by.

Co-authored-by: Cursor <cursoragent@cursor.com>
rocBLAS reaches int8 inputs only with an int32 output and compute type,
which takes int32 alpha and beta, whereas oneMath specifies a float
output and float scalars for this combination. Accumulate the products
exactly in an int32 workspace and apply the float scalars in a scaling
kernel afterwards, for the buffer strided, USM strided and USM grouped
entry points in both layouts.

Bound k so the int32 accumulator cannot wrap, reject sizes whose
workspace or kernel range would overflow, and reject an ldc below the
row count, which the scaling kernel would otherwise fold onto the next
column. Grouped scaling stays a single kernel by locating the group that
owns each entry with a binary search over per-group metadata.

The scaling kernels launch over a flat range, since HIP cannot map every
large multi-dimensional range onto its grid, and the added regression
holds a large prime in n to cover that.

Co-authored-by: Cursor <cursoragent@cursor.com>

@melonakos melonakos left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The rocBLAS work here is better than I expected — it's a genuinely different and harder problem than the cuBLAS side, and you handled the tricky parts correctly. Holding approval only on the rebase, plus one cross-backend consequence that needs deciding.

First: please rebase onto #761

This PR carries #761's cublas_batch.cpp change (+5/−3) and the same test-harness edits verbatim, on top of the rocBLAS work. I've approved #761, so once it lands, rebasing should reduce this to just the rocBLAS delta. Reviewing them as independent PRs means reviewing the same cuBLAS change and the same tolerance machinery twice.

The rocBLAS strategy is fundamentally different, and you documented why

rocBLAS reaches int8 inputs only with an int32 output and compute type, so the int8-to-float combination accumulates exactly in int32 and applies oneMath's float alpha and beta afterwards.

That's the crux, and it's worth spelling out for anyone else reading this: cuBLAS does int8→float natively with an fp32 execution type, so #761 is five lines. rocBLAS can't, so here the gemm runs into an int32 workspace and a separate kernel applies alpha and beta to produce the float output. Hence 400 lines instead of 5.

An interesting side effect: the rocBLAS path is more accurate than the cuBLAS one. Integer accumulation is exact, so the only rounding is the final scaling, whereas the cuBLAS path rounds throughout the fp32 reduction.

What I verified

beta == 0 is handled correctly, which is the classic bug in this shape:

float result = 0.0f;
if (alpha != 0.0f) result = alpha * static_cast<float>(accum[offset]);
if (beta != 0.0f)  result += beta * c[offset];
c[offset] = result;

Skipping the read of C entirely when beta == 0 — rather than computing beta * c[offset] and relying on 0 * x == 0 — is what makes beta = 0 work when the caller passes an uninitialised or NaN-filled C, as BLAS explicitly permits. Getting this wrong produces NaNs that are maddening to trace. You also handle alpha == 0 and the both-zero case correctly.

The workspace lifetime is right, and better than the pattern in your LAPACK PRs:

queue.submit([&](sycl::handler& cgh) {
    cgh.depends_on(done);
    cgh.host_task([=]() { sycl::free(accum, queue); });
});

The free is deferred behind the scaling kernel rather than being sequenced by a blocking wait, so the entry point stays asynchronous. That's exactly what #767's sycl::free(ipiv32, queue) should be doing instead of leaning on lapack_info_check_batch's internal queue.wait(). Same author, better pattern here — worth propagating.

Allocation failure is checked, with device_bad_alloc rather than a null dereference, and std::max<std::size_t>(accum_size_size, 1) avoids a zero-size malloc_device. The overflow-checked size arithmetic (checked_int8_float_product/_sum) is thorough.

The ldc < rows check earns its comment:

rocBLAS rejects such a call itself, but only from its host task, which does not hold back the kernel.

That's a sharp observation. The validation genuinely has to happen on the host before submission, because rocBLAS's own rejection happens too late to stop your scaling kernel from reading past the workspace. Not obvious, and correctly reasoned.

Each work-item touches one distinct offset, so the read-then-write of c[offset] in the scaling kernel is race-free.

Needs a decision: the k limit is a user-visible backend difference

constexpr int64_t max_product = 128 * 128;
constexpr int64_t max_safe_k = std::numeric_limits<std::int32_t>::max() / max_product;

That's k ≤ 131071, above which you throw unimplemented. The bound is correct — |a·b| ≤ 128² since int8 reaches −128 — and throwing is far better than silently wrapping the accumulator.

But the consequence is that the same oneMath call succeeds on NVIDIA and throws on AMD once k exceeds 131071. cuBLAS accumulates in fp32 and has no such ceiling. A user who develops against CUDA and deploys on MI300 finds out at runtime.

That needs to be documented wherever backend limitations are listed, not only in the exception text. It's also worth a sentence in the PR description so it's visible to whoever reviews this second.

Consequence for #761's test tolerance

Worth thinking about together, since these two PRs share the test harness. The abs_bound model in #761 — eps × k × 128² — is derived from fp32 accumulation, which is the cuBLAS behaviour. The rocBLAS path accumulates exactly, so its only error is the final alpha scaling: dramatically tighter than the model allows.

So on AMD that tolerance is a very loose upper bound, and a genuine kernel error there could pass comfortably inside it. Since you went to the trouble of building an error model that reports how much of the budget is consumed, it'd be a shame for the AMD path to sit at 1% of budget and never catch anything. Either derive the bound per-backend, or at minimum note in the test that it's the fp32-path bound and the integer path should be checked far more tightly.

Minor

  • When alpha == 0 you set accumulate = 0 but still run the full batched gemm, whose result is then discarded. An early-out that only scales C by beta would skip the multiply entirely. Small, and the case is rare.
  • The helper names (checked_int8_float_matrix_elements, checked_int8_float_size_t, add_int8_float_workspace_elements) are a mouthful. They're clear, so no objection — but they're local to one type combination and could lose the int8_float prefix inside an anonymous namespace.

zjin-lcf and others added 3 commits September 18, 2026 15:53
uxlfoundation#761 has landed, so this keeps the rocBLAS int8-to-float gemm_batch work
and the large-prime range regression on top of that test harness.

Co-authored-by: Cursor <cursoragent@cursor.com>
The same gemm_batch call is unimplemented on rocBLAS once k exceeds 131071,
while cuBLAS accumulates in float with no such ceiling. Record that in the
developer reference and apply a per-backend error model so the integer path
is not judged against the looser fp32 bound.

Co-authored-by: Cursor <cursoragent@cursor.com>
@zjin-lcf
zjin-lcf requested a review from a team as a code owner September 18, 2026 23:08
@zjin-lcf

Copy link
Copy Markdown
Contributor Author

Thanks for the careful review.

Rebase onto #761. That PR is already on develop. I merged current develop into this branch instead of rebasing, so the unique diff is the rocBLAS backend, the large-prime USM regression, the BLAS limitation note, and the AMD error-model change. The cuBLAS five-line change is no longer in the reviewable delta.

k ≤ 131071 as a user-visible backend difference. Documented in docs/domains/blas.rst (Developer Reference) and called out in the exception-path comment in check_int8_float_accumulation_size. The PR description now states that the same call succeeds on NVIDIA and is unimplemented on AMD past that limit.

Test tolerance. The generic abs_bound of eps × |alpha| × k × 128² stays, because those tests compare against CBLAS (fp32). Int8Int8SinglePrecisionErrorModel now branches on AMD_ID: rocBLAS is judged against a few ulps of |alpha * dot| + |beta * C| from an exact integer reference, not the fp32 reduction budget.

Minor. Left the alpha == 0 gemm and the int8_float helper names as they are; both were optional.

Int8 unit tests still abort on this DPC++ checkout (ProgramManager::getDeviceKernelInfo when the kernel is in the backend .so). Direct compilation of rocblas_batch.cpp into a standalone driver on MI210 covers the strided alpha/beta combinations, the k-limit rejection, the large-prime flat launch, and a 512³×4 timing run (~143 GFLOP/s).

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants