Skip to content

Commit be668c6

Browse files
qutaoclaude
andcommitted
Scope shared-memory bit-exact sizing to packed scalar NVFP4
The exact-bits shared storage sizing previously applied to every dtype with element_bits % 8 != 0 (int4/uint4/uint2/...), halving e.g. int4 shared allocations while SharedByteOffsetToLogicalIndexOffset kept the legacy bytes-per-element conversion for those dtypes — an inconsistent size/index pair for any non-FP4 sub-byte buffer in a merged arena. Restrict the new sizing to float4_e2m1fn scalar so it mirrors the byte-offset -> logical-index special case exactly: NVFP4 buffers get packed two-per-byte sizing and indexing, every other dtype keeps bit-exact upstream behavior. This pass now changes nothing outside the NVFP4 path. Validated after rebuilding libtilelang.so: nvf4 language + example CLI + quantize layout + access_ptr codegen + atom mma tests (129 passed), example and WS benchmark --verify pass with a fresh JIT cache. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
1 parent 3e4984e commit be668c6

1 file changed

Lines changed: 18 additions & 7 deletions

File tree

‎src/transform/merge_shared_memory_allocations.cc‎

Lines changed: 18 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -83,19 +83,30 @@ static DataType GetStorageSizeExprDType(const Buffer &buffer) {
8383
return size_dtype;
8484
}
8585

86+
// Scalar packed NVFP4 stores two logical elements per byte, and its buffer
87+
// shapes are expressed in logical elements. It is the only dtype whose
88+
// storage sizing and byte-offset -> logical-index conversion deviate from the
89+
// legacy DataType::bytes() rule; all other (sub-)byte dtypes keep the
90+
// upstream semantics so this pass changes behavior for NVFP4 buffers only.
91+
static bool IsPackedScalarFp4(DataType dtype) {
92+
return dtype.is_float4_e2m1fn() && dtype.is_scalar();
93+
}
94+
8695
static int64_t GetSharedStorageBitsPerLogicalElement(DataType dtype) {
87-
return static_cast<int64_t>(dtype.bits()) * dtype.lanes();
96+
if (IsPackedScalarFp4(dtype)) {
97+
return 4;
98+
}
99+
return static_cast<int64_t>(dtype.bytes()) * dtype.lanes() * 8;
88100
}
89101

90102
static PrimExpr GetSharedStorageSizeBytes(const Buffer &buffer) {
91103
DataType size_dtype = GetStorageSizeExprDType(buffer);
92104
int64_t element_bits = GetSharedStorageBitsPerLogicalElement(buffer->dtype);
93105

94-
// Buffer shapes are expressed in logical elements. Do not use
95-
// DataType::bytes() here: it rounds sub-byte scalar types up per element, so
96-
// scalar packed FP4 would be charged as one byte per value instead of two
97-
// values per byte. Compute total bits first, then round the whole allocation
98-
// up to bytes.
106+
// Do not use DataType::bytes() for packed FP4: it rounds sub-byte scalar
107+
// types up per element, so scalar packed FP4 would be charged as one byte
108+
// per value instead of two values per byte. Compute total bits first, then
109+
// round the whole allocation up to bytes.
99110
PrimExpr size_bits = make_const(size_dtype, element_bits);
100111
for (const PrimExpr &extent : buffer->shape) {
101112
PrimExpr e = extent;
@@ -118,7 +129,7 @@ static PrimExpr SharedByteOffsetToLogicalIndexOffset(PrimExpr byte_offset,
118129
// the non-alias rewrite still indexes the original typed buffer. Convert the
119130
// byte offset back to that buffer's logical element index. Packed scalar
120131
// NVFP4 has two logical elements per byte.
121-
if (dtype.is_float4_e2m1fn() && dtype.is_scalar()) {
132+
if (IsPackedScalarFp4(dtype)) {
122133
return byte_offset * make_const(byte_offset.dtype(), 2);
123134
}
124135
return indexdiv(byte_offset, dtype.bytes() * dtype.lanes());

0 commit comments

Comments
 (0)