Conversation
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
left a comment
There was a problem hiding this comment.
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 == 0you setaccumulate = 0but still run the full batched gemm, whose result is then discarded. An early-out that only scalesCbybetawould 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 theint8_floatprefix inside an anonymous namespace.
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>
|
Thanks for the careful review. Rebase onto #761. That PR is already on k ≤ 131071 as a user-visible backend difference. Documented in Test tolerance. The generic Minor. Left the Int8 unit tests still abort on this DPC++ checkout ( |
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.
kis bounded so the int32 accumulator cannot wrap. Every int8 magnitude is at most 128,so
128 * 128bounds a product andkaboveINT32_MAX / (128 * 128)(k > 131071)reports
unimplemented. The cuBLAS backend accumulates this combination in float and hasno such ceiling, so the same call succeeds on NVIDIA and throws on AMD past that limit.
This is documented in
docs/domains/blas.rst.allocated, using checked multiplication and addition rather than a check on the product.
ldcbelow the row count is rejected. rocBLAS rejects such a call itself, but onlyfrom 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.
locating the group that owns each entry with a binary search over per-group metadata.
multi-dimensional range onto its grid. The added
Int8Int8SinglePrecisionLargePrimeRangeregression holds a large prime in
nto cover that.#761is already ondevelop; this branch is merged onto currentdeveloprather thanrebased, 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++
ProgramManagerkernel-map assertion whenlaunching SYCL kernels from a backend
.so) passes in both dispatch modes:Int8 coverage is from a standalone program compiled together with
rocblas_batch.cppsothe scaling kernels live in the same image:
Made with Cursor